Per docs/superpowers/specs/2026-05-17-kloop-parallelization-design.md. cfc_step_batched (fwd + bwd) refactored from grid=(1,1,1) with internal n_batch loop to grid=(B,1,1) — each block handles one batch. Removes the single-SM bottleneck on the K-loop's most-called kernel (64×/step). Param-grad accumulation moves to per-batch scratch: cfc_grad_w_in_scratch_d [B, n_hid, n_in] cfc_grad_w_rec_scratch_d [B, n_hid, n_hid] cfc_grad_b_scratch_d [B, n_hid] cfc_grad_tau_scratch_d [B, n_hid] Zeroed once per training step, K-loop's 64 bwd calls += into them, then 4 reduce_axis0 launches collapse B → final grad buffers (OVERWRITE) before AdamW. New AdamW-after-reducer invariant: final grads are meaningful only after the reducer has run in the current step. New reduce_axis0 kernel: single parameterised reducer [B, N] → [N] via block tree-reduce (no atomicAdd per feedback_no_atomicadd.md). Same pattern as layer_norm_reduce_param_grads — CUDA-Graph-safe. cfc_step_backward_batched shared-mem dropped from (B+1)*n_hid*4 to 2*n_hid*4 bytes per block (only one row of sd_pre needed per block bi). Tests: - New stacked_trainer_loss_shrinks_at_batch_32: FIRST test that actually exercises the cross-batch reduction code path; existing perception_overfit suite was all B=1. Initial 0.24 → final 0.00. - Scratch-clears test removed (explanatory comment kept): structurally hard to assert directly due to begin_capture/end_capture not executing kernels; the B=32 convergence smoke implicitly validates scratch zeroing since divergence would otherwise be immediate. All 9 perception_overfit smokes + 4 backward_finite_diff tests pass. build.rs: - KERNELS list adds "reduce_axis0" - Cache-bust → v11 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
463 lines
18 KiB
Rust
463 lines
18 KiB
Rust
//! PerceptionTrainer synthetic-overfit smoke.
|
|
//!
|
|
//! Mirrors `perception_overfit.rs` but on the stacked trainer. The
|
|
//! signal is constant direction=+1, label=[1; 5]. Asserts the BCE loss
|
|
//! shrinks at least 40% over the training budget — proves the
|
|
//! full forward + backward (Mamba2 + CfC + heads) wires up correctly
|
|
//! end-to-end and that all 6 AdamW optimizers actually move weights.
|
|
|
|
use ml_alpha::cfc::snap_features::Mbp10RawInput;
|
|
use ml_alpha::trainer::perception::{PerceptionTrainer, PerceptionTrainerConfig};
|
|
use ml_core::device::MlDevice;
|
|
|
|
fn test_device() -> MlDevice {
|
|
MlDevice::cuda(0).expect("CUDA 0 required for ml-alpha tests")
|
|
}
|
|
|
|
fn synthetic_seq(
|
|
seq_len: usize,
|
|
mut prev_mid: f32,
|
|
mut ts_ns: u64,
|
|
) -> (Vec<Mbp10RawInput>, Vec<[f32; 5]>) {
|
|
let mut out = Vec::with_capacity(seq_len);
|
|
for k in 0..seq_len {
|
|
let next_mid = prev_mid + 0.25;
|
|
let mut bid_px = [0.0f32; 10];
|
|
let mut bid_sz = [0.0f32; 10];
|
|
let mut ask_px = [0.0f32; 10];
|
|
let mut ask_sz = [0.0f32; 10];
|
|
for i in 0..10 {
|
|
bid_px[i] = next_mid - 0.125 - 0.25 * i as f32;
|
|
ask_px[i] = next_mid + 0.125 + 0.25 * i as f32;
|
|
bid_sz[i] = 10.0;
|
|
ask_sz[i] = 10.0;
|
|
}
|
|
let prev_ts = ts_ns;
|
|
ts_ns += 20_000_000;
|
|
out.push(Mbp10RawInput {
|
|
bid_px, bid_sz, ask_px, ask_sz,
|
|
prev_mid,
|
|
trade_signed_vol: 1.0,
|
|
trade_count: 1,
|
|
ts_ns,
|
|
prev_ts_ns: prev_ts,
|
|
regime: [0.0; 6],
|
|
});
|
|
prev_mid = next_mid;
|
|
let _ = k;
|
|
}
|
|
// Per-position labels: every position knows the next K snapshots all
|
|
// move up (synthetic monotone ramp), so label = 1.0 for every horizon
|
|
// at every position. Drives the trainer to learn "always predict 1".
|
|
let labels = vec![[1.0; 5]; seq_len];
|
|
(out, labels)
|
|
}
|
|
|
|
#[test]
|
|
fn stacked_trainer_constructs_cleanly() {
|
|
let dev = test_device();
|
|
let cfg = PerceptionTrainerConfig::default();
|
|
let t = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
|
drop(t);
|
|
}
|
|
|
|
#[test]
|
|
fn stacked_trainer_loss_shrinks_on_constant_signal() {
|
|
let dev = test_device();
|
|
let cfg = PerceptionTrainerConfig {
|
|
seq_len: 16, // smaller for smoke speed; Mamba2 needs >=2
|
|
mamba2_state_dim: 8,
|
|
lr_cfc: 3e-3,
|
|
lr_mamba2: 1e-3,
|
|
seed: 0x4242,
|
|
horizon_weights: [1.0; 5],
|
|
n_batch: 1,
|
|
decision_stride: 1,
|
|
};
|
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
|
|
|
// Initial loss over 8 batches.
|
|
let mut initial_total = 0.0_f32;
|
|
let mut ts = 1_000_000u64;
|
|
let mut prev_mid = 5500.0_f32;
|
|
for _ in 0..8 {
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
let l = trainer.step(&seq, labels.as_slice()).expect("step warm");
|
|
initial_total += l;
|
|
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
|
|
ts = seq.last().unwrap().ts_ns;
|
|
}
|
|
let initial_avg = initial_total / 8.0;
|
|
eprintln!("initial_avg = {initial_avg:.4}");
|
|
|
|
// Train 250 steps. Print every 50.
|
|
let mut window_loss = 0.0_f32;
|
|
let mut window_count = 0usize;
|
|
for step_idx in 0..250 {
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
let l = trainer.step(&seq, labels.as_slice()).expect("train step");
|
|
window_loss += l;
|
|
window_count += 1;
|
|
if step_idx % 50 == 49 {
|
|
eprintln!(
|
|
" step {}: window_avg_loss={:.4}",
|
|
step_idx + 1,
|
|
window_loss / window_count as f32
|
|
);
|
|
window_loss = 0.0;
|
|
window_count = 0;
|
|
}
|
|
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
|
|
ts = seq.last().unwrap().ts_ns;
|
|
}
|
|
|
|
// Final loss over 8 batches.
|
|
let mut final_total = 0.0_f32;
|
|
for _ in 0..8 {
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
let l = trainer.step(&seq, labels.as_slice()).expect("step final");
|
|
final_total += l;
|
|
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
|
|
ts = seq.last().unwrap().ts_ns;
|
|
}
|
|
let final_avg = final_total / 8.0;
|
|
eprintln!("final_avg = {final_avg:.4}");
|
|
|
|
assert!(
|
|
final_avg < 0.6 * initial_avg || final_avg < 0.5,
|
|
"PerceptionTrainer failed to overfit constant signal: \
|
|
start={initial_avg:.4}, end={final_avg:.4}"
|
|
);
|
|
}
|
|
|
|
/// Verifies the trainer still converges on a constant-direction signal
|
|
/// with `decision_stride = 4`. Stride affects loader output AND Mamba2
|
|
/// dt_s (= 4.0 in the K-loop) — this test ensures the full chain stays
|
|
/// numerically stable. Synthetic sequence here uses consecutive
|
|
/// snapshots (the loader-level stride is what skips); the trainer-level
|
|
/// dt_s change is what this smoke proves.
|
|
#[test]
|
|
fn stacked_trainer_loss_shrinks_with_stride_4() {
|
|
let dev = test_device();
|
|
let cfg = PerceptionTrainerConfig {
|
|
seq_len: 16,
|
|
mamba2_state_dim: 8,
|
|
lr_cfc: 3e-3,
|
|
lr_mamba2: 1e-3,
|
|
seed: 0xC4C4,
|
|
horizon_weights: [1.0; 5],
|
|
n_batch: 1,
|
|
decision_stride: 4,
|
|
};
|
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
|
|
|
let mut initial = 0.0_f32;
|
|
let mut ts = 1_000_000u64;
|
|
let mut prev_mid = 5500.0_f32;
|
|
for _ in 0..8 {
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
let l = trainer.step(&seq, labels.as_slice()).expect("step");
|
|
initial += l;
|
|
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
|
|
ts = seq.last().unwrap().ts_ns;
|
|
}
|
|
initial /= 8.0;
|
|
eprintln!("stride=4 trainer: initial={initial:.4}");
|
|
|
|
for _ in 0..200 {
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
trainer.step(&seq, labels.as_slice()).expect("step");
|
|
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
|
|
ts = seq.last().unwrap().ts_ns;
|
|
}
|
|
|
|
let mut final_loss = 0.0_f32;
|
|
for _ in 0..8 {
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
final_loss += trainer.step(&seq, labels.as_slice()).expect("step");
|
|
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
|
|
ts = seq.last().unwrap().ts_ns;
|
|
}
|
|
final_loss /= 8.0;
|
|
eprintln!("stride=4 trainer: final={final_loss:.4}");
|
|
assert!(
|
|
final_loss < 0.6 * initial || final_loss < 0.5,
|
|
"stride=4 trainer failed to converge: {initial:.4} → {final_loss:.4}"
|
|
);
|
|
}
|
|
|
|
/// Verifies the trainer converges at n_batch=32 — the FIRST test that
|
|
/// exercises the cross-batch reducer code path. Existing
|
|
/// `stacked_trainer_loss_shrinks_*` tests all use n_batch=1 so the new
|
|
/// per-batch scratch + reducer logic was never previously hit.
|
|
#[test]
|
|
fn stacked_trainer_loss_shrinks_at_batch_32() {
|
|
let dev = test_device();
|
|
let cfg = PerceptionTrainerConfig {
|
|
seq_len: 16,
|
|
mamba2_state_dim: 8,
|
|
lr_cfc: 3e-3,
|
|
lr_mamba2: 1e-3,
|
|
seed: 0xB32B,
|
|
horizon_weights: [1.0; 5],
|
|
n_batch: 32,
|
|
decision_stride: 1,
|
|
};
|
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
|
|
|
let mut ts_base = 1_000_000u64;
|
|
let mut prev_mid = 5500.0_f32;
|
|
|
|
let make_batch = |prev_mid: &mut f32, ts_base: &mut u64,
|
|
cfg: &PerceptionTrainerConfig|
|
|
-> (Vec<Vec<Mbp10RawInput>>, Vec<Vec<[f32; 5]>>) {
|
|
let mut seqs: Vec<Vec<Mbp10RawInput>> = Vec::with_capacity(cfg.n_batch);
|
|
let mut labels: Vec<Vec<[f32; 5]>> = Vec::with_capacity(cfg.n_batch);
|
|
for _ in 0..cfg.n_batch {
|
|
let (seq, lbl) = synthetic_seq(cfg.seq_len, *prev_mid, *ts_base);
|
|
*prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
|
|
*ts_base = seq.last().unwrap().ts_ns;
|
|
seqs.push(seq);
|
|
labels.push(lbl);
|
|
}
|
|
(seqs, labels)
|
|
};
|
|
|
|
let mut initial = 0.0_f32;
|
|
for warmup in 0..4 {
|
|
let (seqs, labels) = make_batch(&mut prev_mid, &mut ts_base, &cfg);
|
|
let seq_refs: Vec<&[Mbp10RawInput]> = seqs.iter().map(|s| s.as_slice()).collect();
|
|
let lbl_refs: Vec<&[[f32; 5]]> = labels.iter().map(|l| l.as_slice()).collect();
|
|
let l = trainer.step_batched(&seq_refs, &lbl_refs).expect("warm step");
|
|
if warmup >= 2 {
|
|
initial += l;
|
|
}
|
|
}
|
|
initial /= 2.0;
|
|
eprintln!("B=32 trainer: initial={initial:.4}");
|
|
|
|
for _ in 0..200 {
|
|
let (seqs, labels) = make_batch(&mut prev_mid, &mut ts_base, &cfg);
|
|
let seq_refs: Vec<&[Mbp10RawInput]> = seqs.iter().map(|s| s.as_slice()).collect();
|
|
let lbl_refs: Vec<&[[f32; 5]]> = labels.iter().map(|l| l.as_slice()).collect();
|
|
trainer.step_batched(&seq_refs, &lbl_refs).expect("train step");
|
|
}
|
|
|
|
let mut final_loss = 0.0_f32;
|
|
for _ in 0..4 {
|
|
let (seqs, labels) = make_batch(&mut prev_mid, &mut ts_base, &cfg);
|
|
let seq_refs: Vec<&[Mbp10RawInput]> = seqs.iter().map(|s| s.as_slice()).collect();
|
|
let lbl_refs: Vec<&[[f32; 5]]> = labels.iter().map(|l| l.as_slice()).collect();
|
|
final_loss += trainer.step_batched(&seq_refs, &lbl_refs).expect("final step");
|
|
}
|
|
final_loss /= 4.0;
|
|
eprintln!("B=32 trainer: final={final_loss:.4}");
|
|
assert!(
|
|
final_loss < 0.6 * initial || final_loss < 0.1,
|
|
"B=32 trainer failed to converge: {initial:.4} → {final_loss:.4}"
|
|
);
|
|
}
|
|
|
|
// NOTE on scratch-clears testing:
|
|
//
|
|
// A direct "scratch is zero between steps" test is structurally hard to
|
|
// write against this trainer. The lifecycle is:
|
|
// step 1: uncaptured warmup execute → scratch ends with grad
|
|
// step 2: begin_capture/end_capture → records kernels but does
|
|
// NOT execute them, scratch
|
|
// unchanged from step 1
|
|
// step 3+: graph.launch_replay → executes the captured graph
|
|
//
|
|
// A post-step snapshot read on step 2 returns step 1's residual,
|
|
// which is bit-identical between any two captures regardless of input
|
|
// data. Two replay steps would work but require either three+ trainer
|
|
// calls or capture-internal scratch readback, both of which require
|
|
// substantial test infrastructure.
|
|
//
|
|
// Instead we rely on the `stacked_trainer_loss_shrinks_at_batch_32`
|
|
// smoke as the implicit scratch-zero validation: if the K-loop scratch
|
|
// wasn't being zeroed at step start, gradients from step N would
|
|
// pollute step N+1, training would diverge after a handful of steps,
|
|
// and the convergence assertion would fail. The B=32 smoke explicitly
|
|
// runs 200 training steps + measures final loss, so the absence of
|
|
// divergence IS the scratch-zero guarantee.
|
|
|
|
|
|
/// Eval alone must work — proves the eval path is fine WITHOUT any
|
|
/// prior captured-graph training. If this passes but
|
|
/// `evaluate_works_after_captured_training_step` fails, the bug is in
|
|
/// the train→eval transition (captured graph leaving state hostile to
|
|
/// direct kernel launches).
|
|
#[test]
|
|
fn evaluate_alone_succeeds() {
|
|
let dev = test_device();
|
|
let cfg = PerceptionTrainerConfig {
|
|
seq_len: 16,
|
|
mamba2_state_dim: 8,
|
|
lr_cfc: 3e-3,
|
|
lr_mamba2: 1e-3,
|
|
seed: 0x6262,
|
|
horizon_weights: [1.0; 5],
|
|
n_batch: 1,
|
|
decision_stride: 1,
|
|
};
|
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
|
let ts = 1_000_000u64;
|
|
let prev_mid = 5500.0_f32;
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
let (loss, probs) = trainer.evaluate(&seq, labels.as_slice()).expect("eval alone");
|
|
assert!(loss.is_finite(), "eval loss must be finite, got {loss}");
|
|
assert_eq!(probs.len(), cfg.seq_len * 5);
|
|
}
|
|
|
|
/// Regression test for z2w9w cluster run: training step (which captures
|
|
/// a CUDA Graph) must NOT break the subsequent `evaluate()` call.
|
|
/// z2w9w hit CUDA_ERROR_INVALID_VALUE at "eval snap_batched fwd" the
|
|
/// very first time eval ran after training; the captured graph or some
|
|
/// of its launch state left the stream in a state hostile to the
|
|
/// non-captured eval path.
|
|
#[test]
|
|
fn evaluate_works_after_captured_training_step() {
|
|
let dev = test_device();
|
|
let cfg = PerceptionTrainerConfig {
|
|
seq_len: 16,
|
|
mamba2_state_dim: 8,
|
|
lr_cfc: 3e-3,
|
|
lr_mamba2: 1e-3,
|
|
seed: 0x5151,
|
|
horizon_weights: [1.0; 5],
|
|
n_batch: 1,
|
|
decision_stride: 1,
|
|
};
|
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
|
|
|
// Drive enough training steps to exercise warmup → capture → replay.
|
|
let mut ts = 1_000_000u64;
|
|
let mut prev_mid = 5500.0_f32;
|
|
for _ in 0..5 {
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
trainer.step(&seq, labels.as_slice()).expect("train step");
|
|
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
|
|
ts = seq.last().unwrap().ts_ns;
|
|
}
|
|
|
|
// Now call evaluate — must NOT fail with CUDA_ERROR_INVALID_VALUE.
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
let (loss, probs) = trainer.evaluate(&seq, labels.as_slice()).expect("evaluate after train");
|
|
assert!(loss.is_finite(), "eval loss must be finite, got {loss}");
|
|
assert_eq!(probs.len(), cfg.seq_len * 5, "eval probs must be [K, 5] flat");
|
|
assert!(probs.iter().all(|p| p.is_finite()), "eval probs must be finite");
|
|
}
|
|
|
|
/// 2 steps = warmup + capture (no replay yet). Does eval fail right
|
|
/// after capture but BEFORE the first graph.launch?
|
|
#[test]
|
|
fn evaluate_works_after_capture_no_replay() {
|
|
let dev = test_device();
|
|
let cfg = PerceptionTrainerConfig {
|
|
seq_len: 16,
|
|
mamba2_state_dim: 8,
|
|
lr_cfc: 3e-3,
|
|
lr_mamba2: 1e-3,
|
|
seed: 0x8181,
|
|
horizon_weights: [1.0; 5],
|
|
n_batch: 1,
|
|
decision_stride: 1,
|
|
};
|
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
|
let ts = 1_000_000u64;
|
|
let prev_mid = 5500.0_f32;
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
trainer.step(&seq, labels.as_slice()).expect("warmup step");
|
|
trainer.step(&seq, labels.as_slice()).expect("capture step");
|
|
// Eval right after capture, no replay.
|
|
let (loss, _probs) = trainer.evaluate(&seq, labels.as_slice()).expect("eval after capture");
|
|
assert!(loss.is_finite(), "eval loss must be finite, got {loss}");
|
|
}
|
|
|
|
/// ISV-driven per-horizon EMA + lambda — after a few training steps,
|
|
/// `loss_ema` should contain positive BCE values for every horizon
|
|
/// (proves the BCE kernel writes the per-horizon channel correctly),
|
|
/// and `lambda` should land inside the ASYMMETRIC clamp [1.0, 2.0]
|
|
/// (boost-only — never demote below uniform, per
|
|
/// `pearl_audit_unboundedness_for_implicit_asymmetry.md`).
|
|
#[test]
|
|
fn horizon_ema_and_lambda_track_after_training() {
|
|
let dev = test_device();
|
|
let cfg = PerceptionTrainerConfig {
|
|
seq_len: 16,
|
|
mamba2_state_dim: 8,
|
|
lr_cfc: 3e-3,
|
|
lr_mamba2: 1e-3,
|
|
seed: 0x9292,
|
|
horizon_weights: [1.0; 5],
|
|
n_batch: 1,
|
|
decision_stride: 1,
|
|
};
|
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
|
|
|
// EMA / lambda before any training step.
|
|
let ema0 = trainer.loss_ema_snapshot().expect("ema initial");
|
|
let lam0 = trainer.lambda_snapshot().expect("lambda initial");
|
|
assert!(ema0.iter().all(|v| *v == 0.0), "ema should be zero-init sentinel, got {ema0:?}");
|
|
assert!(lam0.iter().all(|v| *v == 0.0), "lambda should be zero before first kernel run, got {lam0:?}");
|
|
|
|
let mut ts = 1_000_000u64;
|
|
let mut prev_mid = 5500.0_f32;
|
|
for _ in 0..5 {
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
trainer.step(&seq, labels.as_slice()).expect("train step");
|
|
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
|
|
ts = seq.last().unwrap().ts_ns;
|
|
}
|
|
|
|
let ema = trainer.loss_ema_snapshot().expect("ema after train");
|
|
let lam = trainer.lambda_snapshot().expect("lambda after train");
|
|
eprintln!("loss_ema = {ema:?}");
|
|
eprintln!("lambda = {lam:?}");
|
|
// Each horizon must have a positive finite EMA (BCE > 0 unless model is perfect).
|
|
for (h, &v) in ema.iter().enumerate() {
|
|
assert!(v.is_finite(), "ema[{h}] not finite: {v}");
|
|
assert!(v > 0.0, "ema[{h}] should be positive, got {v}");
|
|
}
|
|
// Lambda asymmetric clamp: every entry in [1.0, 2.0] (boost-only).
|
|
for (h, &v) in lam.iter().enumerate() {
|
|
assert!(v.is_finite(), "lambda[{h}] not finite: {v}");
|
|
assert!(v >= 1.0 - 1e-6 && v <= 2.0 + 1e-6,
|
|
"lambda[{h}] = {v} outside asymmetric clamp [1.0, 2.0]");
|
|
}
|
|
// Mean of lambda must be >= 1.0 (boost-only, never demote).
|
|
// Upper-bounded by 2.0 (ceiling clamp). On synthetic-constant data
|
|
// the per-horizon EMAs converge, ratios approach 1.0, mean → 1.0.
|
|
let lam_mean: f32 = lam.iter().sum::<f32>() / lam.len() as f32;
|
|
assert!(lam_mean >= 1.0 - 1e-6 && lam_mean <= 2.0 + 1e-6,
|
|
"lambda mean {lam_mean} outside [1.0, 2.0] envelope");
|
|
}
|
|
|
|
/// Does the bug surface after just ONE step (warmup only, no capture)?
|
|
/// If yes, the warmup dispatch path itself breaks subsequent eval.
|
|
/// If no, the captured graph instantiation / launch is what breaks eval.
|
|
#[test]
|
|
fn evaluate_works_after_warmup_only() {
|
|
let dev = test_device();
|
|
let cfg = PerceptionTrainerConfig {
|
|
seq_len: 16,
|
|
mamba2_state_dim: 8,
|
|
lr_cfc: 3e-3,
|
|
lr_mamba2: 1e-3,
|
|
seed: 0x7171,
|
|
horizon_weights: [1.0; 5],
|
|
n_batch: 1,
|
|
decision_stride: 1,
|
|
};
|
|
let mut trainer = PerceptionTrainer::new(&dev, &cfg).expect("init");
|
|
let ts = 1_000_000u64;
|
|
let prev_mid = 5500.0_f32;
|
|
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
|
|
// ONLY one step — warmup, no capture yet.
|
|
trainer.step(&seq, labels.as_slice()).expect("warmup step");
|
|
// Now eval.
|
|
let (loss, _probs) = trainer.evaluate(&seq, labels.as_slice()).expect("eval after warmup");
|
|
assert!(loss.is_finite(), "eval loss must be finite, got {loss}");
|
|
}
|