cnjfl evidence: val_loss and mean_auc disagree.
val_loss best at e3 (0.5592)
mean_auc best at e4 (0.7670 — new h300 + h6000 peaks)
For downstream trading, ranking quality (AUC) matters more than
probability calibration (BCE loss). New default is mean_auc-based
early stopping, but val_loss/none remain selectable.
AUC is noisier than loss epoch-to-epoch (1-2pt bounces are common
even when long-horizon AUCs are still drifting up under
auto-horizon-weights), so patience defaults bump from 3 → 5.
CLI: --early-stop-metric {val_loss|mean_auc|none} default mean_auc
--early-stop-patience N default 5
Argo template parameters added with matching defaults.
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
458 lines
18 KiB
Rust
458 lines
18 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,
|
||
|
||
#[arg(long, default_value_t = 32)]
|
||
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: "val_loss" (stable, tracks
|
||
/// probability calibration) or "mean_auc" (noisier, tracks the
|
||
/// ranking quality used for downstream trading). mean_auc needs
|
||
/// higher patience (~5) because it bounces 1-2pt epoch-to-epoch
|
||
/// even when the model is still genuinely improving on the long
|
||
/// horizons that the auto-horizon-weights down-weight.
|
||
#[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>,
|
||
|
||
/// Compute per-horizon weights as `min(1, K/h)` — short horizons
|
||
/// at weight 1, long horizons down-weighted to their independent-
|
||
/// sample density inside the K-snapshot window. Recommended when
|
||
/// horizons include values larger than seq_len.
|
||
#[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,
|
||
}
|
||
|
||
#[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],
|
||
/// 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" | "none"),
|
||
"--early-stop-metric must be one of: val_loss, mean_auc, 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];
|
||
|
||
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;
|
||
|
||
for epoch in 0..cli.epochs {
|
||
let mut train_loader = MultiHorizonLoader::new(&MultiHorizonLoaderConfig {
|
||
mbp10_root: cli.mbp10_data_dir.clone(),
|
||
predecoded_dir: cli.predecoded_dir.clone(),
|
||
seq_len: cli.seq_len,
|
||
horizons,
|
||
n_max_sequences: cli.n_train_seqs,
|
||
seed: cli.seed.wrapping_add(epoch as u64),
|
||
})
|
||
.context("train loader")?;
|
||
|
||
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.
|
||
let mut val_loader = MultiHorizonLoader::new(&MultiHorizonLoaderConfig {
|
||
mbp10_root: cli.mbp10_data_dir.clone(),
|
||
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 + epoch as u64),
|
||
})
|
||
.context("val loader")?;
|
||
|
||
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"
|
||
);
|
||
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;
|
||
}
|
||
|
||
// 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),
|
||
_ => (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,
|
||
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(())
|
||
}
|