modified: Figure_1.png
modified: checkpoints/forward/best_tcn_model.pt modified: checkpoints/forward/training_history.csv new file: checkpoints_mlp/task1_feature_mlp/best_feature_mlp.pt new file: checkpoints_mlp/task1_feature_mlp/training_history.csv new file: checkpoints_rms/forward_rms/best_rms_model.pt new file: checkpoints_rms/forward_rms/training_history.csv modified: evaluation_outputs/forward/evaluation_forward_test_b0_s0.png new file: evaluation_outputs/forward/evaluation_forward_val_b0_s0.png new file: evaluation_outputs/forward/evaluation_forward_val_b3_s0.png new file: evaluation_outputs/forward_rms/evaluation_test_all_samples.csv new file: evaluation_outputs/forward_rms/evaluation_test_s0.png new file: evaluation_outputs/forward_rms/evaluation_train_all_samples.csv new file: evaluation_outputs/forward_rms/evaluation_val_all_samples.csv new file: evaluation_outputs/forward_rms/evaluation_val_s0.png new file: evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png new file: evaluation_outputs/task1_feature_mlp/evaluation_train_all_samples.csv new file: evaluation_outputs/task1_feature_mlp/evaluation_train_curve.png new file: evaluation_outputs/task1_feature_mlp/evaluation_val_all_samples.csv new file: evaluation_outputs/task1_feature_mlp/evaluation_val_curve.png new file: evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png new file: evaluation_outputs/task1_feature_mlp/harmonic_5mm_0.75Hz_prediction.png new file: evaluation_outputs/task1_feature_mlp/harmonic_5mm_1.55Hz_prediction.png new file: scripts/__pycache__/config.cpython-310.pyc new file: scripts/__pycache__/dataset.cpython-310.pyc new file: scripts/__pycache__/model.cpython-310.pyc new file: scripts/config.py new file: scripts/dataset.py new file: scripts/evaluate.py new file: scripts/model.py new file: scripts/predict_single.py new file: scripts/train.py modified: src/__pycache__/config.cpython-310.pyc modified: src/__pycache__/dataset.cpython-310.pyc modified: src/__pycache__/model.cpython-310.pyc modified: src/config.py modified: src/dataset.py modified: src/model.py new file: src_new/__pycache__/config.cpython-310.pyc new file: src_new/__pycache__/dataset.cpython-310.pyc new file: src_new/__pycache__/evaluate.cpython-310.pyc new file: src_new/__pycache__/model.cpython-310.pyc new file: src_new/__pycache__/train.cpython-310.pyc new file: src_new/config.py new file: src_new/dataset.py new file: src_new/evaluate.py new file: src_new/model.py new file: src_new/train.py
This commit is contained in:
BIN
scripts/__pycache__/config.cpython-310.pyc
Normal file
BIN
scripts/__pycache__/config.cpython-310.pyc
Normal file
Binary file not shown.
BIN
scripts/__pycache__/dataset.cpython-310.pyc
Normal file
BIN
scripts/__pycache__/dataset.cpython-310.pyc
Normal file
Binary file not shown.
BIN
scripts/__pycache__/model.cpython-310.pyc
Normal file
BIN
scripts/__pycache__/model.cpython-310.pyc
Normal file
Binary file not shown.
109
scripts/config.py
Normal file
109
scripts/config.py
Normal file
@@ -0,0 +1,109 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
CORE_FEATURE_NAMES: tuple[str, ...] = (
|
||||
"dominant_frequency_hz",
|
||||
"frequency_squared",
|
||||
"inverse_frequency_hz",
|
||||
"log_frequency_hz",
|
||||
"input_rms",
|
||||
"input_peak_abs",
|
||||
"input_peak_to_peak",
|
||||
"crest_factor",
|
||||
"middle_length_ratio",
|
||||
"dominant_amplitude",
|
||||
"dominant_energy_ratio",
|
||||
"harmonic_fit_amplitude",
|
||||
"harmonic_fit_residual_ratio",
|
||||
"spectral_peak_prominence",
|
||||
"half_power_bandwidth_hz",
|
||||
"spectral_centroid_hz",
|
||||
"signal_mean",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DataConfig:
|
||||
project_root: Path = field(default_factory=lambda: Path(__file__).resolve().parents[1])
|
||||
scenario: str = "Non_TMD"
|
||||
train_split_name: str = "train"
|
||||
val_split_name: str = "val"
|
||||
test_split_name: str = "test"
|
||||
csv_pattern: str = "*.csv"
|
||||
|
||||
code_column: str = "code"
|
||||
time_column: str = "time"
|
||||
base_sensor_code: str = "WSMS00012"
|
||||
base_axis: str = "value1"
|
||||
response_sensor_code: str = "WSMS00007"
|
||||
response_axis: str = "value3"
|
||||
|
||||
middle_segment_start_ratio: float = 0.20
|
||||
middle_segment_end_ratio: float = 0.80
|
||||
min_segment_length: int = 512
|
||||
steady_window_ratio: float = 0.25
|
||||
steady_window_stride_ratio: float = 0.05
|
||||
stability_subwindow_count: int = 4
|
||||
interpolation_method: str = "linear"
|
||||
normalization_eps: float = 1e-6
|
||||
|
||||
downloads_dir: Path = field(init=False)
|
||||
scenario_dir: Path = field(init=False)
|
||||
train_dir: Path = field(init=False)
|
||||
val_dir: Path = field(init=False)
|
||||
test_dir: Path = field(init=False)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.project_root = Path(self.project_root).resolve()
|
||||
self.downloads_dir = self.project_root / "downloads"
|
||||
self.scenario_dir = self.downloads_dir / self.scenario
|
||||
self.train_dir = self.scenario_dir / self.train_split_name
|
||||
self.val_dir = self.scenario_dir / self.val_split_name
|
||||
self.test_dir = self.scenario_dir / self.test_split_name
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelConfig:
|
||||
input_dim: int = len(CORE_FEATURE_NAMES)
|
||||
hidden_dims: tuple[int, ...] = (96, 64, 32)
|
||||
dropout: float = 0.08
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrainConfig:
|
||||
epochs: int = 400
|
||||
batch_size: int = 16
|
||||
learning_rate: float = 1e-3
|
||||
weight_decay: float = 1e-4
|
||||
seed: int = 42
|
||||
device: str = "cuda"
|
||||
grad_clip_norm: float = 1.0
|
||||
lr_scheduler_patience: int = 20
|
||||
lr_scheduler_factor: float = 0.5
|
||||
min_learning_rate: float = 1e-6
|
||||
early_stop_patience: int = 50
|
||||
checkpoint_dir: str = "checkpoints_mlp"
|
||||
history_name: str = "training_history.csv"
|
||||
best_model_name: str = "best_feature_mlp.pt"
|
||||
|
||||
|
||||
@dataclass
|
||||
class LossConfig:
|
||||
relative_rms_weight: float = 1.0
|
||||
log_rms_huber_weight: float = 0.75
|
||||
mae_weight: float = 0.15
|
||||
|
||||
|
||||
@dataclass
|
||||
class ExperimentConfig:
|
||||
data: DataConfig = field(default_factory=DataConfig)
|
||||
model: ModelConfig = field(default_factory=ModelConfig)
|
||||
train: TrainConfig = field(default_factory=TrainConfig)
|
||||
loss: LossConfig = field(default_factory=LossConfig)
|
||||
|
||||
|
||||
def make_experiment_config() -> ExperimentConfig:
|
||||
return ExperimentConfig()
|
||||
478
scripts/dataset.py
Normal file
478
scripts/dataset.py
Normal file
@@ -0,0 +1,478 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
|
||||
try:
|
||||
from .config import CORE_FEATURE_NAMES, DataConfig, ExperimentConfig
|
||||
except ImportError:
|
||||
from config import CORE_FEATURE_NAMES, DataConfig, ExperimentConfig
|
||||
|
||||
|
||||
def calculate_rms(signal: np.ndarray) -> float:
|
||||
signal = np.asarray(signal, dtype=np.float64).reshape(-1)
|
||||
return float(np.sqrt(np.mean(np.square(signal))))
|
||||
|
||||
|
||||
def extract_frequency_hz(file_name: str) -> float | None:
|
||||
match = re.search(r"(\d+(?:\.\d+)?)Hz", file_name, flags=re.IGNORECASE)
|
||||
if match is None:
|
||||
return None
|
||||
return float(match.group(1))
|
||||
|
||||
|
||||
def estimate_sampling_rate(time_values: np.ndarray) -> float:
|
||||
if time_values.size < 2:
|
||||
return 100.0
|
||||
dt = np.diff(time_values)
|
||||
dt = dt[np.isfinite(dt)]
|
||||
dt = dt[dt > 0.0]
|
||||
if dt.size == 0:
|
||||
return 100.0
|
||||
return float(1.0 / np.median(dt))
|
||||
|
||||
|
||||
def get_middle_segment(
|
||||
time_values: np.ndarray,
|
||||
x_values: np.ndarray,
|
||||
y_values: np.ndarray,
|
||||
config: DataConfig,
|
||||
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
|
||||
length = len(x_values)
|
||||
search_start = int(length * config.middle_segment_start_ratio)
|
||||
search_end = int(length * config.middle_segment_end_ratio)
|
||||
search_start = max(0, min(search_start, length - 1))
|
||||
search_end = max(search_start + 1, min(search_end, length))
|
||||
|
||||
if search_end - search_start < config.min_segment_length:
|
||||
center = length // 2
|
||||
half = config.min_segment_length // 2
|
||||
search_start = max(0, center - half)
|
||||
search_end = min(length, search_start + config.min_segment_length)
|
||||
search_start = max(0, search_end - config.min_segment_length)
|
||||
|
||||
search_length = search_end - search_start
|
||||
window_length = max(config.min_segment_length, int(length * config.steady_window_ratio))
|
||||
window_length = min(window_length, search_length)
|
||||
if search_length <= window_length:
|
||||
return (
|
||||
time_values[search_start:search_end],
|
||||
x_values[search_start:search_end],
|
||||
y_values[search_start:search_end],
|
||||
)
|
||||
|
||||
stride = max(1, int(length * config.steady_window_stride_ratio))
|
||||
candidate_y = y_values[search_start:search_end]
|
||||
candidate_rms = calculate_rms(candidate_y)
|
||||
best_score = float("inf")
|
||||
best_slice = slice(search_start, search_start + window_length)
|
||||
|
||||
for window_start in range(search_start, search_end - window_length + 1, stride):
|
||||
window_end = window_start + window_length
|
||||
window_y = y_values[window_start:window_end]
|
||||
split_windows = np.array_split(window_y, config.stability_subwindow_count)
|
||||
split_rms = np.asarray([calculate_rms(chunk) for chunk in split_windows], dtype=np.float64)
|
||||
split_peaks = np.asarray([float(np.max(np.abs(chunk))) for chunk in split_windows], dtype=np.float64)
|
||||
|
||||
rms_cv = float(split_rms.std() / max(split_rms.mean(), config.normalization_eps))
|
||||
peak_cv = float(split_peaks.std() / max(split_peaks.mean(), config.normalization_eps))
|
||||
window_rms = calculate_rms(window_y)
|
||||
|
||||
# Prefer windows that are stable inside the middle candidate region while keeping enough energy.
|
||||
score = rms_cv + 0.35 * peak_cv - 0.05 * (window_rms / max(candidate_rms, config.normalization_eps))
|
||||
if score < best_score:
|
||||
best_score = score
|
||||
best_slice = slice(window_start, window_end)
|
||||
|
||||
return time_values[best_slice], x_values[best_slice], y_values[best_slice]
|
||||
|
||||
|
||||
def _compute_windowed_spectrum(signal: np.ndarray, sampling_rate: float) -> tuple[np.ndarray, np.ndarray]:
|
||||
signal = np.asarray(signal, dtype=np.float64).reshape(-1)
|
||||
if signal.size < 4:
|
||||
return np.asarray([], dtype=np.float64), np.asarray([], dtype=np.float64)
|
||||
centered = signal - np.mean(signal)
|
||||
window = np.hanning(signal.size)
|
||||
scale = max(np.sum(window), 1e-12)
|
||||
fft_values = np.fft.rfft(centered * window)
|
||||
freqs = np.fft.rfftfreq(signal.size, d=1.0 / sampling_rate)
|
||||
magnitudes = (2.0 / scale) * np.abs(fft_values)
|
||||
return freqs, magnitudes
|
||||
|
||||
|
||||
def _parabolic_peak_frequency(freqs: np.ndarray, magnitudes: np.ndarray, peak_index: int) -> float:
|
||||
if peak_index <= 0 or peak_index >= magnitudes.size - 1:
|
||||
return float(freqs[peak_index])
|
||||
alpha = magnitudes[peak_index - 1]
|
||||
beta = magnitudes[peak_index]
|
||||
gamma = magnitudes[peak_index + 1]
|
||||
denominator = alpha - 2.0 * beta + gamma
|
||||
if abs(denominator) < 1e-12:
|
||||
return float(freqs[peak_index])
|
||||
offset = 0.5 * (alpha - gamma) / denominator
|
||||
bin_width = float(freqs[1] - freqs[0])
|
||||
return float(freqs[peak_index] + offset * bin_width)
|
||||
|
||||
|
||||
def get_dominant_frequency(signal: np.ndarray, sampling_rate: float) -> float:
|
||||
freqs, magnitudes = _compute_windowed_spectrum(signal, sampling_rate)
|
||||
if magnitudes.size == 0:
|
||||
return 0.0
|
||||
magnitudes[0] = 0.0
|
||||
band_mask = (freqs >= 0.1) & (freqs <= 5.0)
|
||||
if not np.any(band_mask):
|
||||
return 0.0
|
||||
band_magnitudes = np.where(band_mask, magnitudes, 0.0)
|
||||
peak_index = int(np.argmax(band_magnitudes))
|
||||
return _parabolic_peak_frequency(freqs, magnitudes, peak_index)
|
||||
|
||||
|
||||
def compute_harmonic_fit_features(signal: np.ndarray, sampling_rate: float, frequency_hz: float) -> tuple[float, float]:
|
||||
signal = np.asarray(signal, dtype=np.float64).reshape(-1)
|
||||
if signal.size < 4 or frequency_hz <= 0.0:
|
||||
return 0.0, 1.0
|
||||
time_axis = np.arange(signal.size, dtype=np.float64) / max(sampling_rate, 1e-12)
|
||||
omega_t = 2.0 * np.pi * frequency_hz * time_axis
|
||||
design = np.stack([np.sin(omega_t), np.cos(omega_t), np.ones_like(omega_t)], axis=1)
|
||||
coefficients, _, _, _ = np.linalg.lstsq(design, signal, rcond=None)
|
||||
fitted = design @ coefficients
|
||||
harmonic_amplitude = float(np.sqrt(coefficients[0] ** 2 + coefficients[1] ** 2))
|
||||
residual = signal - fitted
|
||||
residual_ratio = calculate_rms(residual) / max(calculate_rms(signal), 1e-6)
|
||||
return harmonic_amplitude, float(residual_ratio)
|
||||
|
||||
|
||||
def compute_spectral_features(signal: np.ndarray, sampling_rate: float) -> tuple[float, float, float, float, float]:
|
||||
freqs, amplitudes = _compute_windowed_spectrum(signal, sampling_rate)
|
||||
if amplitudes.size == 0:
|
||||
return 0.0, 0.0, 0.0, 0.0, 0.0
|
||||
|
||||
powers = np.square(amplitudes)
|
||||
amplitudes[0] = 0.0
|
||||
powers[0] = 0.0
|
||||
band_mask = (freqs >= 0.1) & (freqs <= 5.0)
|
||||
if not np.any(band_mask):
|
||||
return 0.0, 0.0, 0.0, 0.0, 0.0
|
||||
|
||||
band_amplitudes = np.where(band_mask, amplitudes, 0.0)
|
||||
dominant_index = int(np.argmax(band_amplitudes))
|
||||
dominant_amplitude = float(amplitudes[dominant_index])
|
||||
dominant_freq = _parabolic_peak_frequency(freqs, amplitudes, dominant_index)
|
||||
|
||||
local_mask = np.abs(freqs - dominant_freq) <= 0.10
|
||||
dominant_energy = float(powers[local_mask].sum())
|
||||
total_band_energy = float(powers[band_mask].sum())
|
||||
dominant_energy_ratio = dominant_energy / max(total_band_energy, 1e-12)
|
||||
spectral_centroid = float((freqs[band_mask] * powers[band_mask]).sum() / max(total_band_energy, 1e-12))
|
||||
|
||||
background_mask = band_mask & (~local_mask)
|
||||
background_level = float(np.median(amplitudes[background_mask])) if np.any(background_mask) else 0.0
|
||||
spectral_peak_prominence = dominant_amplitude / max(background_level, 1e-6)
|
||||
|
||||
half_power_level = dominant_amplitude / np.sqrt(2.0)
|
||||
left_index = dominant_index
|
||||
right_index = dominant_index
|
||||
while left_index > 0 and amplitudes[left_index] >= half_power_level:
|
||||
left_index -= 1
|
||||
while right_index < amplitudes.size - 1 and amplitudes[right_index] >= half_power_level:
|
||||
right_index += 1
|
||||
half_power_bandwidth = float(freqs[right_index] - freqs[left_index]) if right_index > left_index else 0.0
|
||||
return dominant_amplitude, dominant_energy_ratio, spectral_centroid, spectral_peak_prominence, half_power_bandwidth
|
||||
|
||||
|
||||
def extract_core_features(
|
||||
signal: np.ndarray,
|
||||
sampling_rate: float,
|
||||
original_length: int,
|
||||
config: DataConfig,
|
||||
) -> np.ndarray:
|
||||
x_rms = calculate_rms(signal)
|
||||
x_peak = float(np.max(np.abs(signal))) if signal.size > 0 else 0.0
|
||||
x_peak_to_peak = float(np.max(signal) - np.min(signal)) if signal.size > 0 else 0.0
|
||||
crest_factor = x_peak / max(x_rms, config.normalization_eps)
|
||||
dominant_frequency = get_dominant_frequency(signal, sampling_rate)
|
||||
frequency_squared = dominant_frequency * dominant_frequency
|
||||
inverse_frequency = 1.0 / max(dominant_frequency, 1e-6)
|
||||
log_frequency = float(np.log(max(dominant_frequency, 1e-6)))
|
||||
middle_length_ratio = float(signal.size) / float(max(original_length, 1))
|
||||
(
|
||||
dominant_amplitude,
|
||||
dominant_energy_ratio,
|
||||
spectral_centroid,
|
||||
spectral_peak_prominence,
|
||||
half_power_bandwidth,
|
||||
) = compute_spectral_features(signal, sampling_rate)
|
||||
harmonic_fit_amplitude, harmonic_fit_residual_ratio = compute_harmonic_fit_features(
|
||||
signal,
|
||||
sampling_rate,
|
||||
dominant_frequency,
|
||||
)
|
||||
signal_mean = float(np.mean(signal)) if signal.size > 0 else 0.0
|
||||
return np.asarray(
|
||||
[
|
||||
dominant_frequency,
|
||||
frequency_squared,
|
||||
inverse_frequency,
|
||||
log_frequency,
|
||||
x_rms,
|
||||
x_peak,
|
||||
x_peak_to_peak,
|
||||
crest_factor,
|
||||
middle_length_ratio,
|
||||
dominant_amplitude,
|
||||
dominant_energy_ratio,
|
||||
harmonic_fit_amplitude,
|
||||
harmonic_fit_residual_ratio,
|
||||
spectral_peak_prominence,
|
||||
half_power_bandwidth,
|
||||
spectral_centroid,
|
||||
signal_mean,
|
||||
],
|
||||
dtype=np.float32,
|
||||
)
|
||||
|
||||
|
||||
def _value_frame(df: pd.DataFrame, config: DataConfig, sensor_code: str, value_column: str) -> pd.DataFrame:
|
||||
sensor_df = df.loc[df[config.code_column] == sensor_code, [config.time_column, value_column]].copy()
|
||||
sensor_df = sensor_df.sort_values(config.time_column)
|
||||
sensor_df = sensor_df.drop_duplicates(subset=config.time_column, keep="first")
|
||||
sensor_df[config.time_column] = sensor_df[config.time_column].astype("float64")
|
||||
sensor_df[value_column] = sensor_df[value_column].astype("float32")
|
||||
return sensor_df
|
||||
|
||||
|
||||
def load_aligned_signals(file_path: Path, config: DataConfig) -> tuple[np.ndarray, np.ndarray, np.ndarray, int]:
|
||||
df = pd.read_csv(file_path)
|
||||
base_df = _value_frame(df, config, config.base_sensor_code, config.base_axis)
|
||||
response_df = _value_frame(df, config, config.response_sensor_code, config.response_axis)
|
||||
if base_df.empty or response_df.empty:
|
||||
raise ValueError(f"Missing required sensor in {file_path.name}")
|
||||
|
||||
aligned = base_df.rename(columns={config.base_axis: "base_signal"}).merge(
|
||||
response_df.rename(columns={config.response_axis: "response_signal"}),
|
||||
on=config.time_column,
|
||||
how="left",
|
||||
)
|
||||
interpolation_count = int(aligned["response_signal"].isna().sum())
|
||||
aligned["response_signal"] = aligned["response_signal"].interpolate(
|
||||
method=config.interpolation_method,
|
||||
limit_direction="both",
|
||||
).ffill().bfill()
|
||||
if aligned["response_signal"].isna().any():
|
||||
raise ValueError(f"Remaining NaN after interpolation in {file_path.name}")
|
||||
|
||||
time_values = aligned[config.time_column].to_numpy(dtype=np.float64)
|
||||
x_values = aligned["base_signal"].to_numpy(dtype=np.float32)
|
||||
y_values = aligned["response_signal"].to_numpy(dtype=np.float32)
|
||||
return time_values, x_values, y_values, interpolation_count
|
||||
|
||||
|
||||
@dataclass
|
||||
class SampleRecord:
|
||||
file_path: Path
|
||||
split: str
|
||||
frequency_hz: float
|
||||
features: torch.Tensor
|
||||
x_rms: torch.Tensor
|
||||
y_rms: torch.Tensor
|
||||
target_y_rms: torch.Tensor
|
||||
target_log_y_rms: torch.Tensor
|
||||
time_middle: torch.Tensor
|
||||
x_middle: torch.Tensor
|
||||
y_middle: torch.Tensor
|
||||
sampling_rate: float
|
||||
interpolation_count: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class SplitLoadReport:
|
||||
split: str
|
||||
loaded_files: list[str] = field(default_factory=list)
|
||||
skipped_files: list[tuple[str, str]] = field(default_factory=list)
|
||||
interpolated_files: dict[str, int] = field(default_factory=dict)
|
||||
|
||||
def to_lines(self) -> list[str]:
|
||||
lines = [f"[{self.split}] loaded={len(self.loaded_files)} skipped={len(self.skipped_files)}"]
|
||||
for file_name, reason in self.skipped_files:
|
||||
lines.append(f" - skipped {file_name}: {reason}")
|
||||
for file_name, count in self.interpolated_files.items():
|
||||
if count > 0:
|
||||
lines.append(f" - interpolated {file_name}: missing_points={count}")
|
||||
return lines
|
||||
|
||||
|
||||
@dataclass
|
||||
class NormalizationStats:
|
||||
feature_mean: torch.Tensor
|
||||
feature_std: torch.Tensor
|
||||
target_mean: torch.Tensor
|
||||
target_std: torch.Tensor
|
||||
|
||||
|
||||
class FeatureDataset(Dataset):
|
||||
def __init__(self, records: list[SampleRecord], normalization: NormalizationStats) -> None:
|
||||
self.records = records
|
||||
self.normalization = normalization
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.records)
|
||||
|
||||
def __getitem__(self, index: int) -> dict[str, Any]:
|
||||
record = self.records[index]
|
||||
feature_norm = (record.features - self.normalization.feature_mean) / self.normalization.feature_std
|
||||
target_norm = (record.target_y_rms - self.normalization.target_mean) / self.normalization.target_std
|
||||
return {
|
||||
"x": feature_norm,
|
||||
"target": target_norm,
|
||||
"target_raw": record.target_y_rms,
|
||||
"target_log_raw": record.target_log_y_rms,
|
||||
"x_rms_raw": record.x_rms,
|
||||
"y_rms_raw": record.y_rms,
|
||||
"frequency_hz": torch.tensor(record.frequency_hz, dtype=torch.float32),
|
||||
"features_raw": record.features,
|
||||
"file_name": record.file_path.name,
|
||||
}
|
||||
|
||||
|
||||
def list_split_files(config: DataConfig, split: str) -> list[Path]:
|
||||
split_dir = getattr(config, f"{split}_dir")
|
||||
return sorted(split_dir.glob(config.csv_pattern))
|
||||
|
||||
|
||||
def build_record_from_file(file_path: Path, split: str, config: DataConfig) -> SampleRecord:
|
||||
time_values, x_values, y_values, interpolation_count = load_aligned_signals(file_path, config)
|
||||
time_middle, x_middle, y_middle = get_middle_segment(time_values, x_values, y_values, config)
|
||||
sampling_rate = estimate_sampling_rate(time_middle)
|
||||
frequency_hz = extract_frequency_hz(file_path.name)
|
||||
if frequency_hz is None:
|
||||
frequency_hz = get_dominant_frequency(x_middle, sampling_rate)
|
||||
|
||||
x_rms = calculate_rms(x_middle)
|
||||
y_rms = calculate_rms(y_middle)
|
||||
transmission_ratio = float(y_rms / max(x_rms, config.normalization_eps))
|
||||
target_y_rms = transmission_ratio
|
||||
target_log_y_rms = float(np.log(max(transmission_ratio, config.normalization_eps)))
|
||||
features = extract_core_features(x_middle, sampling_rate, len(x_values), config)
|
||||
|
||||
return SampleRecord(
|
||||
file_path=file_path,
|
||||
split=split,
|
||||
frequency_hz=frequency_hz,
|
||||
features=torch.tensor(features, dtype=torch.float32),
|
||||
x_rms=torch.tensor([x_rms], dtype=torch.float32),
|
||||
y_rms=torch.tensor([y_rms], dtype=torch.float32),
|
||||
target_y_rms=torch.tensor([target_y_rms], dtype=torch.float32),
|
||||
target_log_y_rms=torch.tensor([target_log_y_rms], dtype=torch.float32),
|
||||
time_middle=torch.tensor(time_middle, dtype=torch.float64),
|
||||
x_middle=torch.tensor(x_middle[:, None], dtype=torch.float32),
|
||||
y_middle=torch.tensor(y_middle[:, None], dtype=torch.float32),
|
||||
sampling_rate=sampling_rate,
|
||||
interpolation_count=interpolation_count,
|
||||
)
|
||||
|
||||
|
||||
def load_split_records(config: DataConfig, split: str) -> tuple[list[SampleRecord], SplitLoadReport]:
|
||||
records: list[SampleRecord] = []
|
||||
report = SplitLoadReport(split=split)
|
||||
for file_path in list_split_files(config, split):
|
||||
try:
|
||||
record = build_record_from_file(file_path, split, config)
|
||||
except ValueError as error:
|
||||
report.skipped_files.append((file_path.name, str(error)))
|
||||
continue
|
||||
records.append(record)
|
||||
report.loaded_files.append(file_path.name)
|
||||
if record.interpolation_count > 0:
|
||||
report.interpolated_files[file_path.name] = record.interpolation_count
|
||||
|
||||
if not records:
|
||||
raise RuntimeError(f"No usable records found for split='{split}'.")
|
||||
return records, report
|
||||
|
||||
|
||||
def fit_normalization(records: list[SampleRecord], config: DataConfig) -> NormalizationStats:
|
||||
feature_all = torch.stack([record.features for record in records], dim=0)
|
||||
target_all = torch.cat([record.target_y_rms for record in records], dim=0)
|
||||
|
||||
feature_std = torch.clamp(feature_all.std(dim=0, unbiased=False), min=config.normalization_eps)
|
||||
target_std = torch.clamp(target_all.std(dim=0, unbiased=False).view(1), min=config.normalization_eps)
|
||||
|
||||
return NormalizationStats(
|
||||
feature_mean=feature_all.mean(dim=0),
|
||||
feature_std=feature_std,
|
||||
target_mean=target_all.mean(dim=0, keepdim=True),
|
||||
target_std=target_std,
|
||||
)
|
||||
|
||||
|
||||
def build_datasets(
|
||||
config: ExperimentConfig | DataConfig,
|
||||
) -> tuple[dict[str, FeatureDataset], dict[str, list[SampleRecord]], dict[str, SplitLoadReport]]:
|
||||
data_config = config.data if isinstance(config, ExperimentConfig) else config
|
||||
train_records, train_report = load_split_records(data_config, "train")
|
||||
normalization = fit_normalization(train_records, data_config)
|
||||
val_records, val_report = load_split_records(data_config, "val")
|
||||
test_records, test_report = load_split_records(data_config, "test")
|
||||
|
||||
raw_records = {
|
||||
"train": train_records,
|
||||
"val": val_records,
|
||||
"test": test_records,
|
||||
}
|
||||
datasets = {
|
||||
split: FeatureDataset(records, normalization)
|
||||
for split, records in raw_records.items()
|
||||
}
|
||||
reports = {
|
||||
"train": train_report,
|
||||
"val": val_report,
|
||||
"test": test_report,
|
||||
}
|
||||
return datasets, raw_records, reports
|
||||
|
||||
|
||||
def build_dataloaders(
|
||||
config: ExperimentConfig | DataConfig,
|
||||
) -> tuple[dict[str, DataLoader], dict[str, FeatureDataset], dict[str, list[SampleRecord]], dict[str, SplitLoadReport]]:
|
||||
experiment_config = config if isinstance(config, ExperimentConfig) else ExperimentConfig(data=config)
|
||||
datasets, raw_records, reports = build_datasets(experiment_config)
|
||||
batch_size = experiment_config.train.batch_size
|
||||
loaders = {
|
||||
"train": DataLoader(datasets["train"], batch_size=batch_size, shuffle=True),
|
||||
"val": DataLoader(datasets["val"], batch_size=batch_size, shuffle=False),
|
||||
"test": DataLoader(datasets["test"], batch_size=batch_size, shuffle=False),
|
||||
}
|
||||
return loaders, datasets, raw_records, reports
|
||||
|
||||
|
||||
def report_to_text(reports: dict[str, SplitLoadReport]) -> str:
|
||||
lines: list[str] = []
|
||||
for split in ("train", "val", "test"):
|
||||
lines.extend(reports[split].to_lines())
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def checkpoint_normalization_payload(normalization: NormalizationStats) -> dict[str, list[float]]:
|
||||
return {
|
||||
"feature_mean": normalization.feature_mean.tolist(),
|
||||
"feature_std": normalization.feature_std.tolist(),
|
||||
"target_mean": normalization.target_mean.tolist(),
|
||||
"target_std": normalization.target_std.tolist(),
|
||||
"feature_names": list(CORE_FEATURE_NAMES),
|
||||
}
|
||||
|
||||
|
||||
def normalization_from_payload(payload: dict[str, Any]) -> NormalizationStats:
|
||||
return NormalizationStats(
|
||||
feature_mean=torch.tensor(payload["feature_mean"], dtype=torch.float32),
|
||||
feature_std=torch.tensor(payload["feature_std"], dtype=torch.float32),
|
||||
target_mean=torch.tensor(payload["target_mean"], dtype=torch.float32),
|
||||
target_std=torch.tensor(payload["target_std"], dtype=torch.float32),
|
||||
)
|
||||
224
scripts/evaluate.py
Normal file
224
scripts/evaluate.py
Normal file
@@ -0,0 +1,224 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
|
||||
try:
|
||||
from .config import CORE_FEATURE_NAMES, ExperimentConfig, make_experiment_config
|
||||
from .dataset import build_dataloaders, normalization_from_payload, report_to_text
|
||||
from .model import build_model
|
||||
except ImportError:
|
||||
from config import CORE_FEATURE_NAMES, ExperimentConfig, make_experiment_config
|
||||
from dataset import build_dataloaders, normalization_from_payload, report_to_text
|
||||
from model import build_model
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Evaluate feature MLP on train/val/test splits.")
|
||||
parser.add_argument("--split", choices=("train", "val", "test"), default="val")
|
||||
parser.add_argument("--sample-index", type=int, default=0)
|
||||
parser.add_argument("--all-samples", action="store_true")
|
||||
parser.add_argument("--device", type=str, default=None)
|
||||
parser.add_argument("--checkpoint", type=str, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def resolve_device(device_name: str | None, config: ExperimentConfig) -> torch.device:
|
||||
requested = device_name or config.train.device
|
||||
if requested.startswith("cuda") and not torch.cuda.is_available():
|
||||
return torch.device("cpu")
|
||||
return torch.device(requested)
|
||||
|
||||
|
||||
def resolve_checkpoint_path(config: ExperimentConfig, checkpoint_arg: str | None) -> Path:
|
||||
if checkpoint_arg:
|
||||
return Path(checkpoint_arg).resolve()
|
||||
return config.data.project_root / config.train.checkpoint_dir / "task1_feature_mlp" / config.train.best_model_name
|
||||
|
||||
|
||||
def relative_percent_error(true_value: float, pred_value: float) -> float:
|
||||
denominator = max(abs(true_value), 1e-12)
|
||||
return abs(pred_value - true_value) / denominator * 100.0
|
||||
|
||||
|
||||
def predict_rms(
|
||||
model: torch.nn.Module,
|
||||
feature_norm: torch.Tensor,
|
||||
x_rms_raw: torch.Tensor,
|
||||
target_mean: torch.Tensor,
|
||||
target_std: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
pred_norm = model(feature_norm)
|
||||
pred_tr = torch.clamp(pred_norm * target_std + target_mean, min=1e-6)
|
||||
return pred_tr * x_rms_raw
|
||||
|
||||
|
||||
def evaluate_sample(
|
||||
model: torch.nn.Module,
|
||||
dataset,
|
||||
raw_record,
|
||||
sample_index: int,
|
||||
device: torch.device,
|
||||
) -> dict[str, float | str]:
|
||||
sample = dataset[sample_index]
|
||||
normalization = dataset.normalization
|
||||
feature_norm = sample["x"].unsqueeze(0).to(device)
|
||||
x_rms_raw = sample["x_rms_raw"].unsqueeze(0).to(device)
|
||||
target_mean = normalization.target_mean.to(device).view(1, 1)
|
||||
target_std = normalization.target_std.to(device).view(1, 1)
|
||||
|
||||
with torch.no_grad():
|
||||
pred_rms = predict_rms(model, feature_norm, x_rms_raw, target_mean, target_std)
|
||||
|
||||
true_rms = float(sample["y_rms_raw"].item())
|
||||
pred_rms_value = float(pred_rms.detach().cpu().numpy().reshape(-1)[0])
|
||||
row = {
|
||||
"file_name": str(sample["file_name"]),
|
||||
"frequency_hz": float(sample["frequency_hz"].item()),
|
||||
"true_rms": true_rms,
|
||||
"pred_rms": pred_rms_value,
|
||||
"relative_error_percent": relative_percent_error(true_rms, pred_rms_value),
|
||||
}
|
||||
features_raw = sample["features_raw"].numpy().reshape(-1)
|
||||
for name, value in zip(CORE_FEATURE_NAMES, features_raw):
|
||||
row[name] = float(value)
|
||||
row["sampling_rate"] = float(raw_record.sampling_rate)
|
||||
return row
|
||||
|
||||
|
||||
def save_all_samples_plot(result_df: pd.DataFrame, split: str, save_dir: Path) -> Path | None:
|
||||
if result_df["frequency_hz"].isna().any():
|
||||
return None
|
||||
plot_df = result_df.sort_values("frequency_hz").reset_index(drop=True)
|
||||
fig, axes = plt.subplots(2, 1, figsize=(12, 8))
|
||||
fig.suptitle(f"Task1 Feature-MLP Evaluation | {split}")
|
||||
|
||||
axes[0].plot(plot_df["frequency_hz"], plot_df["true_rms"], marker="o", label="True RMS")
|
||||
axes[0].plot(plot_df["frequency_hz"], plot_df["pred_rms"], marker="o", label="Pred RMS")
|
||||
axes[0].set_xlabel("Frequency (Hz)")
|
||||
axes[0].set_ylabel("RMS")
|
||||
axes[0].grid(True, alpha=0.3)
|
||||
axes[0].legend()
|
||||
|
||||
axes[1].bar(plot_df["frequency_hz"].astype(str), plot_df["relative_error_percent"], color="tab:orange")
|
||||
axes[1].set_xlabel("Frequency (Hz)")
|
||||
axes[1].set_ylabel("Relative Error (%)")
|
||||
axes[1].grid(True, axis="y", alpha=0.3)
|
||||
axes[1].tick_params(axis="x", labelrotation=45)
|
||||
|
||||
plt.tight_layout()
|
||||
save_dir.mkdir(parents=True, exist_ok=True)
|
||||
figure_path = save_dir / f"evaluation_{split}_curve.png"
|
||||
plt.savefig(figure_path, dpi=180, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
return figure_path
|
||||
|
||||
|
||||
def save_single_sample_plot(raw_record, result: dict[str, float | str], split: str, save_dir: Path, sample_index: int) -> Path:
|
||||
time_middle = raw_record.time_middle.detach().cpu().numpy().reshape(-1)
|
||||
x_middle = raw_record.x_middle.detach().cpu().numpy().reshape(-1)
|
||||
y_middle = raw_record.y_middle.detach().cpu().numpy().reshape(-1)
|
||||
frequency_hz = float(result["frequency_hz"])
|
||||
pred_rms = float(result["pred_rms"])
|
||||
true_rms = float(result["true_rms"])
|
||||
|
||||
fig, axes = plt.subplots(3, 1, figsize=(12, 10))
|
||||
fig.suptitle(f"Task1 Feature-MLP | {split} | {result['file_name']}")
|
||||
|
||||
axes[0].plot(time_middle, x_middle, color="tab:blue")
|
||||
axes[0].set_title("Input Base Excitation (Middle Segment)")
|
||||
axes[0].set_xlabel("Time")
|
||||
axes[0].set_ylabel("Acceleration")
|
||||
axes[0].grid(True, alpha=0.3)
|
||||
|
||||
axes[1].plot(time_middle, y_middle, color="tab:green")
|
||||
axes[1].set_title("True Top Response (Middle Segment)")
|
||||
axes[1].set_xlabel("Time")
|
||||
axes[1].set_ylabel("Acceleration")
|
||||
axes[1].grid(True, alpha=0.3)
|
||||
|
||||
axes[2].bar(["True RMS", "Pred RMS"], [true_rms, pred_rms], color=["tab:green", "tab:orange"])
|
||||
axes[2].set_title(
|
||||
f"Freq: {frequency_hz:.4f} Hz | True RMS: {true_rms:.6f} | Pred RMS: {pred_rms:.6f} | "
|
||||
f"Error: {float(result['relative_error_percent']):.2f}%"
|
||||
)
|
||||
axes[2].set_ylabel("RMS")
|
||||
axes[2].grid(True, axis="y", alpha=0.3)
|
||||
|
||||
plt.tight_layout()
|
||||
save_dir.mkdir(parents=True, exist_ok=True)
|
||||
figure_path = save_dir / f"evaluation_{split}_s{sample_index}.png"
|
||||
plt.savefig(figure_path, dpi=180, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
return figure_path
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
config = make_experiment_config()
|
||||
device = resolve_device(args.device, config)
|
||||
checkpoint_path = resolve_checkpoint_path(config, args.checkpoint)
|
||||
if not checkpoint_path.exists():
|
||||
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
|
||||
|
||||
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
if "normalization" in checkpoint:
|
||||
normalization = normalization_from_payload(checkpoint["normalization"])
|
||||
else:
|
||||
normalization = None
|
||||
|
||||
loaders, datasets, raw_records, reports = build_dataloaders(config)
|
||||
del loaders
|
||||
print(report_to_text(reports))
|
||||
|
||||
dataset = datasets[args.split]
|
||||
if normalization is not None:
|
||||
dataset.normalization = normalization
|
||||
|
||||
model = build_model(config).to(device)
|
||||
model.load_state_dict(checkpoint["model_state_dict"])
|
||||
model.eval()
|
||||
|
||||
save_dir = config.data.project_root / "evaluation_outputs" / "task1_feature_mlp"
|
||||
if args.all_samples:
|
||||
rows = [
|
||||
evaluate_sample(model, dataset, raw_records[args.split][index], index, device)
|
||||
for index in range(len(dataset))
|
||||
]
|
||||
result_df = pd.DataFrame(rows).sort_values(["relative_error_percent", "file_name"]).reset_index(drop=True)
|
||||
save_dir.mkdir(parents=True, exist_ok=True)
|
||||
csv_path = save_dir / f"evaluation_{args.split}_all_samples.csv"
|
||||
result_df.to_csv(csv_path, index=False)
|
||||
figure_path = save_all_samples_plot(result_df, args.split, save_dir)
|
||||
summary = {
|
||||
"count": len(result_df),
|
||||
"mean_error_percent": float(result_df["relative_error_percent"].mean()),
|
||||
"median_error_percent": float(result_df["relative_error_percent"].median()),
|
||||
"max_error_percent": float(result_df["relative_error_percent"].max()),
|
||||
"min_error_percent": float(result_df["relative_error_percent"].min()),
|
||||
}
|
||||
print(f"Checkpoint: {checkpoint_path}")
|
||||
print(f"Summary: {summary}")
|
||||
print(f"CSV saved to: {csv_path}")
|
||||
if figure_path is not None:
|
||||
print(f"Figure saved to: {figure_path}")
|
||||
print(result_df.to_string(index=False))
|
||||
return
|
||||
|
||||
result = evaluate_sample(model, dataset, raw_records[args.split][args.sample_index], args.sample_index, device)
|
||||
figure_path = save_single_sample_plot(raw_records[args.split][args.sample_index], result, args.split, save_dir, args.sample_index)
|
||||
print(f"Checkpoint: {checkpoint_path}")
|
||||
print(f"Sample file: {result['file_name']}")
|
||||
print(f"True RMS: {float(result['true_rms']):.6f}")
|
||||
print(f"Pred RMS: {float(result['pred_rms']):.6f}")
|
||||
print(f"Relative RMS Error (%): {float(result['relative_error_percent']):.4f}")
|
||||
print(f"Figure saved to: {figure_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
31
scripts/model.py
Normal file
31
scripts/model.py
Normal file
@@ -0,0 +1,31 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
try:
|
||||
from .config import ExperimentConfig, ModelConfig
|
||||
except ImportError:
|
||||
from config import ExperimentConfig, ModelConfig
|
||||
|
||||
|
||||
class FeatureMLP(nn.Module):
|
||||
def __init__(self, config: ModelConfig) -> None:
|
||||
super().__init__()
|
||||
layers: list[nn.Module] = []
|
||||
input_dim = config.input_dim
|
||||
for hidden_dim in config.hidden_dims:
|
||||
layers.append(nn.Linear(input_dim, hidden_dim))
|
||||
layers.append(nn.GELU())
|
||||
layers.append(nn.Dropout(config.dropout))
|
||||
input_dim = hidden_dim
|
||||
layers.append(nn.Linear(input_dim, 1))
|
||||
self.network = nn.Sequential(*layers)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.network(x)
|
||||
|
||||
|
||||
def build_model(config: ExperimentConfig | ModelConfig) -> FeatureMLP:
|
||||
model_config = config.model if isinstance(config, ExperimentConfig) else config
|
||||
return FeatureMLP(model_config)
|
||||
143
scripts/predict_single.py
Normal file
143
scripts/predict_single.py
Normal file
@@ -0,0 +1,143 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
try:
|
||||
from .config import CORE_FEATURE_NAMES, make_experiment_config
|
||||
from .dataset import build_record_from_file, normalization_from_payload
|
||||
from .model import build_model
|
||||
except ImportError:
|
||||
from config import CORE_FEATURE_NAMES, make_experiment_config
|
||||
from dataset import build_record_from_file, normalization_from_payload
|
||||
from model import build_model
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Predict task1 RMS from a single waveform CSV.")
|
||||
parser.add_argument("--file", type=str, required=True)
|
||||
parser.add_argument("--device", type=str, default=None)
|
||||
parser.add_argument("--checkpoint", type=str, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def resolve_device(device_name: str | None, default_device: str) -> torch.device:
|
||||
requested = device_name or default_device
|
||||
if requested.startswith("cuda") and not torch.cuda.is_available():
|
||||
return torch.device("cpu")
|
||||
return torch.device(requested)
|
||||
|
||||
|
||||
def resolve_checkpoint_path(project_root: Path, checkpoint_dir: str, best_model_name: str, checkpoint_arg: str | None) -> Path:
|
||||
if checkpoint_arg:
|
||||
return Path(checkpoint_arg).resolve()
|
||||
return project_root / checkpoint_dir / "task1_feature_mlp" / best_model_name
|
||||
|
||||
|
||||
def save_prediction_figure(
|
||||
file_path: Path,
|
||||
record,
|
||||
pred_rms: float,
|
||||
true_rms: float | None,
|
||||
output_dir: Path,
|
||||
) -> Path:
|
||||
time_middle = record.time_middle.detach().cpu().numpy().reshape(-1)
|
||||
x_middle = record.x_middle.detach().cpu().numpy().reshape(-1)
|
||||
y_middle = record.y_middle.detach().cpu().numpy().reshape(-1)
|
||||
sampling_rate = record.sampling_rate
|
||||
fft_values = np.fft.rfft(x_middle)
|
||||
freqs = np.fft.rfftfreq(x_middle.size, d=1.0 / sampling_rate)
|
||||
amplitudes = np.abs(fft_values)
|
||||
|
||||
fig, axes = plt.subplots(3, 1, figsize=(12, 10))
|
||||
fig.suptitle(f"Task1 Single-File Prediction | {file_path.name}")
|
||||
|
||||
axes[0].plot(time_middle, x_middle, color="tab:blue")
|
||||
axes[0].set_title("Input Base Excitation (Middle Segment)")
|
||||
axes[0].set_xlabel("Time")
|
||||
axes[0].set_ylabel("Acceleration")
|
||||
axes[0].grid(True, alpha=0.3)
|
||||
|
||||
axes[1].plot(freqs, amplitudes, color="tab:purple")
|
||||
axes[1].axvline(record.frequency_hz, color="tab:red", linestyle="--", label=f"Dominant freq = {record.frequency_hz:.4f} Hz")
|
||||
axes[1].set_xlim(0.0, 5.0)
|
||||
axes[1].set_title("Input Spectrum")
|
||||
axes[1].set_xlabel("Frequency (Hz)")
|
||||
axes[1].set_ylabel("Amplitude")
|
||||
axes[1].grid(True, alpha=0.3)
|
||||
axes[1].legend()
|
||||
|
||||
labels = ["Pred RMS"] if true_rms is None else ["True RMS", "Pred RMS"]
|
||||
values = [pred_rms] if true_rms is None else [true_rms, pred_rms]
|
||||
colors = ["tab:orange"] if true_rms is None else ["tab:green", "tab:orange"]
|
||||
axes[2].bar(labels, values, color=colors)
|
||||
title = f"Predicted RMS = {pred_rms:.6f}"
|
||||
if true_rms is not None:
|
||||
error_percent = abs(pred_rms - true_rms) / max(abs(true_rms), 1e-12) * 100.0
|
||||
title = f"True RMS = {true_rms:.6f} | Pred RMS = {pred_rms:.6f} | Error = {error_percent:.2f}%"
|
||||
axes[2].set_title(title)
|
||||
axes[2].set_ylabel("RMS")
|
||||
axes[2].grid(True, axis="y", alpha=0.3)
|
||||
|
||||
plt.tight_layout()
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
figure_path = output_dir / f"{file_path.stem}_prediction.png"
|
||||
plt.savefig(figure_path, dpi=180, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
return figure_path
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
config = make_experiment_config()
|
||||
device = resolve_device(args.device, config.train.device)
|
||||
checkpoint_path = resolve_checkpoint_path(
|
||||
config.data.project_root,
|
||||
config.train.checkpoint_dir,
|
||||
config.train.best_model_name,
|
||||
args.checkpoint,
|
||||
)
|
||||
if not checkpoint_path.exists():
|
||||
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
|
||||
|
||||
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
normalization = normalization_from_payload(checkpoint["normalization"])
|
||||
model = build_model(config).to(device)
|
||||
model.load_state_dict(checkpoint["model_state_dict"])
|
||||
model.eval()
|
||||
|
||||
file_path = Path(args.file).resolve()
|
||||
record = build_record_from_file(file_path, split="predict", config=config.data)
|
||||
feature_norm = ((record.features - normalization.feature_mean) / normalization.feature_std).unsqueeze(0).to(device)
|
||||
x_rms_raw = record.x_rms.unsqueeze(0).to(device)
|
||||
target_mean = normalization.target_mean.to(device).view(1, 1)
|
||||
target_std = normalization.target_std.to(device).view(1, 1)
|
||||
|
||||
with torch.no_grad():
|
||||
pred_norm = model(feature_norm)
|
||||
pred_tr = torch.clamp(pred_norm * target_std + target_mean, min=1e-6)
|
||||
pred_rms = float((pred_tr * x_rms_raw).detach().cpu().numpy().reshape(-1)[0])
|
||||
|
||||
true_rms = float(record.y_rms.item()) if record.y_rms.numel() > 0 else None
|
||||
output_dir = config.data.project_root / "evaluation_outputs" / "task1_feature_mlp"
|
||||
figure_path = save_prediction_figure(file_path, record, pred_rms, true_rms, output_dir)
|
||||
|
||||
print(f"Checkpoint: {checkpoint_path}")
|
||||
print(f"Input file: {file_path}")
|
||||
print(f"Extracted frequency (Hz): {record.frequency_hz:.6f}")
|
||||
for feature_name, feature_value in zip(CORE_FEATURE_NAMES, record.features.tolist()):
|
||||
print(f"{feature_name}: {feature_value:.6f}")
|
||||
print(f"Predicted RMS: {pred_rms:.6f}")
|
||||
if true_rms is not None:
|
||||
error_percent = abs(pred_rms - true_rms) / max(abs(true_rms), 1e-12) * 100.0
|
||||
print(f"True RMS: {true_rms:.6f}")
|
||||
print(f"Relative RMS Error (%): {error_percent:.4f}")
|
||||
print(f"Figure saved to: {figure_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
274
scripts/train.py
Normal file
274
scripts/train.py
Normal file
@@ -0,0 +1,274 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import random
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pandas as pd
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from torch.optim import AdamW
|
||||
from torch.optim.lr_scheduler import ReduceLROnPlateau
|
||||
|
||||
try:
|
||||
from .config import ExperimentConfig, make_experiment_config
|
||||
from .dataset import build_dataloaders, checkpoint_normalization_payload, report_to_text
|
||||
from .model import build_model
|
||||
except ImportError:
|
||||
from config import ExperimentConfig, make_experiment_config
|
||||
from dataset import build_dataloaders, checkpoint_normalization_payload, report_to_text
|
||||
from model import build_model
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Train feature-based MLP for task1 RMS prediction.")
|
||||
parser.add_argument("--epochs", type=int, default=None)
|
||||
parser.add_argument("--batch-size", type=int, default=None)
|
||||
parser.add_argument("--device", type=str, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def set_seed(seed: int) -> None:
|
||||
random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
|
||||
def resolve_device(device_name: str) -> torch.device:
|
||||
if device_name.startswith("cuda") and not torch.cuda.is_available():
|
||||
return torch.device("cpu")
|
||||
return torch.device(device_name)
|
||||
|
||||
|
||||
def serialize_for_checkpoint(value: Any) -> Any:
|
||||
if isinstance(value, Path):
|
||||
return str(value)
|
||||
if isinstance(value, dict):
|
||||
return {key: serialize_for_checkpoint(sub_value) for key, sub_value in value.items()}
|
||||
if isinstance(value, tuple):
|
||||
return [serialize_for_checkpoint(item) for item in value]
|
||||
if isinstance(value, list):
|
||||
return [serialize_for_checkpoint(item) for item in value]
|
||||
return value
|
||||
|
||||
|
||||
def denormalize_target(pred_norm: torch.Tensor, target_mean: torch.Tensor, target_std: torch.Tensor) -> torch.Tensor:
|
||||
return pred_norm * target_std + target_mean
|
||||
|
||||
|
||||
def relative_rms_error(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
|
||||
return torch.abs(pred - target) / torch.clamp(target.abs(), min=1e-6)
|
||||
|
||||
|
||||
class RMSLoss(nn.Module):
|
||||
def __init__(self, relative_rms_weight: float, log_rms_huber_weight: float, mae_weight: float) -> None:
|
||||
super().__init__()
|
||||
self.relative_rms_weight = relative_rms_weight
|
||||
self.log_rms_huber_weight = log_rms_huber_weight
|
||||
self.mae_weight = mae_weight
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pred_tr: torch.Tensor,
|
||||
target_tr: torch.Tensor,
|
||||
target_log_tr: torch.Tensor,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
relative_loss = relative_rms_error(pred_tr, target_tr).mean()
|
||||
log_huber_loss = F.huber_loss(torch.log(torch.clamp(pred_tr, min=1e-6)), target_log_tr)
|
||||
mae_loss = F.l1_loss(pred_tr, target_tr)
|
||||
total = (
|
||||
self.relative_rms_weight * relative_loss
|
||||
+ self.log_rms_huber_weight * log_huber_loss
|
||||
+ self.mae_weight * mae_loss
|
||||
)
|
||||
return {
|
||||
"total": total,
|
||||
"relative": relative_loss,
|
||||
"log_huber": log_huber_loss,
|
||||
"mae": mae_loss,
|
||||
}
|
||||
|
||||
|
||||
def run_epoch(
|
||||
model: nn.Module,
|
||||
dataloader: torch.utils.data.DataLoader,
|
||||
optimizer: AdamW | None,
|
||||
criterion: RMSLoss,
|
||||
target_mean: torch.Tensor,
|
||||
target_std: torch.Tensor,
|
||||
device: torch.device,
|
||||
grad_clip_norm: float,
|
||||
) -> dict[str, float]:
|
||||
is_train = optimizer is not None
|
||||
model.train(is_train)
|
||||
|
||||
total_loss_sum = 0.0
|
||||
relative_loss_sum = 0.0
|
||||
log_huber_sum = 0.0
|
||||
mae_loss_sum = 0.0
|
||||
rms_error_sum = 0.0
|
||||
sample_count = 0
|
||||
|
||||
for batch in dataloader:
|
||||
x = batch["x"].to(device)
|
||||
target_tr = batch["target_raw"].to(device)
|
||||
target_log_tr = batch["target_log_raw"].to(device)
|
||||
x_rms_raw = batch["x_rms_raw"].to(device)
|
||||
y_rms_raw = batch["y_rms_raw"].to(device)
|
||||
|
||||
if is_train:
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
pred_norm = model(x)
|
||||
pred_tr = torch.clamp(denormalize_target(pred_norm, target_mean, target_std), min=1e-6)
|
||||
pred_rms = pred_tr * x_rms_raw
|
||||
losses = criterion(pred_tr, target_tr, target_log_tr)
|
||||
|
||||
if is_train:
|
||||
losses["total"].backward()
|
||||
torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm)
|
||||
optimizer.step()
|
||||
|
||||
batch_size = x.shape[0]
|
||||
total_loss_sum += losses["total"].detach().item() * batch_size
|
||||
relative_loss_sum += losses["relative"].detach().item() * batch_size
|
||||
log_huber_sum += losses["log_huber"].detach().item() * batch_size
|
||||
mae_loss_sum += losses["mae"].detach().item() * batch_size
|
||||
rms_error_sum += relative_rms_error(pred_rms.detach(), y_rms_raw.detach()).mean().item() * batch_size
|
||||
sample_count += batch_size
|
||||
|
||||
return {
|
||||
"loss": total_loss_sum / sample_count,
|
||||
"relative_loss": relative_loss_sum / sample_count,
|
||||
"log_huber_loss": log_huber_sum / sample_count,
|
||||
"mae_loss": mae_loss_sum / sample_count,
|
||||
"rms_error": rms_error_sum / sample_count,
|
||||
}
|
||||
|
||||
|
||||
def checkpoint_paths(config: ExperimentConfig) -> tuple[Path, Path]:
|
||||
root = config.data.project_root / config.train.checkpoint_dir / "task1_feature_mlp"
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
return root / config.train.best_model_name, root / config.train.history_name
|
||||
|
||||
|
||||
def train_model(config: ExperimentConfig) -> None:
|
||||
set_seed(config.train.seed)
|
||||
device = resolve_device(config.train.device)
|
||||
loaders, datasets, raw_records, reports = build_dataloaders(config)
|
||||
del raw_records
|
||||
normalization = datasets["train"].normalization
|
||||
target_mean = normalization.target_mean.to(device).view(1, 1)
|
||||
target_std = normalization.target_std.to(device).view(1, 1)
|
||||
|
||||
model = build_model(config).to(device)
|
||||
optimizer = AdamW(model.parameters(), lr=config.train.learning_rate, weight_decay=config.train.weight_decay)
|
||||
scheduler = ReduceLROnPlateau(
|
||||
optimizer,
|
||||
mode="min",
|
||||
factor=config.train.lr_scheduler_factor,
|
||||
patience=config.train.lr_scheduler_patience,
|
||||
min_lr=config.train.min_learning_rate,
|
||||
)
|
||||
criterion = RMSLoss(
|
||||
relative_rms_weight=config.loss.relative_rms_weight,
|
||||
log_rms_huber_weight=config.loss.log_rms_huber_weight,
|
||||
mae_weight=config.loss.mae_weight,
|
||||
)
|
||||
|
||||
best_val_error = float("inf")
|
||||
epochs_without_improvement = 0
|
||||
history: list[dict[str, float]] = []
|
||||
best_model_path, history_path = checkpoint_paths(config)
|
||||
|
||||
print(f"Device: {device}")
|
||||
print(report_to_text(reports))
|
||||
|
||||
for epoch in range(1, config.train.epochs + 1):
|
||||
train_metrics = run_epoch(
|
||||
model=model,
|
||||
dataloader=loaders["train"],
|
||||
optimizer=optimizer,
|
||||
criterion=criterion,
|
||||
target_mean=target_mean,
|
||||
target_std=target_std,
|
||||
device=device,
|
||||
grad_clip_norm=config.train.grad_clip_norm,
|
||||
)
|
||||
val_metrics = run_epoch(
|
||||
model=model,
|
||||
dataloader=loaders["val"],
|
||||
optimizer=None,
|
||||
criterion=criterion,
|
||||
target_mean=target_mean,
|
||||
target_std=target_std,
|
||||
device=device,
|
||||
grad_clip_norm=config.train.grad_clip_norm,
|
||||
)
|
||||
scheduler.step(val_metrics["rms_error"])
|
||||
current_lr = optimizer.param_groups[0]["lr"]
|
||||
|
||||
history_row = {
|
||||
"epoch": epoch,
|
||||
"lr": current_lr,
|
||||
"train_loss": train_metrics["loss"],
|
||||
"train_relative_loss": train_metrics["relative_loss"],
|
||||
"train_log_rms_huber_loss": train_metrics["log_huber_loss"],
|
||||
"train_mae_loss": train_metrics["mae_loss"],
|
||||
"train_rms_error": train_metrics["rms_error"],
|
||||
"val_loss": val_metrics["loss"],
|
||||
"val_relative_loss": val_metrics["relative_loss"],
|
||||
"val_log_rms_huber_loss": val_metrics["log_huber_loss"],
|
||||
"val_mae_loss": val_metrics["mae_loss"],
|
||||
"val_rms_error": val_metrics["rms_error"],
|
||||
}
|
||||
history.append(history_row)
|
||||
|
||||
print(
|
||||
f"Epoch {epoch:03d} | train_loss={train_metrics['loss']:.6f} | "
|
||||
f"val_loss={val_metrics['loss']:.6f} | val_rms_error={val_metrics['rms_error']:.6f} | lr={current_lr:.2e}"
|
||||
)
|
||||
|
||||
if val_metrics["rms_error"] < best_val_error:
|
||||
best_val_error = val_metrics["rms_error"]
|
||||
epochs_without_improvement = 0
|
||||
torch.save(
|
||||
{
|
||||
"epoch": epoch,
|
||||
"model_state_dict": model.state_dict(),
|
||||
"optimizer_state_dict": optimizer.state_dict(),
|
||||
"best_val_rms_error": best_val_error,
|
||||
"config": serialize_for_checkpoint(asdict(config)),
|
||||
"normalization": checkpoint_normalization_payload(normalization),
|
||||
},
|
||||
best_model_path,
|
||||
)
|
||||
else:
|
||||
epochs_without_improvement += 1
|
||||
if epochs_without_improvement >= config.train.early_stop_patience:
|
||||
print(f"Early stopping triggered after {epoch} epochs.")
|
||||
break
|
||||
|
||||
history_df = pd.DataFrame(history)
|
||||
history_df.to_csv(history_path, index=False)
|
||||
print(f"Best model saved to: {best_model_path}")
|
||||
print(f"Training history saved to: {history_path}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
config = make_experiment_config()
|
||||
if args.epochs is not None:
|
||||
config.train.epochs = args.epochs
|
||||
if args.batch_size is not None:
|
||||
config.train.batch_size = args.batch_size
|
||||
if args.device is not None:
|
||||
config.train.device = args.device
|
||||
train_model(config)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user