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
src_new/__pycache__/config.cpython-310.pyc
Normal file
BIN
src_new/__pycache__/config.cpython-310.pyc
Normal file
Binary file not shown.
BIN
src_new/__pycache__/dataset.cpython-310.pyc
Normal file
BIN
src_new/__pycache__/dataset.cpython-310.pyc
Normal file
Binary file not shown.
BIN
src_new/__pycache__/evaluate.cpython-310.pyc
Normal file
BIN
src_new/__pycache__/evaluate.cpython-310.pyc
Normal file
Binary file not shown.
BIN
src_new/__pycache__/model.cpython-310.pyc
Normal file
BIN
src_new/__pycache__/model.cpython-310.pyc
Normal file
Binary file not shown.
BIN
src_new/__pycache__/train.cpython-310.pyc
Normal file
BIN
src_new/__pycache__/train.cpython-310.pyc
Normal file
Binary file not shown.
103
src_new/config.py
Normal file
103
src_new/config.py
Normal file
@@ -0,0 +1,103 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
@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"
|
||||
|
||||
max_sequence_length: int = 4096
|
||||
min_sequence_length: int = 512
|
||||
batch_size: int = 8
|
||||
num_workers: int = 0
|
||||
pin_memory: bool = True
|
||||
use_weighted_train_sampler: bool = True
|
||||
train_weight_power: float = 1.0
|
||||
train_weight_min: float = 0.5
|
||||
train_weight_max: float = 8.0
|
||||
low_frequency_emphasis_power: float = 1.25
|
||||
low_frequency_reference_hz: float = 1.0
|
||||
|
||||
use_steady_state_only: bool = True
|
||||
steady_state_start_ratio: float = 0.50
|
||||
steady_state_min_samples: int = 256
|
||||
|
||||
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_channels: int = 1
|
||||
tcn_channels: tuple[int, ...] = (32, 32, 64, 64)
|
||||
kernel_size: int = 7
|
||||
dropout: float = 0.15
|
||||
dilation_base: int = 2
|
||||
pooled_feature_dim: int = 128
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrainConfig:
|
||||
epochs: int = 120
|
||||
learning_rate: float = 1e-3
|
||||
weight_decay: float = 1e-4
|
||||
seed: int = 42
|
||||
device: str = "cuda"
|
||||
grad_clip_norm: float = 1.0
|
||||
use_amp: bool = True
|
||||
lr_scheduler_patience: int = 8
|
||||
lr_scheduler_factor: float = 0.5
|
||||
min_learning_rate: float = 1e-6
|
||||
early_stop_patience: int = 15
|
||||
checkpoint_dir: str = "checkpoints_rms"
|
||||
best_model_name: str = "best_rms_model.pt"
|
||||
history_name: str = "training_history.csv"
|
||||
|
||||
|
||||
@dataclass
|
||||
class LossConfig:
|
||||
relative_rms_weight: float = 1.0
|
||||
log_rms_weight: float = 0.5
|
||||
mae_weight: float = 0.25
|
||||
waveform_l1_weight: float = 0.03
|
||||
waveform_huber_weight: float = 0.05
|
||||
|
||||
|
||||
@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_rms_forward_config() -> ExperimentConfig:
|
||||
return ExperimentConfig()
|
||||
388
src_new/dataset.py
Normal file
388
src_new/dataset.py
Normal file
@@ -0,0 +1,388 @@
|
||||
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
|
||||
import torch.nn.functional as F
|
||||
from torch.utils.data import DataLoader, Dataset, WeightedRandomSampler
|
||||
|
||||
try:
|
||||
from .config import DataConfig, ExperimentConfig
|
||||
except ImportError:
|
||||
from config import 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))
|
||||
|
||||
|
||||
@dataclass
|
||||
class NormalizationStats:
|
||||
x_mean: torch.Tensor
|
||||
x_std: torch.Tensor
|
||||
y_wave_mean: torch.Tensor
|
||||
y_wave_std: torch.Tensor
|
||||
aux_mean: torch.Tensor
|
||||
aux_std: torch.Tensor
|
||||
target_mean: torch.Tensor
|
||||
target_std: torch.Tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class TensorStandardScaler:
|
||||
mean: torch.Tensor
|
||||
std: torch.Tensor
|
||||
|
||||
def transform(self, array: torch.Tensor | np.ndarray) -> torch.Tensor | np.ndarray:
|
||||
if isinstance(array, torch.Tensor):
|
||||
mean = self.mean.to(array.device, dtype=array.dtype)
|
||||
std = self.std.to(array.device, dtype=array.dtype)
|
||||
return (array - mean) / std
|
||||
np_array = np.asarray(array, dtype=np.float32)
|
||||
return (np_array - self.mean.cpu().numpy()) / self.std.cpu().numpy()
|
||||
|
||||
def inverse_transform(self, array: torch.Tensor | np.ndarray) -> torch.Tensor | np.ndarray:
|
||||
if isinstance(array, torch.Tensor):
|
||||
mean = self.mean.to(array.device, dtype=array.dtype)
|
||||
std = self.std.to(array.device, dtype=array.dtype)
|
||||
return array * std + mean
|
||||
np_array = np.asarray(array, dtype=np.float32)
|
||||
return np_array * self.std.cpu().numpy() + self.mean.cpu().numpy()
|
||||
|
||||
|
||||
@dataclass
|
||||
class RMSRecord:
|
||||
file_path: Path
|
||||
split: str
|
||||
time: torch.Tensor
|
||||
x_full: torch.Tensor
|
||||
x_model_input: torch.Tensor
|
||||
y_model_target: torch.Tensor
|
||||
x_rms: torch.Tensor
|
||||
y_rms: torch.Tensor
|
||||
target_value: torch.Tensor
|
||||
aux_features: torch.Tensor
|
||||
sample_weight: 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
|
||||
|
||||
|
||||
class RMSRegressionDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
records: list[RMSRecord],
|
||||
normalization: NormalizationStats,
|
||||
max_sequence_length: int,
|
||||
) -> None:
|
||||
self.records = records
|
||||
self.normalization = normalization
|
||||
self.max_sequence_length = max_sequence_length
|
||||
self.sample_weights = [record.sample_weight for record in records]
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.records)
|
||||
|
||||
def __getitem__(self, index: int) -> dict[str, Any]:
|
||||
record = self.records[index]
|
||||
x = record.x_model_input
|
||||
if x.shape[0] > self.max_sequence_length:
|
||||
x = x[-self.max_sequence_length :]
|
||||
|
||||
pad_length = self.max_sequence_length - x.shape[0]
|
||||
if pad_length > 0:
|
||||
x = F.pad(x.transpose(0, 1), (pad_length, 0), value=0.0).transpose(0, 1)
|
||||
|
||||
x_norm = (x - self.normalization.x_mean) / self.normalization.x_std
|
||||
y = record.y_model_target
|
||||
if y.shape[0] > self.max_sequence_length:
|
||||
y = y[-self.max_sequence_length :]
|
||||
aux_norm = (record.aux_features - self.normalization.aux_mean) / self.normalization.aux_std
|
||||
target_norm = (record.target_value - self.normalization.target_mean) / self.normalization.target_std
|
||||
if pad_length > 0:
|
||||
y = F.pad(y.transpose(0, 1), (pad_length, 0), value=0.0).transpose(0, 1)
|
||||
y_wave_norm = (y - self.normalization.y_wave_mean) / self.normalization.y_wave_std
|
||||
|
||||
valid_length = min(record.x_model_input.shape[0], self.max_sequence_length)
|
||||
mask = torch.zeros(self.max_sequence_length, dtype=torch.float32)
|
||||
mask[-valid_length:] = 1.0
|
||||
|
||||
return {
|
||||
"x": x_norm,
|
||||
"aux": aux_norm,
|
||||
"y": target_norm,
|
||||
"y_wave": y_wave_norm,
|
||||
"y_raw": record.y_rms,
|
||||
"x_rms_raw": record.x_rms,
|
||||
"frequency_hz": torch.tensor(record.aux_features[0].item(), dtype=torch.float32),
|
||||
"mask": mask,
|
||||
"file_name": record.file_path.name,
|
||||
"file_path": str(record.file_path),
|
||||
"valid_length": torch.tensor(valid_length, dtype=torch.long),
|
||||
"sample_weight": torch.tensor(record.sample_weight, dtype=torch.float32),
|
||||
}
|
||||
|
||||
|
||||
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 _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 select_model_segment(signal: np.ndarray, config: DataConfig) -> np.ndarray:
|
||||
if not config.use_steady_state_only:
|
||||
return signal
|
||||
start_index = int(len(signal) * config.steady_state_start_ratio)
|
||||
max_start = max(0, len(signal) - config.steady_state_min_samples)
|
||||
start_index = min(start_index, max_start)
|
||||
return signal[start_index:]
|
||||
|
||||
|
||||
def estimate_dominant_frequency(signal: np.ndarray, sampling_rate: float) -> float:
|
||||
signal = np.asarray(signal, dtype=np.float64).reshape(-1)
|
||||
if signal.size < 4:
|
||||
return 0.0
|
||||
fft_values = np.fft.rfft(signal)
|
||||
freqs = np.fft.rfftfreq(signal.size, d=1.0 / sampling_rate)
|
||||
magnitudes = np.abs(fft_values)
|
||||
magnitudes[0] = 0.0
|
||||
band_mask = (freqs >= 0.1) & (freqs <= 5.0)
|
||||
if not np.any(band_mask):
|
||||
return 0.0
|
||||
masked_magnitudes = np.where(band_mask, magnitudes, 0.0)
|
||||
return float(freqs[int(np.argmax(masked_magnitudes))])
|
||||
|
||||
|
||||
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 compute_sample_weight(y_rms: float, x_rms: float, freq_hz: float, config: DataConfig) -> float:
|
||||
gain = y_rms / max(x_rms, config.normalization_eps)
|
||||
low_freq_factor = (config.low_frequency_reference_hz / max(freq_hz, config.normalization_eps)) ** config.low_frequency_emphasis_power
|
||||
weight = (gain ** config.train_weight_power) * low_freq_factor
|
||||
return float(np.clip(weight, config.train_weight_min, config.train_weight_max))
|
||||
|
||||
|
||||
def load_split_records(config: DataConfig, split: str) -> tuple[list[RMSRecord], SplitLoadReport]:
|
||||
records: list[RMSRecord] = []
|
||||
report = SplitLoadReport(split=split)
|
||||
|
||||
for file_path in list_split_files(config, split):
|
||||
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:
|
||||
report.skipped_files.append((file_path.name, "missing required sensor"))
|
||||
continue
|
||||
|
||||
aligned = base_df.rename(columns={config.base_axis: "base_signal"})
|
||||
aligned = aligned.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():
|
||||
report.skipped_files.append((file_path.name, "remaining NaN after interpolation"))
|
||||
continue
|
||||
|
||||
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)
|
||||
if len(x_values) < config.min_sequence_length:
|
||||
report.skipped_files.append((file_path.name, f"sequence too short: {len(x_values)}"))
|
||||
continue
|
||||
|
||||
x_model = select_model_segment(x_values, config)
|
||||
y_model = select_model_segment(y_values, config)
|
||||
time_model = select_model_segment(time_values, config)
|
||||
sampling_rate = estimate_sampling_rate(time_model)
|
||||
file_frequency = extract_frequency_hz(file_path.name)
|
||||
x_rms = calculate_rms(x_model)
|
||||
y_rms = calculate_rms(y_model)
|
||||
target_value = float(np.log(max(y_rms / max(x_rms, config.normalization_eps), config.normalization_eps)))
|
||||
dominant_freq = estimate_dominant_frequency(x_model, sampling_rate)
|
||||
feature_frequency = file_frequency if file_frequency is not None else dominant_freq
|
||||
sample_weight = compute_sample_weight(y_rms=y_rms, x_rms=x_rms, freq_hz=feature_frequency, config=config) if split == "train" else 1.0
|
||||
|
||||
aux_features = torch.tensor(
|
||||
[feature_frequency, x_rms, float(len(x_model)) / float(config.max_sequence_length)],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
records.append(
|
||||
RMSRecord(
|
||||
file_path=file_path,
|
||||
split=split,
|
||||
time=torch.tensor(time_model, dtype=torch.float64),
|
||||
x_full=torch.tensor(x_values[:, None], dtype=torch.float32),
|
||||
x_model_input=torch.tensor(x_model[:, None], dtype=torch.float32),
|
||||
y_model_target=torch.tensor(y_model[:, None], dtype=torch.float32),
|
||||
x_rms=torch.tensor([x_rms], dtype=torch.float32),
|
||||
y_rms=torch.tensor([y_rms], dtype=torch.float32),
|
||||
target_value=torch.tensor([target_value], dtype=torch.float32),
|
||||
aux_features=aux_features,
|
||||
sample_weight=sample_weight,
|
||||
interpolation_count=interpolation_count,
|
||||
)
|
||||
)
|
||||
report.loaded_files.append(file_path.name)
|
||||
if interpolation_count > 0:
|
||||
report.interpolated_files[file_path.name] = interpolation_count
|
||||
|
||||
if not records:
|
||||
raise RuntimeError(f"No usable records found for split='{split}'.")
|
||||
return records, report
|
||||
|
||||
|
||||
def fit_normalization(records: list[RMSRecord], config: DataConfig) -> NormalizationStats:
|
||||
x_all = torch.cat([record.x_model_input for record in records], dim=0)
|
||||
y_all = torch.cat([record.y_model_target for record in records], dim=0)
|
||||
aux_all = torch.stack([record.aux_features for record in records], dim=0)
|
||||
target_all = torch.cat([record.target_value for record in records], dim=0)
|
||||
|
||||
def safe_std(tensor: torch.Tensor, dim: int) -> torch.Tensor:
|
||||
std = tensor.std(dim=dim, unbiased=False)
|
||||
return torch.clamp(std, min=config.normalization_eps)
|
||||
|
||||
return NormalizationStats(
|
||||
x_mean=x_all.mean(dim=0),
|
||||
x_std=safe_std(x_all, dim=0),
|
||||
y_wave_mean=y_all.mean(dim=0),
|
||||
y_wave_std=safe_std(y_all, dim=0),
|
||||
aux_mean=aux_all.mean(dim=0),
|
||||
aux_std=safe_std(aux_all, dim=0),
|
||||
target_mean=target_all.mean(dim=0, keepdim=True),
|
||||
target_std=safe_std(target_all, dim=0).view(1),
|
||||
)
|
||||
|
||||
|
||||
def build_datasets(
|
||||
config: ExperimentConfig | DataConfig,
|
||||
) -> tuple[dict[str, RMSRegressionDataset], NormalizationStats, 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")
|
||||
|
||||
datasets = {
|
||||
"train": RMSRegressionDataset(train_records, normalization, data_config.max_sequence_length),
|
||||
"val": RMSRegressionDataset(val_records, normalization, data_config.max_sequence_length),
|
||||
"test": RMSRegressionDataset(test_records, normalization, data_config.max_sequence_length),
|
||||
}
|
||||
return datasets, normalization, {"train": train_report, "val": val_report, "test": test_report}
|
||||
|
||||
|
||||
def build_dataloaders(
|
||||
config: ExperimentConfig | DataConfig,
|
||||
) -> tuple[dict[str, DataLoader], dict[str, RMSRegressionDataset], dict[str, SplitLoadReport]]:
|
||||
data_config = config.data if isinstance(config, ExperimentConfig) else config
|
||||
datasets, _, reports = build_datasets(config)
|
||||
|
||||
train_sampler = None
|
||||
train_shuffle = True
|
||||
if data_config.use_weighted_train_sampler:
|
||||
weights = torch.tensor(datasets["train"].sample_weights, dtype=torch.double)
|
||||
train_sampler = WeightedRandomSampler(weights, num_samples=len(weights), replacement=True)
|
||||
train_shuffle = False
|
||||
|
||||
loaders = {
|
||||
"train": DataLoader(
|
||||
datasets["train"],
|
||||
batch_size=data_config.batch_size,
|
||||
shuffle=train_shuffle,
|
||||
sampler=train_sampler,
|
||||
num_workers=data_config.num_workers,
|
||||
pin_memory=data_config.pin_memory,
|
||||
),
|
||||
"val": DataLoader(
|
||||
datasets["val"],
|
||||
batch_size=data_config.batch_size,
|
||||
shuffle=False,
|
||||
num_workers=data_config.num_workers,
|
||||
pin_memory=data_config.pin_memory,
|
||||
),
|
||||
"test": DataLoader(
|
||||
datasets["test"],
|
||||
batch_size=data_config.batch_size,
|
||||
shuffle=False,
|
||||
num_workers=data_config.num_workers,
|
||||
pin_memory=data_config.pin_memory,
|
||||
),
|
||||
}
|
||||
return loaders, datasets, reports
|
||||
|
||||
|
||||
def get_dataloaders(
|
||||
config: ExperimentConfig | DataConfig,
|
||||
) -> tuple[
|
||||
dict[str, DataLoader],
|
||||
dict[str, RMSRegressionDataset],
|
||||
dict[str, SplitLoadReport],
|
||||
TensorStandardScaler,
|
||||
TensorStandardScaler,
|
||||
TensorStandardScaler,
|
||||
]:
|
||||
loaders, datasets, reports = build_dataloaders(config)
|
||||
normalization = datasets["train"].normalization
|
||||
x_scaler = TensorStandardScaler(normalization.x_mean, normalization.x_std)
|
||||
aux_scaler = TensorStandardScaler(normalization.aux_mean, normalization.aux_std)
|
||||
y_scaler = TensorStandardScaler(normalization.target_mean, normalization.target_std)
|
||||
return loaders, datasets, reports, x_scaler, aux_scaler, y_scaler
|
||||
|
||||
|
||||
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)
|
||||
222
src_new/evaluate.py
Normal file
222
src_new/evaluate.py
Normal file
@@ -0,0 +1,222 @@
|
||||
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 ExperimentConfig, make_rms_forward_config
|
||||
from .dataset import get_dataloaders, report_to_text
|
||||
from .model import build_model
|
||||
except ImportError:
|
||||
from config import ExperimentConfig, make_rms_forward_config
|
||||
from dataset import get_dataloaders, report_to_text
|
||||
from model import build_model
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Evaluate direct RMS regression model on harmonic data.")
|
||||
parser.add_argument("--split", choices=("train", "val", "test"), default="val")
|
||||
parser.add_argument("--sample-index", type=int, default=0)
|
||||
parser.add_argument("--device", type=str, default=None)
|
||||
parser.add_argument("--checkpoint", type=str, default=None)
|
||||
parser.add_argument("--all-samples", action="store_true")
|
||||
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 / "forward_rms" / 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 inverse_waveform(array: torch.Tensor, mean: torch.Tensor, std: torch.Tensor) -> torch.Tensor:
|
||||
mean = mean.to(array.device, dtype=array.dtype)
|
||||
std = std.to(array.device, dtype=array.dtype)
|
||||
return array * std + mean
|
||||
|
||||
|
||||
def predict_sample(
|
||||
model: torch.nn.Module,
|
||||
sample: dict[str, object],
|
||||
device: torch.device,
|
||||
y_scaler,
|
||||
x_scaler,
|
||||
y_wave_mean: torch.Tensor,
|
||||
y_wave_std: torch.Tensor,
|
||||
) -> dict[str, float | str | np.ndarray]:
|
||||
x = sample["x"].unsqueeze(0).to(device)
|
||||
aux = sample["aux"].unsqueeze(0).to(device)
|
||||
mask = sample["mask"].unsqueeze(0).to(device)
|
||||
x_rms_raw = sample["x_rms_raw"].unsqueeze(0).to(device)
|
||||
y_true = float(sample["y_raw"].detach().cpu().numpy().reshape(-1)[0])
|
||||
frequency_hz = float(sample["frequency_hz"].item())
|
||||
file_name = str(sample["file_name"])
|
||||
valid_mask = sample["mask"].detach().cpu().numpy() > 0.5
|
||||
|
||||
with torch.no_grad():
|
||||
outputs = model(x, mask, aux)
|
||||
|
||||
pred_target = y_scaler.inverse_transform(outputs["rms"])
|
||||
y_pred = float((torch.exp(pred_target) * x_rms_raw).detach().cpu().numpy().reshape(-1)[0])
|
||||
pred_wave = inverse_waveform(outputs["waveform"], y_wave_mean, y_wave_std).detach().cpu().numpy().reshape(-1)[valid_mask]
|
||||
true_wave = inverse_waveform(sample["y_wave"], y_wave_mean.cpu(), y_wave_std.cpu()).detach().cpu().numpy().reshape(-1)[valid_mask]
|
||||
x_wave = x_scaler.inverse_transform(sample["x"]).reshape(-1)[valid_mask]
|
||||
return {
|
||||
"file_name": file_name,
|
||||
"frequency_hz": frequency_hz,
|
||||
"true_rms": y_true,
|
||||
"pred_rms": y_pred,
|
||||
"relative_error_percent": float(relative_percent_error(y_true, y_pred)),
|
||||
"x_wave": x_wave.numpy().reshape(-1),
|
||||
"true_wave": true_wave,
|
||||
"pred_wave": pred_wave,
|
||||
}
|
||||
|
||||
|
||||
def evaluate_sample(
|
||||
model: torch.nn.Module,
|
||||
sample: dict[str, object],
|
||||
device: torch.device,
|
||||
y_scaler,
|
||||
x_scaler,
|
||||
y_wave_mean: torch.Tensor,
|
||||
y_wave_std: torch.Tensor,
|
||||
) -> dict[str, float | str]:
|
||||
result = predict_sample(model, sample, device, y_scaler, x_scaler, y_wave_mean, y_wave_std)
|
||||
return {
|
||||
"file_name": str(result["file_name"]),
|
||||
"frequency_hz": float(result["frequency_hz"]),
|
||||
"true_rms": float(result["true_rms"]),
|
||||
"pred_rms": float(result["pred_rms"]),
|
||||
"relative_error_percent": float(result["relative_error_percent"]),
|
||||
}
|
||||
|
||||
|
||||
def save_single_sample_figure(result: dict[str, float | str | np.ndarray], split: str, save_dir: Path, sample_index: int) -> Path:
|
||||
true_rms = float(result["true_rms"])
|
||||
pred_rms = float(result["pred_rms"])
|
||||
error_percent = relative_percent_error(true_rms, pred_rms)
|
||||
x_signal = np.asarray(result["x_wave"], dtype=np.float64).reshape(-1)
|
||||
true_wave = np.asarray(result["true_wave"], dtype=np.float64).reshape(-1)
|
||||
pred_wave = np.asarray(result["pred_wave"], dtype=np.float64).reshape(-1)
|
||||
time_steps = np.arange(x_signal.shape[0], dtype=np.float64)
|
||||
|
||||
fig, axes = plt.subplots(3, 1, figsize=(12, 10))
|
||||
fig.suptitle(f"Forward RMS regression | {split} | {result['file_name']}")
|
||||
axes[0].plot(time_steps, x_signal, linewidth=1.2)
|
||||
axes[0].set_title("Input Base Excitation (Steady-State Segment)")
|
||||
axes[0].set_xlabel("Time Step")
|
||||
axes[0].set_ylabel("Acceleration")
|
||||
axes[0].grid(True, alpha=0.3)
|
||||
axes[1].plot(time_steps, true_wave, label="True response", linewidth=1.2, color="tab:blue")
|
||||
axes[1].plot(time_steps, pred_wave, label="Pred response", linewidth=1.2, color="tab:orange", alpha=0.85)
|
||||
axes[1].set_title("Auxiliary Waveform Head: True vs Predicted Top Response")
|
||||
axes[1].set_xlabel("Time Step")
|
||||
axes[1].set_ylabel("Acceleration")
|
||||
axes[1].grid(True, alpha=0.3)
|
||||
axes[1].legend()
|
||||
axes[2].bar(["True RMS", "Pred RMS"], [true_rms, pred_rms], color=["tab:blue", "tab:orange"])
|
||||
axes[2].set_title(f"True RMS: {true_rms:.4f}, Pred RMS: {pred_rms:.4f}, Error: {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}_waveform.png"
|
||||
plt.savefig(figure_path, dpi=180, bbox_inches="tight")
|
||||
plt.show()
|
||||
plt.close(fig)
|
||||
return figure_path
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
config = make_rms_forward_config()
|
||||
device = resolve_device(args.device, config)
|
||||
checkpoint_path = resolve_checkpoint_path(config, args.checkpoint)
|
||||
loaders, datasets, reports, x_scaler, aux_scaler, y_scaler = get_dataloaders(config)
|
||||
print(report_to_text(reports))
|
||||
|
||||
if not checkpoint_path.exists():
|
||||
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
|
||||
|
||||
dataset = datasets[args.split]
|
||||
sample = dataset[args.sample_index]
|
||||
normalization = datasets["train"].normalization
|
||||
model = build_model(config).to(device)
|
||||
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
model.load_state_dict(checkpoint["model_state_dict"])
|
||||
model.eval()
|
||||
save_dir = config.data.project_root / "evaluation_outputs" / "forward_rms"
|
||||
if args.all_samples:
|
||||
rows = [
|
||||
evaluate_sample(
|
||||
model,
|
||||
dataset[index],
|
||||
device,
|
||||
y_scaler,
|
||||
x_scaler,
|
||||
normalization.y_wave_mean,
|
||||
normalization.y_wave_std,
|
||||
)
|
||||
for index in range(len(dataset))
|
||||
]
|
||||
result_df = pd.DataFrame(rows)
|
||||
result_df = result_df.sort_values(["relative_error_percent", "file_name"]).reset_index(drop=True)
|
||||
csv_path = save_dir / f"evaluation_{args.split}_all_samples.csv"
|
||||
save_dir.mkdir(parents=True, exist_ok=True)
|
||||
result_df.to_csv(csv_path, index=False)
|
||||
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}")
|
||||
print(result_df.to_string(index=False))
|
||||
return
|
||||
|
||||
result = predict_sample(
|
||||
model,
|
||||
sample,
|
||||
device,
|
||||
y_scaler,
|
||||
x_scaler,
|
||||
normalization.y_wave_mean,
|
||||
normalization.y_wave_std,
|
||||
)
|
||||
figure_path = save_single_sample_figure(
|
||||
result=result,
|
||||
split=args.split,
|
||||
save_dir=save_dir,
|
||||
sample_index=args.sample_index,
|
||||
)
|
||||
print(f"Checkpoint: {checkpoint_path}")
|
||||
print(f"Sample file: {result['file_name']}")
|
||||
print(f"True RMS: {result['true_rms']:.6f}")
|
||||
print(f"Pred RMS: {result['pred_rms']:.6f}")
|
||||
print(f"Relative RMS Error (%): {result['relative_error_percent']:.4f}")
|
||||
print(f"Figure saved to: {figure_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
140
src_new/model.py
Normal file
140
src_new/model.py
Normal file
@@ -0,0 +1,140 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn.utils import weight_norm
|
||||
|
||||
try:
|
||||
from .config import ExperimentConfig, ModelConfig
|
||||
except ImportError:
|
||||
from config import ExperimentConfig, ModelConfig
|
||||
|
||||
|
||||
class Chomp1d(nn.Module):
|
||||
def __init__(self, chomp_size: int) -> None:
|
||||
super().__init__()
|
||||
self.chomp_size = chomp_size
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self.chomp_size == 0:
|
||||
return x
|
||||
return x[:, :, :-self.chomp_size].contiguous()
|
||||
|
||||
|
||||
class CausalConv1d(nn.Module):
|
||||
def __init__(self, in_channels: int, out_channels: int, kernel_size: int, dilation: int = 1) -> None:
|
||||
super().__init__()
|
||||
padding = (kernel_size - 1) * dilation
|
||||
self.net = nn.Sequential(
|
||||
weight_norm(
|
||||
nn.Conv1d(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
)
|
||||
),
|
||||
Chomp1d(padding),
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.net(x)
|
||||
|
||||
|
||||
class TemporalBlock(nn.Module):
|
||||
def __init__(self, in_channels: int, out_channels: int, kernel_size: int, dilation: int, dropout: float) -> None:
|
||||
super().__init__()
|
||||
self.conv1 = CausalConv1d(in_channels, out_channels, kernel_size, dilation=dilation)
|
||||
self.act1 = nn.GELU()
|
||||
self.dropout1 = nn.Dropout(dropout)
|
||||
self.conv2 = CausalConv1d(out_channels, out_channels, kernel_size, dilation=dilation)
|
||||
self.act2 = nn.GELU()
|
||||
self.dropout2 = nn.Dropout(dropout)
|
||||
self.residual = nn.Conv1d(in_channels, out_channels, kernel_size=1) if in_channels != out_channels else nn.Identity()
|
||||
self.final_act = nn.GELU()
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
residual = self.residual(x)
|
||||
out = self.dropout1(self.act1(self.conv1(x)))
|
||||
out = self.dropout2(self.act2(self.conv2(out)))
|
||||
return self.final_act(out + residual)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReceptiveFieldInfo:
|
||||
receptive_field: int
|
||||
dilations: tuple[int, ...]
|
||||
|
||||
|
||||
class RMSRegressionTCN(nn.Module):
|
||||
def __init__(self, config: ModelConfig) -> None:
|
||||
super().__init__()
|
||||
blocks: list[nn.Module] = []
|
||||
in_channels = config.input_channels
|
||||
dilations: list[int] = []
|
||||
|
||||
for level, out_channels in enumerate(config.tcn_channels):
|
||||
dilation = config.dilation_base ** level
|
||||
dilations.append(dilation)
|
||||
blocks.append(
|
||||
TemporalBlock(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=config.kernel_size,
|
||||
dilation=dilation,
|
||||
dropout=config.dropout,
|
||||
)
|
||||
)
|
||||
in_channels = out_channels
|
||||
|
||||
self.encoder = nn.Sequential(*blocks)
|
||||
self.rms_head = nn.Sequential(
|
||||
nn.Linear(in_channels * 2 + 3, config.pooled_feature_dim),
|
||||
nn.GELU(),
|
||||
nn.Dropout(config.dropout),
|
||||
nn.Linear(config.pooled_feature_dim, config.pooled_feature_dim),
|
||||
nn.GELU(),
|
||||
nn.Dropout(config.dropout),
|
||||
nn.Linear(config.pooled_feature_dim, 1),
|
||||
)
|
||||
self.waveform_head = nn.Sequential(
|
||||
nn.Conv1d(in_channels, in_channels, kernel_size=1),
|
||||
nn.GELU(),
|
||||
nn.Dropout(config.dropout),
|
||||
nn.Conv1d(in_channels, 1, kernel_size=1),
|
||||
)
|
||||
self.receptive_field_info = compute_receptive_field(config)
|
||||
|
||||
def forward(self, x: torch.Tensor, mask: torch.Tensor, aux: torch.Tensor) -> dict[str, torch.Tensor]:
|
||||
if x.ndim != 3:
|
||||
raise ValueError(f"Expected x shape (batch, seq, channels), got {tuple(x.shape)}")
|
||||
features = self.encoder(x.transpose(1, 2))
|
||||
mask_1d = mask.unsqueeze(1)
|
||||
masked_features = features * mask_1d
|
||||
valid_count = mask_1d.sum(dim=2).clamp_min(1.0)
|
||||
mean_pool = masked_features.sum(dim=2) / valid_count
|
||||
masked_for_max = features.masked_fill(mask_1d == 0.0, float("-inf"))
|
||||
max_pool = masked_for_max.max(dim=2).values
|
||||
max_pool = torch.where(torch.isfinite(max_pool), max_pool, torch.zeros_like(max_pool))
|
||||
fused = torch.cat([mean_pool, max_pool, aux], dim=1)
|
||||
rms_prediction = self.rms_head(fused)
|
||||
waveform_prediction = self.waveform_head(features).transpose(1, 2)
|
||||
return {"rms": rms_prediction, "waveform": waveform_prediction}
|
||||
|
||||
|
||||
def compute_receptive_field(config: ModelConfig) -> ReceptiveFieldInfo:
|
||||
receptive_field = 1
|
||||
dilations: list[int] = []
|
||||
for level, _ in enumerate(config.tcn_channels):
|
||||
dilation = config.dilation_base ** level
|
||||
dilations.append(dilation)
|
||||
receptive_field += 2 * (config.kernel_size - 1) * dilation
|
||||
return ReceptiveFieldInfo(receptive_field=receptive_field, dilations=tuple(dilations))
|
||||
|
||||
|
||||
def build_model(config: ExperimentConfig | ModelConfig) -> RMSRegressionTCN:
|
||||
model_config = config.model if isinstance(config, ExperimentConfig) else config
|
||||
return RMSRegressionTCN(model_config)
|
||||
327
src_new/train.py
Normal file
327
src_new/train.py
Normal file
@@ -0,0 +1,327 @@
|
||||
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_rms_forward_config
|
||||
from .dataset import build_dataloaders, report_to_text
|
||||
from .model import build_model, compute_receptive_field
|
||||
except ImportError:
|
||||
from config import ExperimentConfig, make_rms_forward_config
|
||||
from dataset import build_dataloaders, report_to_text
|
||||
from model import build_model, compute_receptive_field
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Train TCN for direct RMS regression.")
|
||||
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_norm: torch.Tensor,
|
||||
target_mean: torch.Tensor,
|
||||
target_std: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
pred = pred_norm * target_std + target_mean
|
||||
target = target_norm * target_std + target_mean
|
||||
return pred, target
|
||||
|
||||
|
||||
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 RMSRegressionLoss(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
relative_rms_weight: float,
|
||||
log_rms_weight: float,
|
||||
mae_weight: float,
|
||||
waveform_l1_weight: float,
|
||||
waveform_huber_weight: float,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.relative_rms_weight = relative_rms_weight
|
||||
self.log_rms_weight = log_rms_weight
|
||||
self.mae_weight = mae_weight
|
||||
self.waveform_l1_weight = waveform_l1_weight
|
||||
self.waveform_huber_weight = waveform_huber_weight
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pred: torch.Tensor,
|
||||
target: torch.Tensor,
|
||||
pred_wave: torch.Tensor,
|
||||
target_wave: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
rel_loss = relative_rms_error(pred, target).mean()
|
||||
log_loss = F.huber_loss(torch.log(torch.clamp(pred, min=1e-6)), torch.log(torch.clamp(target, min=1e-6)))
|
||||
mae_loss = F.l1_loss(pred, target)
|
||||
mask_expanded = mask.unsqueeze(-1)
|
||||
valid_count = mask_expanded.sum().clamp_min(1.0)
|
||||
wave_residual = (pred_wave - target_wave) * mask_expanded
|
||||
waveform_l1 = torch.abs(wave_residual).sum() / valid_count
|
||||
waveform_huber = F.huber_loss(pred_wave * mask_expanded, target_wave * mask_expanded, reduction="sum") / valid_count
|
||||
total = (
|
||||
self.relative_rms_weight * rel_loss
|
||||
+ self.log_rms_weight * log_loss
|
||||
+ self.mae_weight * mae_loss
|
||||
+ self.waveform_l1_weight * waveform_l1
|
||||
+ self.waveform_huber_weight * waveform_huber
|
||||
)
|
||||
return {
|
||||
"total": total,
|
||||
"relative": rel_loss,
|
||||
"log": log_loss,
|
||||
"mae": mae_loss,
|
||||
"waveform_l1": waveform_l1,
|
||||
"waveform_huber": waveform_huber,
|
||||
}
|
||||
|
||||
|
||||
def run_epoch(
|
||||
model: nn.Module,
|
||||
dataloader: torch.utils.data.DataLoader,
|
||||
optimizer: AdamW | None,
|
||||
criterion: RMSRegressionLoss,
|
||||
target_mean: torch.Tensor,
|
||||
target_std: torch.Tensor,
|
||||
device: torch.device,
|
||||
grad_clip_norm: float,
|
||||
scaler: torch.cuda.amp.GradScaler,
|
||||
amp_enabled: bool,
|
||||
) -> dict[str, float]:
|
||||
is_train = optimizer is not None
|
||||
model.train(is_train)
|
||||
|
||||
total_loss_sum = 0.0
|
||||
relative_loss_sum = 0.0
|
||||
log_loss_sum = 0.0
|
||||
mae_loss_sum = 0.0
|
||||
waveform_l1_sum = 0.0
|
||||
waveform_huber_sum = 0.0
|
||||
rms_error_sum = 0.0
|
||||
sample_count = 0
|
||||
|
||||
for batch in dataloader:
|
||||
x = batch["x"].to(device)
|
||||
aux = batch["aux"].to(device)
|
||||
mask = batch["mask"].to(device)
|
||||
y_norm = batch["y"].to(device)
|
||||
y_wave = batch["y_wave"].to(device)
|
||||
x_rms_raw = batch["x_rms_raw"].to(device)
|
||||
y_rms_raw = batch["y_raw"].to(device)
|
||||
|
||||
if is_train:
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
with torch.amp.autocast(device_type=device.type, enabled=amp_enabled):
|
||||
outputs = model(x, mask, aux)
|
||||
pred_target, _ = denormalize_target(outputs["rms"], y_norm, target_mean, target_std)
|
||||
pred = torch.exp(pred_target) * x_rms_raw
|
||||
target = y_rms_raw
|
||||
losses = criterion(pred, target, outputs["waveform"], y_wave, mask)
|
||||
|
||||
if is_train:
|
||||
scaler.scale(losses["total"]).backward()
|
||||
scaler.unscale_(optimizer)
|
||||
torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm)
|
||||
scaler.step(optimizer)
|
||||
scaler.update()
|
||||
|
||||
batch_size = x.shape[0]
|
||||
total_loss_sum += losses["total"].detach().item() * batch_size
|
||||
relative_loss_sum += losses["relative"].detach().item() * batch_size
|
||||
log_loss_sum += losses["log"].detach().item() * batch_size
|
||||
mae_loss_sum += losses["mae"].detach().item() * batch_size
|
||||
waveform_l1_sum += losses["waveform_l1"].detach().item() * batch_size
|
||||
waveform_huber_sum += losses["waveform_huber"].detach().item() * batch_size
|
||||
rms_error_sum += relative_rms_error(pred.detach(), target.detach()).mean().item() * batch_size
|
||||
sample_count += batch_size
|
||||
|
||||
return {
|
||||
"loss": total_loss_sum / sample_count,
|
||||
"relative_loss": relative_loss_sum / sample_count,
|
||||
"log_loss": log_loss_sum / sample_count,
|
||||
"mae_loss": mae_loss_sum / sample_count,
|
||||
"waveform_l1": waveform_l1_sum / sample_count,
|
||||
"waveform_huber": waveform_huber_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 / "forward_rms"
|
||||
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, reports = build_dataloaders(config)
|
||||
train_dataset = datasets["train"]
|
||||
normalization = train_dataset.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)
|
||||
receptive_field = compute_receptive_field(config.model)
|
||||
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 = RMSRegressionLoss(
|
||||
relative_rms_weight=config.loss.relative_rms_weight,
|
||||
log_rms_weight=config.loss.log_rms_weight,
|
||||
mae_weight=config.loss.mae_weight,
|
||||
waveform_l1_weight=config.loss.waveform_l1_weight,
|
||||
waveform_huber_weight=config.loss.waveform_huber_weight,
|
||||
)
|
||||
amp_enabled = config.train.use_amp and device.type == "cuda"
|
||||
scaler = torch.cuda.amp.GradScaler(enabled=amp_enabled)
|
||||
|
||||
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(f"TCN receptive field: {receptive_field.receptive_field} samples, dilations={receptive_field.dilations}")
|
||||
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,
|
||||
scaler=scaler,
|
||||
amp_enabled=amp_enabled,
|
||||
)
|
||||
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,
|
||||
scaler=scaler,
|
||||
amp_enabled=amp_enabled,
|
||||
)
|
||||
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_loss": train_metrics["log_loss"],
|
||||
"train_mae_loss": train_metrics["mae_loss"],
|
||||
"train_waveform_l1": train_metrics["waveform_l1"],
|
||||
"train_waveform_huber": train_metrics["waveform_huber"],
|
||||
"train_rms_error": train_metrics["rms_error"],
|
||||
"val_loss": val_metrics["loss"],
|
||||
"val_relative_loss": val_metrics["relative_loss"],
|
||||
"val_log_loss": val_metrics["log_loss"],
|
||||
"val_mae_loss": val_metrics["mae_loss"],
|
||||
"val_waveform_l1": val_metrics["waveform_l1"],
|
||||
"val_waveform_huber": val_metrics["waveform_huber"],
|
||||
"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)),
|
||||
},
|
||||
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_rms_forward_config()
|
||||
if args.epochs is not None:
|
||||
config.train.epochs = args.epochs
|
||||
if args.batch_size is not None:
|
||||
config.data.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