Files
foxhunt/crates/ml-alpha/examples/alpha_train.rs
jgrusewski 00da163078 feat(ml-alpha): h6000-aligned ISV — uniform BCE + z-score lambda + K=64
Three correlated fixes addressing the architectural inconsistency
surfaced by the 3-fold ISV CV: we built a horizon-aware gradient
controller (ISV) but suppressed its target horizon (h6000) to 0.36%
of the loss via auto-horizon-weights, then used a ratio formula
that never approached its own clamp ceiling. ISV's lambda was
operating on rounding error.

(1) Uniform BCE weights as auto-default
    trainer/perception.rs: `auto_horizon_weights` now returns
    [1.0; 5] regardless of seq_len. Prior schedule `min(1, K/h)`
    gave h6000 weight 0.0053 at K=32 — combined with lambda ~1.04,
    h6000's effective loss contribution was ~0.37%, indistinguishable
    from zero. With uniform weights, each horizon contributes 20% and
    ISV's lambda actually has something to scale.

(2) Z-score lambda derivation
    cuda/horizon_lambda.cu: replace `ratio = ema_h / mean(ema)` with
    `z_h = (ema_h - mean) / std(ema); lambda = clamp(1.0 + 0.5*z, 1.0, 2.0)`.
    Per `pearl_zscore_normalization_for_magnitude_asymmetric_signals.md`
    z-score makes lambda spread scale-invariant of the absolute EMA
    level. The ratio formula gave lambdas ≤ 1.04 in our data because
    per-horizon BCE clusters tightly (range ~0.04) while mean is
    ~0.65. Z-score fills the [1.0, 2.0] envelope: 1σ → 1.5, 2σ →
    ceiling. Boost-only asymmetric clamp preserved.

    Test verification on the existing smoke (after 5 steps):
      ema    = [0.526, 0.522, 0.608, 0.641, 0.553]
      lambda = [1.00, 1.00, 1.40, 1.76, 1.00]
    Previously with ratio formula, max lambda on the same data
    would have been ~1.05. h1000 (1.5σ above mean BCE here) now
    gets a 76% trunk-gradient boost vs uniform.

(3) Default --seq-len 32 → 64
    examples/alpha_train.rs: K=32 gives the model 0.5% of the
    h6000 prediction window as in-window context. K=64 doubles
    that, giving Mamba2's SSM state more material to build
    long-horizon predictions. Within the kernel's MAMBA2_KERNEL_SEQ_MAX
    cap of 96. Per-epoch wall scales ~K (more K-loop launches in
    the captured graph, ~2× wall at K=64 vs K=32 for the K-loop
    portion of dispatch).

Cache-bust v5 in build.rs to force nvcc recompile against the new
horizon_lambda.cu formula on the cluster's /cargo-target PVC. Old
cubins compute a numerically different lambda; running them against
the new Rust loop would silently apply the wrong gradient scaler.

Validation: 7 perception_overfit tests + 26 lib + 23 integration
ml-alpha tests pass. Synthetic overfit still converges to 0.0006.
horizon_ema_and_lambda_track_after_training observes the new
lambda spread (1.0-1.76) and asserts the asymmetric clamp envelope.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-17 20:58:44 +02:00

633 lines
27 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 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,
}
#[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,
};
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,
})
.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),
})
.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(())
}