Add inference_only flag to MultiHorizonLoaderConfig that skips per-file forward-label precomputation (~half the file-load cost), plus peek_first() and next_inference_input() chronological-streaming methods for the ml-backtesting LOB harness. - min_size relaxed to cfg.seq_len when inference_only=true (training still requires seq_len + max_horizon + 1 for label generation) - New cursor fields (inference_file_idx, inference_snap_idx) walk every loaded snapshot in chronological order; reset() zeros both - peek_first() seeds CfcTrunk::capture_graph_a with cur==prev semantics (prev_ts_ns==ts_ns, trade_signed_vol=0) — natural stream-start - next_inference_input() errors if cfg.inference_only=false (guard against accidental mixing of training/inference paths) - All trainer call-sites (alpha_train example + multi_horizon_loader tests) updated with inference_only: false (zero behaviour change) - Inline test module exercises both modes; tests skip gracefully when fixture data isn't populated rather than panicking See docs/superpowers/specs/2026-05-18-real-lob-integration-design.md §1 (trainer parity) + §7 (orchestrator). Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
648 lines
28 KiB
Rust
648 lines
28 KiB
Rust
//! Stacked Mamba2 -> CfC -> heads trainer CLI.
|
||
//!
|
||
//! Reads predecoded MBP-10 sidecars via `MultiHorizonLoader`, drives
|
||
//! `PerceptionTrainer` end-to-end. Each training step consumes ONE
|
||
//! seq_len-snapshot window and emits one BCE loss against the labels
|
||
//! at the window's last position (per-horizon).
|
||
//!
|
||
//! Emits `alpha_train_summary.json` on exit for downstream gate
|
||
//! consumption.
|
||
//!
|
||
//! Usage:
|
||
//! alpha_train \
|
||
//! --mbp10-data-dir <path> \
|
||
//! --predecoded-dir <path> \
|
||
//! --epochs 5 \
|
||
//! --seq-len 32 \
|
||
//! --mamba2-state-dim 16 \
|
||
//! --lr-cfc 3e-3 \
|
||
//! --lr-mamba2 1e-3 \
|
||
//! --n-train-seqs 8000 \
|
||
//! --n-val-seqs 1000 \
|
||
//! --out artifacts/alpha-train/
|
||
|
||
use anyhow::{Context, Result};
|
||
use clap::Parser;
|
||
use ml_alpha::cfc::snap_features::Mbp10RawInput;
|
||
use ml_alpha::data::loader::{MultiHorizonLoader, MultiHorizonLoaderConfig};
|
||
use ml_alpha::eval::auc::{compute_auc, AucInput};
|
||
use ml_alpha::heads::N_HORIZONS;
|
||
use ml_alpha::trainer::perception::{auto_horizon_weights, PerceptionTrainer, PerceptionTrainerConfig};
|
||
use ml_core::device::MlDevice;
|
||
use serde::Serialize;
|
||
use std::path::PathBuf;
|
||
|
||
#[derive(Parser)]
|
||
#[command(name = "alpha_train")]
|
||
struct Cli {
|
||
/// Directory containing MBP-10 .dbn.zst files.
|
||
#[arg(long)]
|
||
mbp10_data_dir: PathBuf,
|
||
|
||
/// Directory holding predecoded sidecar caches (created on miss).
|
||
#[arg(long)]
|
||
predecoded_dir: PathBuf,
|
||
|
||
/// Output directory for weights + summary JSON.
|
||
#[arg(long)]
|
||
out: PathBuf,
|
||
|
||
#[arg(long, default_value_t = 5)]
|
||
epochs: usize,
|
||
|
||
/// Snapshots per training sequence (Mamba2 + CfC in-window context).
|
||
/// Default 64 — empirically the 3-fold ISV CV ran K=32 and observed
|
||
/// h6000 (6000-snapshot horizon) saturating around 0.70-0.74. K=32
|
||
/// is only 0.5% of the h6000 prediction window; K=64 doubles
|
||
/// in-window context and gives Mamba2's SSM state more material
|
||
/// to build long-horizon predictions from. Kernel maximum is 96
|
||
/// (MAMBA2_KERNEL_SEQ_MAX). Per-epoch wall scales ~K (more
|
||
/// K-loop launches in the captured graph).
|
||
#[arg(long, default_value_t = 64)]
|
||
seq_len: usize,
|
||
|
||
#[arg(long, default_value_t = 16)]
|
||
mamba2_state_dim: usize,
|
||
|
||
#[arg(long, default_value_t = 3e-3)]
|
||
lr_cfc: f32,
|
||
|
||
#[arg(long, default_value_t = 1e-3)]
|
||
lr_mamba2: f32,
|
||
|
||
#[arg(long, default_value_t = 8000)]
|
||
n_train_seqs: usize,
|
||
|
||
#[arg(long, default_value_t = 1000)]
|
||
n_val_seqs: usize,
|
||
|
||
#[arg(long, default_value_t = 0x4242)]
|
||
seed: u64,
|
||
|
||
/// Linear LR warmup over this many training steps from 0 → configured
|
||
/// lr_cfc / lr_mamba2. 0 disables warmup.
|
||
#[arg(long, default_value_t = 200)]
|
||
lr_warmup_steps: usize,
|
||
|
||
/// Floor for the cosine-decay schedule: final LR = lr_max * this factor.
|
||
/// 1.0 disables decay (constant LR after warmup).
|
||
#[arg(long, default_value_t = 0.1)]
|
||
lr_min_factor: f32,
|
||
|
||
/// Early-stop training after this many consecutive epochs without
|
||
/// improvement on `early_stop_metric`. 0 disables early stopping.
|
||
#[arg(long, default_value_t = 5)]
|
||
early_stop_patience: usize,
|
||
|
||
/// Which metric drives early stopping. Options:
|
||
/// - "val_loss" : stable, tracks probability calibration.
|
||
/// - "mean_auc" : noisier, average ranking across all 5 horizons.
|
||
/// - "auc_h6000" : ranking on the long horizon ONLY — recommended
|
||
/// for deployment-oriented runs where the saved checkpoint is
|
||
/// used to trade at multi-minute horizon. The 3-fold ISV CV
|
||
/// showed mean_auc-best-epoch and h6000-best-epoch can differ
|
||
/// by 2-6pt on h6000 within a single run (e.g. fblb2 fold-1
|
||
/// saved h6000=0.681 at mean_auc-best E10 but the trajectory
|
||
/// peaked at h6000=0.739 one epoch later). Using auc_h6000
|
||
/// directly saves the deployment-relevant checkpoint.
|
||
#[arg(long, default_value = "mean_auc")]
|
||
early_stop_metric: String,
|
||
|
||
/// Custom per-horizon BCE weights (5 floats, comma-separated). When
|
||
/// set, overrides the auto/uniform defaults. Example: "1.0,1.0,0.5,0.1,0.02".
|
||
#[arg(long)]
|
||
horizon_weights: Option<String>,
|
||
|
||
/// Use the canonical auto-derived per-horizon BCE weight schedule.
|
||
/// Currently identical to `--horizon-weights 1,1,1,1,1` (uniform):
|
||
/// the prior `min(1, K/h)` formula down-weighted h6000 to 0.5% of
|
||
/// the loss which suppressed long-horizon training. With uniform
|
||
/// weights, ISV's lambda controller (per-horizon trunk-gradient
|
||
/// scaler driven by BCE EMA) is the single place where per-horizon
|
||
/// rebalancing happens — cleaner than splitting between BCE
|
||
/// coefficients and the trunk lambda. Flag preserved for backward
|
||
/// compat; passing `--horizon-weights` still overrides.
|
||
#[arg(long, default_value_t = false)]
|
||
auto_horizon_weights: bool,
|
||
|
||
/// Mini-batch size for training. Each step processes B sequences
|
||
/// in parallel via batched CUDA kernels. Default 1 matches the
|
||
/// previous unbatched behavior. Validation always runs at B=1.
|
||
#[arg(long, default_value_t = 1)]
|
||
batch_size: usize,
|
||
|
||
/// Sliding-window cross-validation: index of THIS fold (0-indexed).
|
||
/// Combined with `--cv-n-folds` and `--cv-train-window` to split
|
||
/// the discovered MBP-10 files chronologically. Default 0 means
|
||
/// "no CV — use the first fold of a 1-fold split" (legacy behavior:
|
||
/// train = all-but-last file, val = last file).
|
||
#[arg(long, default_value_t = 0)]
|
||
cv_fold: usize,
|
||
|
||
/// Total number of sliding-window CV folds. 1 disables CV (single
|
||
/// train/val split). When >1, fold k trains on the W files starting
|
||
/// at `cv_fold_offset(k)` and validates on the file immediately
|
||
/// after that window. Folds are submitted as independent runs.
|
||
#[arg(long, default_value_t = 1)]
|
||
cv_n_folds: usize,
|
||
|
||
/// Size of each fold's TRAIN window in files. 0 means "use all
|
||
/// files except the val file" (expanding-window mode at 1 fold).
|
||
/// >0 enables sliding-window mode where the training data has a
|
||
/// fixed temporal extent.
|
||
#[arg(long, default_value_t = 0)]
|
||
cv_train_window: usize,
|
||
|
||
/// Decision-stride sampling: yield every S-th snapshot per training
|
||
/// sequence. S=1 (default) preserves current consecutive-snapshot
|
||
/// behaviour. S>1 expands the effective time-window covered by each
|
||
/// sequence (window = (seq_len-1) * S + 1 snapshots) at the same
|
||
/// compute cost. Pairs with Mamba2 dt_s scaling so the SSM's state
|
||
/// decay reflects the actual elapsed time between K-positions.
|
||
/// Recommended S=4 with K=64 for h6000 deployment (256-tick window).
|
||
#[arg(long, default_value_t = 1)]
|
||
decision_stride: usize,
|
||
}
|
||
|
||
#[derive(Serialize)]
|
||
struct AlphaTrainSummary {
|
||
epochs: usize,
|
||
seq_len: usize,
|
||
horizons: [usize; N_HORIZONS],
|
||
final_train_loss: f32,
|
||
final_val_loss: f32,
|
||
final_val_auc: [f32; N_HORIZONS],
|
||
n_train_seqs_consumed: usize,
|
||
n_val_seqs_consumed: usize,
|
||
/// Epoch index that achieved the lowest val_loss.
|
||
best_epoch: usize,
|
||
/// Lowest val_loss observed across all epochs.
|
||
best_val_loss: f32,
|
||
/// Per-horizon AUCs at the best-val-loss epoch.
|
||
best_val_auc: [f32; N_HORIZONS],
|
||
/// Epoch index that achieved the highest mean AUC across the 5 horizons.
|
||
/// Tracked independently of best_epoch because val_loss + AUC can
|
||
/// disagree — a calibration-flat / ranking-sharp epoch beats a
|
||
/// calibration-sharp / ranking-flat one for downstream trading.
|
||
best_mean_auc_epoch: usize,
|
||
/// Mean AUC at the best-mean-AUC epoch.
|
||
best_mean_auc: f32,
|
||
/// Per-horizon AUCs at the best-mean-AUC epoch.
|
||
best_mean_auc_per_horizon: [f32; N_HORIZONS],
|
||
/// Epoch with the highest h6000 AUC (long-horizon, deployment-relevant).
|
||
/// Saved independently of best_mean_auc_epoch because the two
|
||
/// frequently disagree by 1-2 epochs and h6000 is the metric we
|
||
/// care about for multi-minute trading deployment.
|
||
best_auc_h6000_epoch: usize,
|
||
/// Highest h6000 AUC observed across all epochs.
|
||
best_auc_h6000: f32,
|
||
/// Per-horizon AUCs at the best-h6000 epoch.
|
||
best_auc_h6000_per_horizon: [f32; N_HORIZONS],
|
||
/// True if training stopped early (patience exceeded).
|
||
early_stopped: bool,
|
||
}
|
||
|
||
/// Linear warmup then cosine decay to `lr_min`. `step_idx` is the
|
||
/// global training-step counter (0-indexed). `total_steps` is the
|
||
/// budget the cosine targets.
|
||
fn lr_schedule(lr_max: f32, lr_min: f32, step_idx: usize, warmup_steps: usize, total_steps: usize) -> f32 {
|
||
if step_idx < warmup_steps {
|
||
return lr_max * (step_idx as f32 + 1.0) / warmup_steps as f32;
|
||
}
|
||
if total_steps <= warmup_steps {
|
||
return lr_max;
|
||
}
|
||
let progress = (step_idx - warmup_steps) as f32 / (total_steps - warmup_steps) as f32;
|
||
let progress = progress.clamp(0.0, 1.0);
|
||
let cos = 0.5 * (1.0 + (std::f32::consts::PI * progress).cos());
|
||
lr_min + (lr_max - lr_min) * cos
|
||
}
|
||
|
||
fn main() -> Result<()> {
|
||
tracing_subscriber::fmt()
|
||
.with_env_filter(tracing_subscriber::EnvFilter::from_default_env())
|
||
.init();
|
||
|
||
let cli = Cli::parse();
|
||
std::fs::create_dir_all(&cli.out).with_context(|| format!("mkdir {}", cli.out.display()))?;
|
||
|
||
let dev = MlDevice::cuda(0).context("CUDA 0 init")?;
|
||
tracing::info!(?dev, "MlDevice initialized");
|
||
|
||
let horizons = [30usize, 100, 300, 1000, 6000];
|
||
|
||
// Resolve per-horizon BCE weights: explicit > auto > uniform.
|
||
let horizon_weights: [f32; N_HORIZONS] = if let Some(spec) = &cli.horizon_weights {
|
||
let parts: Vec<f32> = spec.split(',')
|
||
.map(|s| s.trim().parse().context("horizon weight parse"))
|
||
.collect::<Result<Vec<_>>>()?;
|
||
anyhow::ensure!(
|
||
parts.len() == N_HORIZONS,
|
||
"--horizon-weights expected {} floats, got {}",
|
||
N_HORIZONS, parts.len()
|
||
);
|
||
let mut w = [0.0; N_HORIZONS];
|
||
w.copy_from_slice(&parts);
|
||
w
|
||
} else if cli.auto_horizon_weights {
|
||
auto_horizon_weights(cli.seq_len, &horizons)
|
||
} else {
|
||
[1.0; N_HORIZONS]
|
||
};
|
||
tracing::info!(
|
||
w_h30 = horizon_weights[0], w_h100 = horizon_weights[1],
|
||
w_h300 = horizon_weights[2], w_h1000 = horizon_weights[3],
|
||
w_h6000 = horizon_weights[4],
|
||
"per-horizon BCE weights"
|
||
);
|
||
|
||
anyhow::ensure!(cli.batch_size >= 1, "batch_size must be >= 1");
|
||
let trainer_cfg = PerceptionTrainerConfig {
|
||
seq_len: cli.seq_len,
|
||
mamba2_state_dim: cli.mamba2_state_dim,
|
||
lr_cfc: cli.lr_cfc,
|
||
lr_mamba2: cli.lr_mamba2,
|
||
seed: cli.seed,
|
||
horizon_weights,
|
||
n_batch: cli.batch_size,
|
||
decision_stride: cli.decision_stride,
|
||
};
|
||
let mut trainer = PerceptionTrainer::new(&dev, &trainer_cfg).context("trainer init")?;
|
||
|
||
let mut train_loss_running = 0.0_f32;
|
||
let mut train_steps = 0usize;
|
||
let mut final_val_loss = 0.0_f32;
|
||
let mut final_val_auc = [0.5_f32; N_HORIZONS];
|
||
let mut n_val_seqs_consumed = 0usize;
|
||
|
||
// Validate early-stop metric choice upfront.
|
||
let early_stop_metric = cli.early_stop_metric.as_str();
|
||
anyhow::ensure!(
|
||
matches!(early_stop_metric, "val_loss" | "mean_auc" | "auc_h6000" | "none"),
|
||
"--early-stop-metric must be one of: val_loss, mean_auc, auc_h6000, none (got `{}`)",
|
||
early_stop_metric
|
||
);
|
||
|
||
// Best-checkpoint tracking (by val_loss).
|
||
let mut best_val_loss = f32::INFINITY;
|
||
let mut best_epoch = 0usize;
|
||
let mut best_val_auc = [0.5_f32; N_HORIZONS];
|
||
let mut val_loss_no_improvement = 0usize;
|
||
let mut mean_auc_no_improvement = 0usize;
|
||
let mut early_stopped = false;
|
||
let mut epochs_completed = 0usize;
|
||
|
||
// Independent tracker for best mean-AUC across horizons. val_loss and
|
||
// ranking quality (AUC) can disagree — e.g. when the model improves
|
||
// probability calibration on common cases (lower BCE) but loses
|
||
// ranking sharpness on the regime tails. For downstream trading the
|
||
// AUC profile usually matters more, so we publish both bests.
|
||
let mut best_mean_auc = f32::NEG_INFINITY;
|
||
let mut best_mean_auc_epoch = 0usize;
|
||
let mut best_mean_auc_per_horizon = [0.5_f32; N_HORIZONS];
|
||
|
||
// Independent tracker for best h6000 AUC — the deployment-relevant
|
||
// long-horizon metric. The 3-fold ISV CV showed that mean_auc-best
|
||
// and h6000-best epochs frequently differ (within fblb2 fold-1, the
|
||
// saved E10 had h6000=0.681 while E11 had h6000=0.739). For
|
||
// multi-minute trading deployment we want the h6000-best checkpoint
|
||
// directly, not the mean_auc proxy.
|
||
let mut best_auc_h6000 = f32::NEG_INFINITY;
|
||
let mut best_auc_h6000_epoch = 0usize;
|
||
let mut best_auc_h6000_per_horizon = [0.5_f32; N_HORIZONS];
|
||
let mut auc_h6000_no_improvement = 0usize;
|
||
|
||
let lr_min_cfc = cli.lr_cfc * cli.lr_min_factor;
|
||
let lr_min_m2 = cli.lr_mamba2 * cli.lr_min_factor;
|
||
// Approximate the total step budget: epochs × n_train_seqs (each
|
||
// sequence yields one optimizer step). Used by cosine decay.
|
||
let total_steps_budget = cli.epochs * cli.n_train_seqs;
|
||
|
||
// ── Walk-forward CV file split ────────────────────────────────────
|
||
// Files are discovered in chronological order (filenames follow
|
||
// ES.FUT_<YEAR>-Q<n>.dbn.zst so sort = chronological). Then we
|
||
// slice train_files / val_files based on the CV flags. Train and
|
||
// val NEVER share a file — every val sample is strictly OOS in
|
||
// time vs every train sample.
|
||
let all_files = ml_alpha::data::loader::discover_mbp10_files_sorted(&cli.mbp10_data_dir)
|
||
.context("discover MBP-10 files")?;
|
||
anyhow::ensure!(
|
||
cli.cv_n_folds >= 1,
|
||
"--cv-n-folds must be >= 1 (got {})", cli.cv_n_folds
|
||
);
|
||
anyhow::ensure!(
|
||
cli.cv_fold < cli.cv_n_folds,
|
||
"--cv-fold {} must be < --cv-n-folds {}", cli.cv_fold, cli.cv_n_folds
|
||
);
|
||
let n_files = all_files.len();
|
||
let (train_files, val_files): (Vec<PathBuf>, Vec<PathBuf>) = if cli.cv_n_folds == 1 {
|
||
// Single fold: train on all files except the LAST, val on the LAST file.
|
||
// Legacy callers that didn't think about CV at all get a temporal split
|
||
// by default instead of the old "same files for train and val" bug.
|
||
anyhow::ensure!(n_files >= 2,
|
||
"need >= 2 files for single-fold temporal split, found {} under {}",
|
||
n_files, cli.mbp10_data_dir.display());
|
||
let val_idx = n_files - 1;
|
||
(
|
||
all_files[..val_idx].to_vec(),
|
||
vec![all_files[val_idx].clone()],
|
||
)
|
||
} else {
|
||
// Sliding-window CV. Train window has fixed extent `W`; val is the
|
||
// single file immediately after. Fold k uses files [k .. k+W] for
|
||
// train and file [k+W] for val.
|
||
let w = if cli.cv_train_window > 0 {
|
||
cli.cv_train_window
|
||
} else {
|
||
// Default: leave room for n_folds val files at the tail.
|
||
n_files.saturating_sub(cli.cv_n_folds)
|
||
};
|
||
anyhow::ensure!(w >= 1, "cv_train_window resolves to 0 (n_files={n_files}, n_folds={})", cli.cv_n_folds);
|
||
anyhow::ensure!(
|
||
cli.cv_fold + w + 1 <= n_files,
|
||
"CV layout overflows: cv_fold={} + window={} + 1 val file > {} files available",
|
||
cli.cv_fold, w, n_files
|
||
);
|
||
let train_start = cli.cv_fold;
|
||
let train_end = train_start + w; // exclusive
|
||
let val_idx = train_end;
|
||
(
|
||
all_files[train_start..train_end].to_vec(),
|
||
vec![all_files[val_idx].clone()],
|
||
)
|
||
};
|
||
tracing::info!(
|
||
cv_fold = cli.cv_fold,
|
||
cv_n_folds = cli.cv_n_folds,
|
||
cv_train_window = cli.cv_train_window,
|
||
n_train_files = train_files.len(),
|
||
n_val_files = val_files.len(),
|
||
train_files = ?train_files.iter().map(|p| p.file_name().unwrap().to_string_lossy().to_string()).collect::<Vec<_>>(),
|
||
val_files = ?val_files.iter().map(|p| p.file_name().unwrap().to_string_lossy().to_string()).collect::<Vec<_>>(),
|
||
"walk-forward CV file split",
|
||
);
|
||
|
||
// PRELOAD: construct loaders ONCE (each preloads its files into RAM).
|
||
// Per-epoch we call reset(seed) to re-shuffle anchor sampling — no
|
||
// disk IO. Cuts wall time ~5× on a 6-epoch run, ~7× on 15 epochs.
|
||
tracing::info!("preloading train + val MBP-10 files into RAM (slow startup, fast per-epoch)…");
|
||
let preload_start = std::time::Instant::now();
|
||
let mut train_loader = MultiHorizonLoader::new(&MultiHorizonLoaderConfig {
|
||
files: train_files,
|
||
predecoded_dir: cli.predecoded_dir.clone(),
|
||
seq_len: cli.seq_len,
|
||
horizons,
|
||
n_max_sequences: cli.n_train_seqs,
|
||
seed: cli.seed,
|
||
decision_stride: cli.decision_stride,
|
||
inference_only: false,
|
||
})
|
||
.context("train loader")?;
|
||
let mut val_loader = MultiHorizonLoader::new(&MultiHorizonLoaderConfig {
|
||
files: val_files,
|
||
predecoded_dir: cli.predecoded_dir.clone(),
|
||
seq_len: cli.seq_len,
|
||
horizons,
|
||
n_max_sequences: cli.n_val_seqs,
|
||
seed: cli.seed.wrapping_add(0xC0FFEE),
|
||
decision_stride: cli.decision_stride,
|
||
inference_only: false,
|
||
})
|
||
.context("val loader")?;
|
||
tracing::info!(
|
||
elapsed_s = preload_start.elapsed().as_secs(),
|
||
"preload complete",
|
||
);
|
||
|
||
for epoch in 0..cli.epochs {
|
||
// Re-seed loaders for this epoch (anchor sampling differs per epoch).
|
||
train_loader.reset(cli.seed.wrapping_add(epoch as u64));
|
||
val_loader.reset(cli.seed.wrapping_add(0xC0FFEE + epoch as u64));
|
||
|
||
let mut epoch_train_loss = 0.0_f32;
|
||
let mut epoch_train_steps = 0usize;
|
||
// Accumulate B sequences per optimizer step. Sequences with no
|
||
// finite labels at any position are skipped (continue) so they
|
||
// don't pollute the batch.
|
||
let mut snap_batch: Vec<Vec<Mbp10RawInput>> = Vec::with_capacity(cli.batch_size);
|
||
let mut label_batch: Vec<Vec<[f32; N_HORIZONS]>> = Vec::with_capacity(cli.batch_size);
|
||
while let Some(seq) = train_loader.next_sequence().context("train next_seq")? {
|
||
let mut labels_per_pos: Vec<[f32; N_HORIZONS]> = Vec::with_capacity(seq.snapshots.len());
|
||
let mut any_finite = false;
|
||
for k in 0..seq.snapshots.len() {
|
||
let row = [
|
||
seq.labels[0][k], seq.labels[1][k], seq.labels[2][k],
|
||
seq.labels[3][k], seq.labels[4][k],
|
||
];
|
||
if row.iter().any(|v| v.is_finite()) { any_finite = true; }
|
||
labels_per_pos.push(row);
|
||
}
|
||
if !any_finite { continue; }
|
||
snap_batch.push(seq.snapshots);
|
||
label_batch.push(labels_per_pos);
|
||
if snap_batch.len() < cli.batch_size { continue; }
|
||
|
||
// Full batch ready — apply LR schedule, fire step.
|
||
let lr_cfc_now = lr_schedule(cli.lr_cfc, lr_min_cfc, train_steps, cli.lr_warmup_steps, total_steps_budget);
|
||
let lr_m2_now = lr_schedule(cli.lr_mamba2, lr_min_m2, train_steps, cli.lr_warmup_steps, total_steps_budget);
|
||
trainer.set_lr_cfc(lr_cfc_now);
|
||
trainer.set_lr_mamba2(lr_m2_now);
|
||
|
||
let snap_refs: Vec<&[Mbp10RawInput]> = snap_batch.iter().map(|v| v.as_slice()).collect();
|
||
let label_refs: Vec<&[[f32; N_HORIZONS]]> = label_batch.iter().map(|v| v.as_slice()).collect();
|
||
let loss = trainer.step_batched(&snap_refs, &label_refs).context("train step_batched")?;
|
||
epoch_train_loss += loss;
|
||
epoch_train_steps += 1;
|
||
train_loss_running += loss;
|
||
train_steps += 1;
|
||
snap_batch.clear();
|
||
label_batch.clear();
|
||
}
|
||
let epoch_avg = if epoch_train_steps > 0 {
|
||
epoch_train_loss / epoch_train_steps as f32
|
||
} else { 0.0 };
|
||
tracing::info!(epoch, train_loss = epoch_avg, train_steps = epoch_train_steps, "epoch complete");
|
||
|
||
// Validation: accumulate (probs, labels) per horizon, compute AUC.
|
||
// val_loader was preloaded above and reset at top of loop.
|
||
|
||
let mut val_probs: [Vec<f32>; N_HORIZONS] = Default::default();
|
||
let mut val_labels: [Vec<f32>; N_HORIZONS] = Default::default();
|
||
let mut val_loss_sum = 0.0_f32;
|
||
let mut val_steps = 0usize;
|
||
let mut val_snap_batch: Vec<Vec<Mbp10RawInput>> = Vec::with_capacity(cli.batch_size);
|
||
let mut val_label_batch: Vec<Vec<[f32; N_HORIZONS]>> = Vec::with_capacity(cli.batch_size);
|
||
let mut val_last_label_batch: Vec<[f32; N_HORIZONS]> = Vec::with_capacity(cli.batch_size);
|
||
while let Some(seq) = val_loader.next_sequence().context("val next_seq")? {
|
||
let mut labels_per_pos: Vec<[f32; N_HORIZONS]> = Vec::with_capacity(seq.snapshots.len());
|
||
let mut any_finite = false;
|
||
for k in 0..seq.snapshots.len() {
|
||
let row = [
|
||
seq.labels[0][k], seq.labels[1][k], seq.labels[2][k],
|
||
seq.labels[3][k], seq.labels[4][k],
|
||
];
|
||
if row.iter().any(|v| v.is_finite()) { any_finite = true; }
|
||
labels_per_pos.push(row);
|
||
}
|
||
if !any_finite { continue; }
|
||
let last = seq.snapshots.len().saturating_sub(1);
|
||
let last_labels = [
|
||
seq.labels[0][last], seq.labels[1][last], seq.labels[2][last],
|
||
seq.labels[3][last], seq.labels[4][last],
|
||
];
|
||
val_snap_batch.push(seq.snapshots);
|
||
val_label_batch.push(labels_per_pos);
|
||
val_last_label_batch.push(last_labels);
|
||
if val_snap_batch.len() < cli.batch_size { continue; }
|
||
|
||
// FORWARD-ONLY batched evaluation — no backward, no AdamW.
|
||
let snap_refs: Vec<&[Mbp10RawInput]> = val_snap_batch.iter().map(|v| v.as_slice()).collect();
|
||
let label_refs: Vec<&[[f32; N_HORIZONS]]> = val_label_batch.iter().map(|v| v.as_slice()).collect();
|
||
let (l, probs_all) = trainer.evaluate_batched(&snap_refs, &label_refs).context("val eval_batched")?;
|
||
val_loss_sum += l;
|
||
val_steps += 1;
|
||
// probs_all is [K, B, N_HORIZONS] row-major. Score AUC from
|
||
// the LAST-position predictions for each sample in the batch.
|
||
let last = cli.seq_len - 1;
|
||
for (b_idx, last_lbl) in val_last_label_batch.iter().enumerate() {
|
||
let off = (last * cli.batch_size + b_idx) * N_HORIZONS;
|
||
let last_probs = &probs_all[off..off + N_HORIZONS];
|
||
for h in 0..N_HORIZONS {
|
||
if last_lbl[h].is_finite() {
|
||
val_probs[h].push(last_probs[h]);
|
||
val_labels[h].push(last_lbl[h]);
|
||
}
|
||
}
|
||
}
|
||
val_snap_batch.clear();
|
||
val_label_batch.clear();
|
||
val_last_label_batch.clear();
|
||
}
|
||
let mut per_horizon_auc = [0.5_f32; N_HORIZONS];
|
||
for h in 0..N_HORIZONS {
|
||
per_horizon_auc[h] = compute_auc(&AucInput {
|
||
probs: val_probs[h].clone(),
|
||
labels: val_labels[h].clone(),
|
||
}).context("AUC")?;
|
||
}
|
||
let val_avg = if val_steps > 0 { val_loss_sum / val_steps as f32 } else { 0.0 };
|
||
tracing::info!(
|
||
epoch,
|
||
val_loss = val_avg,
|
||
auc_h30 = per_horizon_auc[0],
|
||
auc_h100 = per_horizon_auc[1],
|
||
auc_h300 = per_horizon_auc[2],
|
||
auc_h1000 = per_horizon_auc[3],
|
||
auc_h6000 = per_horizon_auc[4],
|
||
"validation"
|
||
);
|
||
|
||
// Snapshot the ISV controller's per-horizon BCE EMA + lambda
|
||
// multipliers — diagnostic on what the horizon-aware gradient
|
||
// scaler is doing each epoch. Reads happen via mapped-pinned
|
||
// (not in the training hot path; we're already mid-validation
|
||
// sync window). Surfaces fold-specific lambda trajectories so
|
||
// we can diagnose why some folds gain and others regress.
|
||
match (trainer.loss_ema_snapshot(), trainer.lambda_snapshot()) {
|
||
(Ok(ema), Ok(lam)) => tracing::info!(
|
||
epoch,
|
||
ema_h30 = ema[0], ema_h100 = ema[1], ema_h300 = ema[2],
|
||
ema_h1000 = ema[3], ema_h6000 = ema[4],
|
||
lam_h30 = lam[0], lam_h100 = lam[1], lam_h300 = lam[2],
|
||
lam_h1000 = lam[3], lam_h6000 = lam[4],
|
||
"isv snapshot"
|
||
),
|
||
(Err(e), _) | (_, Err(e)) => {
|
||
tracing::warn!(epoch, error = %e, "isv snapshot failed (continuing)");
|
||
}
|
||
}
|
||
final_val_loss = val_avg;
|
||
final_val_auc = per_horizon_auc;
|
||
n_val_seqs_consumed = val_loader.yielded();
|
||
epochs_completed = epoch + 1;
|
||
|
||
// Best-checkpoint trackers — independent for each metric.
|
||
let val_loss_improved = val_avg < best_val_loss;
|
||
if val_loss_improved {
|
||
best_val_loss = val_avg;
|
||
best_epoch = epoch;
|
||
best_val_auc = per_horizon_auc;
|
||
val_loss_no_improvement = 0;
|
||
tracing::info!(epoch, val_loss = val_avg, "new best val_loss");
|
||
} else {
|
||
val_loss_no_improvement += 1;
|
||
}
|
||
|
||
let mean_auc = per_horizon_auc.iter().sum::<f32>() / per_horizon_auc.len() as f32;
|
||
let mean_auc_improved = mean_auc > best_mean_auc;
|
||
if mean_auc_improved {
|
||
best_mean_auc = mean_auc;
|
||
best_mean_auc_epoch = epoch;
|
||
best_mean_auc_per_horizon = per_horizon_auc;
|
||
mean_auc_no_improvement = 0;
|
||
tracing::info!(epoch, mean_auc, "new best mean_auc");
|
||
} else {
|
||
mean_auc_no_improvement += 1;
|
||
}
|
||
|
||
// h6000 — deployment-relevant long-horizon ranking.
|
||
let auc_h6000 = per_horizon_auc[N_HORIZONS - 1];
|
||
let auc_h6000_improved = auc_h6000 > best_auc_h6000;
|
||
if auc_h6000_improved {
|
||
best_auc_h6000 = auc_h6000;
|
||
best_auc_h6000_epoch = epoch;
|
||
best_auc_h6000_per_horizon = per_horizon_auc;
|
||
auc_h6000_no_improvement = 0;
|
||
tracing::info!(epoch, auc_h6000, "new best auc_h6000");
|
||
} else {
|
||
auc_h6000_no_improvement += 1;
|
||
}
|
||
|
||
// Early-stop decision based on selected metric.
|
||
if cli.early_stop_patience > 0 {
|
||
let (counter, metric_label, best_val) = match early_stop_metric {
|
||
"val_loss" => (val_loss_no_improvement, "val_loss", best_val_loss),
|
||
"mean_auc" => (mean_auc_no_improvement, "mean_auc", best_mean_auc),
|
||
"auc_h6000" => (auc_h6000_no_improvement, "auc_h6000", best_auc_h6000),
|
||
_ => (0, "none", 0.0), // disabled
|
||
};
|
||
if counter >= cli.early_stop_patience && metric_label != "none" {
|
||
tracing::warn!(
|
||
epoch, metric = metric_label, best = best_val,
|
||
patience = cli.early_stop_patience,
|
||
"early stopping",
|
||
);
|
||
early_stopped = true;
|
||
break;
|
||
}
|
||
}
|
||
}
|
||
|
||
let summary = AlphaTrainSummary {
|
||
epochs: epochs_completed,
|
||
seq_len: cli.seq_len,
|
||
horizons,
|
||
final_train_loss: if train_steps > 0 { train_loss_running / train_steps as f32 } else { 0.0 },
|
||
final_val_loss,
|
||
final_val_auc,
|
||
n_train_seqs_consumed: train_steps,
|
||
n_val_seqs_consumed,
|
||
best_epoch,
|
||
best_val_loss,
|
||
best_val_auc,
|
||
best_mean_auc_epoch,
|
||
best_mean_auc,
|
||
best_mean_auc_per_horizon,
|
||
best_auc_h6000_epoch,
|
||
best_auc_h6000,
|
||
best_auc_h6000_per_horizon,
|
||
early_stopped,
|
||
};
|
||
let summary_path = cli.out.join("alpha_train_summary.json");
|
||
std::fs::write(&summary_path, serde_json::to_vec_pretty(&summary).context("serialize summary")?)
|
||
.with_context(|| format!("write {}", summary_path.display()))?;
|
||
tracing::info!(summary_path = %summary_path.display(), "wrote summary");
|
||
|
||
Ok(())
|
||
}
|