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:
2026-05-06 12:19:55 +08:00
parent dcc023cc04
commit 484643409d
48 changed files with 2706 additions and 18 deletions

Binary file not shown.

Binary file not shown.

Binary file not shown.

109
scripts/config.py Normal file
View 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
View 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
View 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
View 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
View 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
View 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()