| """ |
| Training loop for CircuitTransformer models. |
| |
| Supports both Model A (plain HH) and Model B (HH+ACh). |
| Includes: |
| - AdamW optimizer with cosine LR + linear warmup |
| - Gradient clipping |
| - Early stopping on validation loss |
| - Per-statistic RΒ² logging |
| - Checkpoint saving (best + final) |
| """ |
|
|
| from __future__ import annotations |
|
|
| import json |
| import logging |
| import math |
| import os |
| import time |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import torch.nn as nn |
| from torch.utils.data import DataLoader |
|
|
| from .config import OUTPUT_STATS, TrainConfig |
| from .dataset import Normalizer, SimDataset, build_datasets |
| from .model import CircuitTransformer, CircuitMLP, build_model_a, build_model_b |
|
|
| logger = logging.getLogger(__name__) |
|
|
|
|
| |
|
|
| class CosineWarmupScheduler: |
| """Linear warmup then cosine decay to 0.""" |
|
|
| def __init__(self, optimizer, warmup_steps: int, total_steps: int): |
| self.optimizer = optimizer |
| self.warmup_steps = warmup_steps |
| self.total_steps = total_steps |
| self.step_count = 0 |
| self.base_lrs = [pg["lr"] for pg in optimizer.param_groups] |
|
|
| def step(self): |
| self.step_count += 1 |
| if self.step_count <= self.warmup_steps: |
| |
| scale = self.step_count / max(1, self.warmup_steps) |
| else: |
| |
| progress = (self.step_count - self.warmup_steps) / max( |
| 1, self.total_steps - self.warmup_steps |
| ) |
| scale = 0.5 * (1.0 + math.cos(math.pi * progress)) |
|
|
| for pg, base_lr in zip(self.optimizer.param_groups, self.base_lrs): |
| pg["lr"] = base_lr * scale |
|
|
| def get_lr(self) -> float: |
| return self.optimizer.param_groups[0]["lr"] |
|
|
|
|
| |
|
|
| def compute_r2_per_stat( |
| y_pred: np.ndarray, y_true: np.ndarray |
| ) -> dict[str, float]: |
| """Compute RΒ² for each of the 11 output statistics. |
| |
| Works in NORMALIZED space (so RΒ²=1 means perfect prediction of z-scores). |
| """ |
| r2s = {} |
| for i, name in enumerate(OUTPUT_STATS): |
| ss_res = np.sum((y_true[:, i] - y_pred[:, i]) ** 2) |
| ss_tot = np.sum((y_true[:, i] - y_true[:, i].mean()) ** 2) |
| r2 = 1.0 - ss_res / max(ss_tot, 1e-8) |
| r2s[name] = round(float(r2), 4) |
| return r2s |
|
|
|
|
| |
|
|
| @torch.no_grad() |
| def evaluate( |
| model: CircuitTransformer, |
| loader: DataLoader, |
| criterion: nn.Module, |
| device: torch.device, |
| ) -> tuple[float, dict[str, float]]: |
| """Evaluate model on a dataset. |
| |
| Returns: |
| (mean_loss, per_stat_r2) |
| """ |
| model.eval() |
| total_loss = 0.0 |
| n_batches = 0 |
| all_preds = [] |
| all_targets = [] |
|
|
| for X_batch, Y_batch in loader: |
| X_batch = X_batch.to(device) |
| Y_batch = Y_batch.to(device) |
|
|
| preds = model(X_batch) |
| loss = criterion(preds, Y_batch) |
| total_loss += loss.item() |
| n_batches += 1 |
|
|
| all_preds.append(preds.cpu().numpy()) |
| all_targets.append(Y_batch.cpu().numpy()) |
|
|
| mean_loss = total_loss / max(n_batches, 1) |
| all_preds = np.concatenate(all_preds, axis=0) |
| all_targets = np.concatenate(all_targets, axis=0) |
| r2s = compute_r2_per_stat(all_preds, all_targets) |
|
|
| return mean_loss, r2s |
|
|
|
|
| def train_xgboost( |
| cfg: TrainConfig, |
| model_variant: str, |
| ) -> dict: |
| """Train XGBoost model (no GPU needed). Uses raw numpy arrays.""" |
| from xgboost import XGBRegressor |
| from sklearn.multioutput import MultiOutputRegressor |
|
|
| train_ds, val_ds, x_norm, y_norm, meta = build_datasets(cfg, model_variant) |
|
|
| X_train = train_ds.X.numpy() |
| Y_train = train_ds.Y.numpy() |
| X_val = val_ds.X.numpy() |
| Y_val = val_ds.Y.numpy() |
|
|
| logger.info(f"XGBoost {model_variant}: {X_train.shape[0]} train, {X_val.shape[0]} val") |
|
|
| import time |
| t0 = time.time() |
|
|
| model = MultiOutputRegressor(XGBRegressor( |
| n_estimators=200, |
| max_depth=6, |
| learning_rate=0.1, |
| subsample=0.8, |
| colsample_bytree=0.8, |
| random_state=cfg.seed, |
| n_jobs=-1, |
| )) |
| model.fit(X_train, Y_train) |
|
|
| |
| Y_pred = model.predict(X_val) |
| r2s = compute_r2_per_stat(Y_pred, Y_val) |
| mean_r2 = float(np.mean(list(r2s.values()))) |
|
|
| |
| mse = float(np.mean((Y_pred - Y_val) ** 2)) |
|
|
| total_time = time.time() - t0 |
|
|
| logger.info(f"XGBoost {model_variant}: Mean RΒ²={mean_r2:.4f}, MSE={mse:.5f}, Time={total_time:.1f}s") |
| for stat, r2 in r2s.items(): |
| logger.info(f" {stat:25s}: {r2:.4f}") |
|
|
| |
| import os, json |
| log_dir = os.path.join(cfg.log_dir) |
| os.makedirs(log_dir, exist_ok=True) |
|
|
| results = { |
| "model_variant": model_variant, |
| "arch": "xgboost", |
| "best_epoch": 200, |
| "best_val_loss": round(mse, 6), |
| "final_val_r2": r2s, |
| "mean_val_r2": round(mean_r2, 4), |
| "n_params": 0, |
| "n_train_samples": meta["n_train_samples"], |
| "n_val_samples": meta["n_val_samples"], |
| "total_time_s": round(total_time, 1), |
| "device": "cpu", |
| "checkpoint_path": "xgboost (no checkpoint)", |
| } |
|
|
| with open(os.path.join(log_dir, f"results_xgboost_{model_variant.lower()}.json"), "w") as f: |
| json.dump(results, f, indent=2) |
|
|
| return results |
|
|
|
|
| def train_one_model( |
| cfg: TrainConfig, |
| model_variant: str, |
| device: torch.device | None = None, |
| arch: str = "transformer", |
| ) -> dict: |
| """Train a single model (A or B) end-to-end. |
| |
| Args: |
| cfg: Training configuration |
| model_variant: "A" or "B" |
| device: PyTorch device (auto-detected if None) |
| arch: "transformer", "mlp", or "xgboost" |
| |
| Returns: |
| Dict with training results, paths, and final metrics. |
| """ |
| |
| if arch == "xgboost": |
| return train_xgboost(cfg, model_variant) |
|
|
| if device is None: |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
|
|
| logger.info(f"\n{'='*70}") |
| logger.info(f"TRAINING MODEL {model_variant} ({arch}) on {device}") |
| logger.info(f"{'='*70}") |
|
|
| |
| train_ds, val_ds, x_norm, y_norm, meta = build_datasets(cfg, model_variant) |
|
|
| train_loader = DataLoader( |
| train_ds, batch_size=cfg.batch_size, shuffle=True, |
| num_workers=0, pin_memory=(device.type == "cuda"), |
| ) |
| val_loader = DataLoader( |
| val_ds, batch_size=cfg.batch_size * 2, shuffle=False, |
| num_workers=0, pin_memory=(device.type == "cuda"), |
| ) |
|
|
| logger.info( |
| f"Data: {meta['n_train_samples']} train, {meta['n_val_samples']} val " |
| f"({meta['n_train_circuits']} / {meta['n_val_circuits']} circuits)" |
| ) |
|
|
| |
| if model_variant == "A": |
| model = build_model_a(cfg, arch=arch) |
| else: |
| model = build_model_b(cfg, arch=arch) |
| model = model.to(device) |
|
|
| n_params = model.count_params() |
| logger.info(f"Model {model_variant} ({arch}): {n_params:,} parameters") |
|
|
| |
| optimizer = torch.optim.AdamW( |
| model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay |
| ) |
|
|
| steps_per_epoch = math.ceil(len(train_ds) / cfg.batch_size) |
| total_steps = steps_per_epoch * cfg.max_epochs |
| scheduler = CosineWarmupScheduler(optimizer, cfg.warmup_steps, total_steps) |
|
|
| criterion = nn.MSELoss() |
|
|
| |
| best_val_loss = float("inf") |
| best_epoch = 0 |
| patience_counter = 0 |
| history = [] |
|
|
| |
| ckpt_dir = Path(cfg.checkpoint_dir) / f"model_{model_variant.lower()}" |
| os.makedirs(ckpt_dir, exist_ok=True) |
|
|
| t_start = time.time() |
|
|
| for epoch in range(1, cfg.max_epochs + 1): |
| model.train() |
| epoch_loss = 0.0 |
| n_batches = 0 |
|
|
| for X_batch, Y_batch in train_loader: |
| X_batch = X_batch.to(device) |
| Y_batch = Y_batch.to(device) |
|
|
| preds = model(X_batch) |
| loss = criterion(preds, Y_batch) |
|
|
| optimizer.zero_grad() |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_clip) |
| optimizer.step() |
| scheduler.step() |
|
|
| epoch_loss += loss.item() |
| n_batches += 1 |
|
|
| train_loss = epoch_loss / max(n_batches, 1) |
|
|
| |
| val_loss, val_r2 = evaluate(model, val_loader, criterion, device) |
|
|
| |
| elapsed = time.time() - t_start |
| mean_r2 = np.mean(list(val_r2.values())) |
| lr_now = scheduler.get_lr() |
|
|
| entry = { |
| "epoch": epoch, |
| "train_loss": round(train_loss, 6), |
| "val_loss": round(val_loss, 6), |
| "mean_val_r2": round(float(mean_r2), 4), |
| "lr": round(lr_now, 8), |
| "elapsed_s": round(elapsed, 1), |
| } |
| history.append(entry) |
|
|
| if epoch % 10 == 0 or epoch <= 5 or epoch == cfg.max_epochs: |
| logger.info( |
| f" Epoch {epoch:3d} | train={train_loss:.5f} val={val_loss:.5f} " |
| f"RΒ²={mean_r2:.4f} lr={lr_now:.2e} [{elapsed:.0f}s]" |
| ) |
|
|
| |
| if val_loss < best_val_loss: |
| best_val_loss = val_loss |
| best_epoch = epoch |
| patience_counter = 0 |
|
|
| |
| model_config = { |
| "model_variant": model_variant, |
| "n_features": model.n_features, |
| "n_outputs": model.n_outputs, |
| "arch": "transformer" if isinstance(model, CircuitTransformer) else "mlp", |
| } |
| if isinstance(model, CircuitTransformer): |
| model_config.update({ |
| "d_model": cfg.d_model, |
| "n_heads": cfg.n_heads, |
| "n_layers": cfg.n_layers, |
| "d_ff": cfg.d_ff, |
| "dropout": cfg.dropout, |
| "has_ach": model.has_ach, |
| }) |
| else: |
| model_config.update({ |
| "hidden_dims": cfg.mlp_hidden, |
| "dropout": cfg.mlp_dropout, |
| }) |
| torch.save( |
| { |
| "epoch": epoch, |
| "model_state_dict": model.state_dict(), |
| "optimizer_state_dict": optimizer.state_dict(), |
| "val_loss": val_loss, |
| "val_r2": val_r2, |
| "config": model_config, |
| "x_norm": x_norm.state_dict(), |
| "y_norm": y_norm.state_dict(), |
| "meta": meta, |
| }, |
| ckpt_dir / "best.pt", |
| ) |
| else: |
| patience_counter += 1 |
| if patience_counter >= cfg.patience: |
| logger.info( |
| f" Early stopping at epoch {epoch} " |
| f"(best val_loss={best_val_loss:.5f} at epoch {best_epoch})" |
| ) |
| break |
|
|
| total_time = time.time() - t_start |
|
|
| |
| best_ckpt = torch.load(ckpt_dir / "best.pt", map_location=device, weights_only=False) |
| model.load_state_dict(best_ckpt["model_state_dict"]) |
| final_val_loss, final_r2 = evaluate(model, val_loader, criterion, device) |
|
|
| logger.info(f"\n{'='*70}") |
| logger.info(f"MODEL {model_variant} TRAINING COMPLETE") |
| logger.info(f" Best epoch: {best_epoch}, Val loss: {final_val_loss:.5f}") |
| logger.info(f" Mean RΒ²: {np.mean(list(final_r2.values())):.4f}") |
| logger.info(f" Per-stat RΒ²:") |
| for stat, r2 in final_r2.items(): |
| logger.info(f" {stat:25s}: {r2:.4f}") |
| logger.info(f" Total time: {total_time:.0f}s ({total_time/60:.1f} min)") |
| logger.info(f" Params: {n_params:,}") |
| logger.info(f" Checkpoint: {ckpt_dir / 'best.pt'}") |
| logger.info(f"{'='*70}\n") |
|
|
| |
| log_dir = Path(cfg.log_dir) |
| os.makedirs(log_dir, exist_ok=True) |
| with open(log_dir / f"history_model_{model_variant.lower()}.json", "w") as f: |
| json.dump(history, f, indent=2) |
|
|
| |
| results = { |
| "model_variant": model_variant, |
| "best_epoch": best_epoch, |
| "best_val_loss": round(best_val_loss, 6), |
| "final_val_r2": final_r2, |
| "mean_val_r2": round(float(np.mean(list(final_r2.values()))), 4), |
| "n_params": n_params, |
| "n_train_samples": meta["n_train_samples"], |
| "n_val_samples": meta["n_val_samples"], |
| "total_time_s": round(total_time, 1), |
| "device": str(device), |
| "checkpoint_path": str(ckpt_dir / "best.pt"), |
| } |
| with open(log_dir / f"results_model_{model_variant.lower()}.json", "w") as f: |
| json.dump(results, f, indent=2) |
|
|
| return results |
|
|
|
|
| def train_both_models(cfg: TrainConfig, device: torch.device | None = None) -> dict: |
| """Train both Model A and Model B sequentially. |
| |
| Returns dict with results for both models. |
| """ |
| logger.info("=" * 70) |
| logger.info("TRAINING BOTH MODELS: A (plain HH) + B (HH+ACh)") |
| logger.info("=" * 70) |
|
|
| results_a = train_one_model(cfg, "A", device) |
| results_b = train_one_model(cfg, "B", device) |
|
|
| |
| logger.info("\n" + "=" * 70) |
| logger.info("COMPARISON: Model A vs Model B") |
| logger.info("=" * 70) |
| logger.info(f" {'Statistic':25s} {'Model A RΒ²':>12s} {'Model B RΒ²':>12s} {'Ξ (B-A)':>10s}") |
| logger.info(f" {'-'*25} {'-'*12} {'-'*12} {'-'*10}") |
| for stat in OUTPUT_STATS: |
| r2_a = results_a["final_val_r2"].get(stat, 0) |
| r2_b = results_b["final_val_r2"].get(stat, 0) |
| delta = r2_b - r2_a |
| logger.info(f" {stat:25s} {r2_a:12.4f} {r2_b:12.4f} {delta:+10.4f}") |
| logger.info(f" {'MEAN':25s} {results_a['mean_val_r2']:12.4f} {results_b['mean_val_r2']:12.4f} {results_b['mean_val_r2'] - results_a['mean_val_r2']:+10.4f}") |
| logger.info("=" * 70) |
|
|
| return {"model_a": results_a, "model_b": results_b} |
|
|