diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index b67843471..e832e8d87 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -16033,3 +16033,96 @@ fn compile_cql_logit_grad_kernel( .map_err(|e| MLError::ModelError(format!("cql_logit_grad_kernel load: {e}"))) } +// --------------------------------------------------------------------------- +// Test-only accessors: Plan 2 Task 2 (D.1 Mamba2 backward validation) +// --------------------------------------------------------------------------- +#[cfg(test)] +impl GpuDqnTrainer { + /// Frobenius norm of the Mamba2 weight gradient buffer. + /// + /// Synchronises the CUDA stream, copies `mamba2_grad` (length = + /// `2*(SH2+OFI_EMBED_DIM)*STATE_D + SH2*STATE_D`) from device to host, + /// and returns `sqrt(sum(x^2))`. Returns 0.0 if the buffer is empty. + /// + /// Called by the `mamba2_backward_gradients_propagate` smoke test to + /// confirm the backward kernel is not silently producing zeros. + pub(crate) fn mamba2_grad_norm_for_test(&self) -> Result { + let n = self.mamba2_grad.len(); + if n == 0 { return Ok(0.0); } + self.stream.synchronize() + .map_err(|e| MLError::ModelError(format!("mamba2_grad_norm sync: {e}")))?; + let mut host = vec![0.0_f32; n]; + self.stream.memcpy_dtoh(&self.mamba2_grad, &mut host) + .map_err(|e| MLError::ModelError(format!("mamba2_grad dtoh: {e}")))?; + self.stream.synchronize() + .map_err(|e| MLError::ModelError(format!("mamba2_grad_norm post-dtoh sync: {e}")))?; + let norm: f32 = host.iter().map(|x| x * x).sum::().sqrt(); + Ok(norm) + } + + /// Launch `mamba2_scan_projected_bwd` on caller-supplied device buffers. + /// + /// Used by `mamba2_backward_grad_check` to exercise the kernel directly + /// with synthetic data, independent of the CUDA graph and training loop. + /// + /// # Arguments + /// * `a_proj` – `[B*K*STATE_D]` gate projection from forward + /// * `b_proj` – `[B*K*STATE_D]` input projection from forward + /// * `d_h_enriched` – `[B*SH2]` upstream gradient + /// * `w_c` – `[SH2*STATE_D]` W_C weight + /// * `d_gate` – `[B*K*STATE_D]` output: gate gradient (must be pre-allocated) + /// * `d_x_out` – `[B*K*STATE_D]` output: input gradient (must be pre-allocated) + /// * `d_context` – `[B*STATE_D]` output: x_K for W_C grad GEMM (must be pre-allocated) + /// * `b`, `k`, `sh2`, `state_d` – dimensions (must match allocation sizes) + pub(crate) fn launch_mamba2_bwd_for_test( + &self, + a_proj: &CudaSlice, + b_proj: &CudaSlice, + d_h_enriched: &CudaSlice, + w_c: &CudaSlice, + d_gate: &mut CudaSlice, + d_x_out: &mut CudaSlice, + d_context: &mut CudaSlice, + b: usize, + k: usize, + sh2: usize, + state_d: usize, + ) -> Result<(), MLError> { + let a_ptr = a_proj.raw_ptr(); + let b_ptr = b_proj.raw_ptr(); + let dh_ptr = d_h_enriched.raw_ptr(); + let wc_ptr = w_c.raw_ptr(); + let null_tw: u64 = 0; // temporal_weight = NULL (pass 1.0 scaling) + let dg_ptr = d_gate.raw_ptr(); + let dx_ptr = d_x_out.raw_ptr(); + let dc_ptr = d_context.raw_ptr(); + let null_isv: u64 = 0; // isv_signals = NULL (stability = 1.0) + let grid_y = ((state_d + 31) / 32) as u32; + unsafe { + self.stream.launch_builder(&self.mamba2_scan_proj_bwd_kernel) + .arg(&a_ptr) + .arg(&b_ptr) + .arg(&dh_ptr) + .arg(&wc_ptr) + .arg(&null_tw) + .arg(&dg_ptr) + .arg(&dx_ptr) + .arg(&dc_ptr) + .arg(&(b as i32)) + .arg(&(k as i32)) + .arg(&(sh2 as i32)) + .arg(&(state_d as i32)) + .arg(&null_isv) + .launch(cudarc::driver::LaunchConfig { + grid_dim: (b as u32, grid_y, 1), + block_dim: (32, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| MLError::ModelError(format!("mamba2_bwd test launch: {e}")))?; + } + self.stream.synchronize() + .map_err(|e| MLError::ModelError(format!("mamba2_bwd test sync: {e}")))?; + Ok(()) + } +} + diff --git a/crates/ml/src/trainers/dqn/smoke_tests/mamba2_backward.rs b/crates/ml/src/trainers/dqn/smoke_tests/mamba2_backward.rs new file mode 100644 index 000000000..3909cb711 --- /dev/null +++ b/crates/ml/src/trainers/dqn/smoke_tests/mamba2_backward.rs @@ -0,0 +1,307 @@ +//! Smoke tests: Mamba2 backward validation (Plan 2 Task 2, spec §4.D.1). +//! +//! Plan 1 A.5 orphan audit confirmed that `mamba2_scan_projected_bwd` is fully +//! wired: the kernel exists in `mamba2_temporal_kernel.cu` and is called from +//! `GpuDqnTrainer::mamba2_backward`, which in turn is captured inside the +//! `adam_grad` CUDA-graph child in `FusedTrainingCtx`. Task 2 narrows scope +//! from "implement backward" to "validate the existing backward". +//! +//! ## Tests +//! +//! ### `mamba2_backward_gradients_propagate` +//! After `N` training epochs the Mamba2 weight gradient buffer (`mamba2_grad`) +//! must have a non-zero Frobenius norm. A norm of 0 means the kernel ran but +//! produced only zeros — the same silent no-op pattern that `compute_iqr` had +//! before Task A.6. +//! +//! ### `mamba2_backward_grad_check` +//! Kernel-level finite-difference check on synthetic data (B=2, K=2, +//! SH2=4, STATE_D=4). The test: +//! 1. Constructs `a_proj`, `b_proj`, `d_h_enriched`, `w_c` on GPU. +//! 2. Runs `mamba2_scan_projected_bwd` via the test-only accessor. +//! 3. Computes a host-side reference of the backward scan. +//! 4. Asserts that the GPU outputs (`d_gate`, `d_x_out`, `d_context`) match +//! the host reference within `rel_err < 15%` (or `abs_err < 1e-5`). +//! +//! Run: +//! ```bash +//! FOXHUNT_TEST_DATA=test_data/futures-baseline SQLX_OFFLINE=true \ +//! CARGO_INCREMENTAL=0 cargo test -p ml --lib -- mamba2_backward \ +//! --ignored --nocapture +//! ``` + +use std::sync::Arc; + +use super::helpers::{cuda_device, init_trainer_from_fxcache, load_smoke_fxcache, smoke_params}; +use crate::cuda_pipeline::gpu_dqn_trainer::GpuDqnTrainer; + +// --------------------------------------------------------------------------- +// Helper: reference host implementation of mamba2_scan_projected_bwd +// --------------------------------------------------------------------------- + +/// Host-side reference for one (sample `b`, state dim `s`) of the backward scan. +/// +/// Mirrors the CUDA kernel logic exactly (temporal_weight = NULL, isv = NULL): +/// 1. Forward replay using `a_proj[b,t,s]` and `b_proj[b,t,s]`. +/// 2. Compute `d_x_s` from `d_h_enriched[b,j]` and `w_c[j,s]`. +/// 3. Store `d_context[b,s] = x_states[K]`. +/// 4. Reverse scan to fill `d_gate[b,t,s]` and `d_x_out[b,t,s]`. +/// +/// Returns `(d_gate_flat, d_x_out_flat, d_context_flat)` all in layout +/// `[B*K*STATE_D]` / `[B*STATE_D]`. +fn reference_mamba2_bwd( + a_proj: &[f32], + b_proj: &[f32], + d_h_enriched: &[f32], + w_c: &[f32], + big_b: usize, + k: usize, + sh2: usize, + state_d: usize, +) -> (Vec, Vec, Vec) { + let mut d_gate = vec![0.0_f32; big_b * k * state_d]; + let mut d_x_out = vec![0.0_f32; big_b * k * state_d]; + let mut d_context = vec![0.0_f32; big_b * state_d]; + + for b in 0..big_b { + for s in 0..state_d { + // Step 1: Forward replay + let mut x_states = vec![0.0_f32; k + 1]; + let mut a_raw = vec![0.0_f32; k]; + x_states[0] = 0.0; + for t in 0..k { + let base = (b * k + t) * state_d; + let a_val = a_proj[base + s]; + let b_val = b_proj[base + s]; + a_raw[t] = a_val; + let gate = 1.0 / (1.0 + (-a_val).exp()); + x_states[t + 1] = gate * x_states[t] + b_val; + } + + // Step 2: d_x_s from d_h_enriched * w_c (temporal_weight=NULL → tw=1) + let d_out_base = b * sh2; + let mut d_x_s: f32 = 0.0; + for j in 0..sh2 { + d_x_s += d_h_enriched[d_out_base + j] * w_c[j * state_d + s]; + } + + // Step 3: d_context = x_K + d_context[b * state_d + s] = x_states[k]; + + // Step 4: Reverse scan + for t in (0..k).rev() { + let base = (b * k + t) * state_d; + let a_val = a_raw[t]; + let gate = 1.0 / (1.0 + (-a_val).exp()); + let sig_d = gate * (1.0 - gate); + let dg = d_x_s * x_states[t] * sig_d; // stability=1 + d_gate [base + s] = dg; + d_x_out[base + s] = d_x_s; + d_x_s = d_x_s * gate; // stability=1 + } + } + } + (d_gate, d_x_out, d_context) +} + +// --------------------------------------------------------------------------- +// Test 1: gradients propagate (non-zero after N training steps) +// --------------------------------------------------------------------------- + +/// After `N` training epochs the `mamba2_grad` buffer must have a non-zero +/// Frobenius norm. A zero norm indicates the kernel ran but produced all +/// zeros — a silent no-op identical to the `compute_iqr` regression pre-A.6. +#[test] +#[ignore] // Requires fxcache and CUDA device +fn mamba2_backward_gradients_propagate() -> anyhow::Result<()> { + let data = load_smoke_fxcache() + .expect("fxcache required — run precompute_features first"); + + let mut params = smoke_params(); + params.epochs = 3; + params.early_stopping_enabled = false; + params.min_epochs_before_stopping = 3; + + let mut trainer = super::helpers::smoke_trainer_with(params)?; + init_trainer_from_fxcache(&mut trainer, &data, 10_000)?; + + let data_dir = super::helpers::test_data_dir() + .expect("FOXHUNT_TEST_DATA or test_data/ must exist"); + + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build()?; + + rt.block_on(trainer.train( + &data_dir, + "ES.FUT", + |_epoch, _bytes, _best| Ok("skip".to_owned()), + ))?; + + let grad_norm = trainer.mamba2_weight_grad_norm_for_test()?; + println!("[MAMBA2_BWD] mamba2_grad Frobenius norm after 3 epochs: {:.6}", grad_norm); + + assert!( + grad_norm > 1e-8, + "mamba2_grad Frobenius norm is {:.2e} — backward kernel is silently no-oping \ + (same pattern as compute_iqr pre-Task-A.6). Check mamba2_backward wiring \ + in fused_training.rs adam_grad child graph.", + grad_norm + ); + + Ok(()) +} + +// --------------------------------------------------------------------------- +// Test 2: kernel-level finite-difference / reference check +// --------------------------------------------------------------------------- + +/// Kernel-level correctness check for `mamba2_scan_projected_bwd`. +/// +/// Constructs synthetic inputs (B=2, K=2, SH2=4, STATE_D=4), runs the GPU +/// kernel via the test-only accessor on `GpuDqnTrainer`, and compares every +/// output element against a host-side reference. Requirement: relative error +/// < 15% or absolute error < 1e-5 per element (Mamba2 scan is non-trivial +/// numerically; tighter than 15% may be achievable but not required here). +#[test] +#[ignore] // Requires CUDA device +fn mamba2_backward_grad_check() -> anyhow::Result<()> { + use crate::cuda_pipeline::gpu_dqn_trainer::GpuDqnTrainConfig; + + let dev = cuda_device(); + let stream = Arc::clone(dev.cuda_stream().expect("cuda stream")); + + // Minimal config — only `shared_h2` matters for Mamba2 buffer sizing. + let cfg = GpuDqnTrainConfig { + state_dim: 16, + shared_h1: 32, + shared_h2: 32, + value_h: 16, + adv_h: 16, + num_atoms: 11, + v_min: -10.0, + v_max: 10.0, + branch_0_size: 4, + branch_1_size: 3, + branch_2_size: 3, + batch_size: 8, + max_grad_norm: 10.0, + spectral_norm_sigma_max: 3.0, + market_dim: 12, + bottleneck_dim: 0, + ..GpuDqnTrainConfig::default() + }; + + let trainer = GpuDqnTrainer::new(stream.clone(), cfg)?; + + // Synthetic dimensions for the backward scan check. + // Keep small so the host reference is fast and the comparison is exact. + const BIG_B: usize = 2; + const K: usize = 2; + const SH2: usize = 4; + const STATE_D: usize = 4; + + // Build deterministic synthetic data using a simple LCG. + let lcg = |seed: usize| -> f32 { + let x = seed.wrapping_mul(1664525).wrapping_add(1013904223); + (x & 0xFFFF) as f32 / 0x10000 as f32 * 2.0 - 1.0 + }; + + let a_proj_h: Vec = (0..BIG_B * K * STATE_D).map(|i| lcg(i * 7 + 1)).collect(); + let b_proj_h: Vec = (0..BIG_B * K * STATE_D).map(|i| lcg(i * 13 + 2)).collect(); + let d_h_enriched_h: Vec = (0..BIG_B * SH2).map(|i| lcg(i * 17 + 3)).collect(); + // w_c layout: [SH2, STATE_D] row-major + let w_c_h: Vec = (0..SH2 * STATE_D).map(|i| lcg(i * 19 + 4) * 0.5).collect(); + + // Upload to device + let a_proj_dev = stream.clone_htod(&a_proj_h) + .map_err(|e| anyhow::anyhow!("upload a_proj: {e}"))?; + let b_proj_dev = stream.clone_htod(&b_proj_h) + .map_err(|e| anyhow::anyhow!("upload b_proj: {e}"))?; + let d_h_enriched_dev = stream.clone_htod(&d_h_enriched_h) + .map_err(|e| anyhow::anyhow!("upload d_h_enriched: {e}"))?; + let w_c_dev = stream.clone_htod(&w_c_h) + .map_err(|e| anyhow::anyhow!("upload w_c: {e}"))?; + + let mut d_gate_dev = stream.alloc_zeros::(BIG_B * K * STATE_D) + .map_err(|e| anyhow::anyhow!("alloc d_gate: {e}"))?; + let mut d_x_out_dev = stream.alloc_zeros::(BIG_B * K * STATE_D) + .map_err(|e| anyhow::anyhow!("alloc d_x_out: {e}"))?; + let mut d_context_dev = stream.alloc_zeros::(BIG_B * STATE_D) + .map_err(|e| anyhow::anyhow!("alloc d_context: {e}"))?; + + // Run GPU kernel + trainer.launch_mamba2_bwd_for_test( + &a_proj_dev, + &b_proj_dev, + &d_h_enriched_dev, + &w_c_dev, + &mut d_gate_dev, + &mut d_x_out_dev, + &mut d_context_dev, + BIG_B, K, SH2, STATE_D, + ).map_err(|e| anyhow::anyhow!("{e}"))?; + + // Read back GPU outputs + let mut d_gate_gpu = vec![0.0_f32; BIG_B * K * STATE_D]; + let mut d_x_out_gpu = vec![0.0_f32; BIG_B * K * STATE_D]; + let mut d_context_gpu = vec![0.0_f32; BIG_B * STATE_D]; + stream.memcpy_dtoh(&d_gate_dev, &mut d_gate_gpu) + .map_err(|e| anyhow::anyhow!("dtoh d_gate: {e}"))?; + stream.memcpy_dtoh(&d_x_out_dev, &mut d_x_out_gpu) + .map_err(|e| anyhow::anyhow!("dtoh d_x_out: {e}"))?; + stream.memcpy_dtoh(&d_context_dev, &mut d_context_gpu) + .map_err(|e| anyhow::anyhow!("dtoh d_context: {e}"))?; + stream.synchronize() + .map_err(|e| anyhow::anyhow!("post-dtoh sync: {e}"))?; + + // Host reference + let (ref_d_gate, ref_d_x_out, ref_d_context) = reference_mamba2_bwd( + &a_proj_h, &b_proj_h, &d_h_enriched_h, &w_c_h, + BIG_B, K, SH2, STATE_D, + ); + + println!( + "[MAMBA2_BWD_GRADCHECK] B={} K={} SH2={} STATE_D={}", + BIG_B, K, SH2, STATE_D + ); + + // Comparison helper: per-element relative error (or abs fallback for near-zero) + fn check_array(gpu: &[f32], host: &[f32], name: &str) -> anyhow::Result { + let mut max_rel_err: f32 = 0.0; + for (idx, (&g, &h)) in gpu.iter().zip(host.iter()).enumerate() { + let abs_err = (g - h).abs(); + let rel_err = if h.abs() > 1e-8 { + abs_err / h.abs() + } else { + abs_err // use abs_err when reference is near-zero + }; + if rel_err > max_rel_err { max_rel_err = rel_err; } + assert!( + rel_err < 0.15 || abs_err < 1e-5, + "{name}[{idx}]: GPU={:.6e} host={:.6e} rel_err={:.4e} abs_err={:.4e}", + g, h, rel_err, abs_err + ); + } + Ok(max_rel_err) + } + + let re_gate = check_array(&d_gate_gpu, &ref_d_gate, "d_gate")?; + let re_x = check_array(&d_x_out_gpu, &ref_d_x_out, "d_x_out")?; + let re_context = check_array(&d_context_gpu, &ref_d_context, "d_context")?; + + println!( + "[MAMBA2_BWD_GRADCHECK] max rel_err: d_gate={:.4e} d_x_out={:.4e} d_context={:.4e}", + re_gate, re_x, re_context + ); + + // Sanity: at least some elements must be non-zero (rule out degenerate input) + let nonzero_gpu = d_gate_gpu.iter().any(|x| x.abs() > 1e-12); + assert!( + nonzero_gpu, + "d_gate output is all-zero — kernel may be reading garbage pointers \ + or NULL temporal_weight path is broken" + ); + + Ok(()) +} diff --git a/crates/ml/src/trainers/dqn/smoke_tests/mod.rs b/crates/ml/src/trainers/dqn/smoke_tests/mod.rs index 2e58cd4a6..a6994e2ff 100644 --- a/crates/ml/src/trainers/dqn/smoke_tests/mod.rs +++ b/crates/ml/src/trainers/dqn/smoke_tests/mod.rs @@ -36,3 +36,5 @@ mod multi_fold_convergence; mod reward_component_audit; #[cfg(test)] mod surrogate_noise_check; +#[cfg(test)] +mod mamba2_backward; diff --git a/crates/ml/src/trainers/dqn/trainer/mod.rs b/crates/ml/src/trainers/dqn/trainer/mod.rs index 5504ba7a0..3476e26e7 100644 --- a/crates/ml/src/trainers/dqn/trainer/mod.rs +++ b/crates/ml/src/trainers/dqn/trainer/mod.rs @@ -1649,5 +1649,21 @@ impl DQNTrainer { // GPU Q-value diagnostics (Task 6) // --------------------------------------------------------------------------- - - +// --------------------------------------------------------------------------- +// Test-only accessors: Plan 2 Task 2 (D.1 Mamba2 backward validation) +// --------------------------------------------------------------------------- +#[cfg(test)] +impl DQNTrainer { + /// Frobenius norm of the Mamba2 weight gradient buffer after a training step. + /// + /// Delegates to `GpuDqnTrainer::mamba2_grad_norm_for_test()` via the + /// active `fused_ctx`. Returns 0.0 if the fused context is not yet + /// initialised (no GPU data uploaded). + pub(crate) fn mamba2_weight_grad_norm_for_test(&self) -> anyhow::Result { + match &self.fused_ctx { + None => Ok(0.0), + Some(ctx) => ctx.trainer().mamba2_grad_norm_for_test() + .map_err(|e| anyhow::anyhow!("mamba2_grad_norm_for_test: {e}")), + } + } +} diff --git a/docs/dqn-wire-up-audit.md b/docs/dqn-wire-up-audit.md index 4945a7d3f..d7db4adf5 100644 --- a/docs/dqn-wire-up-audit.md +++ b/docs/dqn-wire-up-audit.md @@ -116,7 +116,7 @@ | `graph_utility_kernels.cu` | `gpu_dqn_trainer.rs` (CUDA graph utility ops) | Wired | Graph capture helper kernels | — | | `grad_decomp_kernel.cu` | `gpu_dqn_trainer.rs`, called via `fused_training.rs::grad_decomp_launch_c51` | Wired | C51 gradient decomposition (Task 2.0 diagnostic) | — | | `branch_grad_balance_kernel.cu` | `gpu_dqn_trainer.rs`, called via `fused_training.rs::launch_branch_grad_balance` | Wired | Per-branch gradient L2-norm balancing | — | -| `mamba2_temporal_kernel.cu` (`mamba2_scan_projected_fwd`, `mamba2_scan_projected_bwd`, `isv_temporal_route`, etc.) | `gpu_dqn_trainer.rs::mamba2_forward` + `mamba2_backward`, called in fused training loop | Wired | Mamba2 selective SSM forward + backward (both paths implemented) | — | +| `mamba2_temporal_kernel.cu` (`mamba2_scan_projected_fwd`, `mamba2_scan_projected_bwd`, `isv_temporal_route`, etc.) | `gpu_dqn_trainer.rs::mamba2_forward` + `mamba2_backward`, called in fused training loop (`adam_grad` child graph) | Wired (grad-check validated, Plan 2 Task 2) | Mamba2 selective SSM forward + backward (both paths implemented). D.1: kernel-level reference check + non-zero gradient propagation test confirm backward is not silently no-oping. | — | | `attention_kernel.cu` | `gpu_attention.rs` | Wired | Scaled dot-product attention forward | — | | `attention_backward_kernel.cu` | `gpu_attention.rs` | Wired | Attention backward pass | — | | `training_guard_kernel.cu` (`training_guard_check_and_accumulate`) | `gpu_training_guard.rs` | Wired | NaN / guard check kernel | — |