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()