From deed15b34ead01be1db0454f9c570b728dcf83f5 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 16 May 2026 22:59:36 +0200 Subject: [PATCH] feat(ml-alpha): cfc_step_backward emits grad_x for upstream chain MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds grad_x[k] = sum_i d_pre[i] * W_in[i,k] computed by thread 0 of the cfc_step_backward kernel (after the existing __syncthreads in the shared-mem sd_pre relay). Required by the stacked Mamba2 -> CfC design: Mamba2.backward_from_h_enriched needs grad on h_enriched, which is the CfC's "x" input in the stacked topology. For the existing CfC-only PerceptionTrainer (x = snap_features, no upstream learnable layer), grad_x is computed but discarded into a preallocated buffer. backward_finite_diff tests still pass (4/4) — the new arg is the 14th positional kernel arg; existing callers updated. perception_ overfit smoke still passes (loss 0.5669 -> 0.0665 in 200 steps). Co-Authored-By: Claude Opus 4.7 --- crates/ml-alpha/cuda/cfc_step.cu | 19 ++++++++++++++++++- crates/ml-alpha/src/cfc/step.rs | 13 ++++++++++--- crates/ml-alpha/src/trainer/perception.rs | 10 +++++++++- crates/ml-alpha/tests/backward_finite_diff.rs | 4 ++-- 4 files changed, 39 insertions(+), 7 deletions(-) diff --git a/crates/ml-alpha/cuda/cfc_step.cu b/crates/ml-alpha/cuda/cfc_step.cu index a9b40a39c..40057818e 100644 --- a/crates/ml-alpha/cuda/cfc_step.cu +++ b/crates/ml-alpha/cuda/cfc_step.cu @@ -71,7 +71,8 @@ extern "C" __global__ void cfc_step_backward( float* __restrict__ grad_w_in, // [n_hid, n_in] float* __restrict__ grad_w_rec, // [n_hid, n_hid] float* __restrict__ grad_b, // [n_hid] - float* __restrict__ grad_h_old // [n_hid] + float* __restrict__ grad_h_old, // [n_hid] + float* __restrict__ grad_x // [n_in] — set to nullptr to skip ) { extern __shared__ float smem[]; float* sd_pre = smem; // [n_hid] @@ -107,4 +108,20 @@ extern "C" __global__ void cfc_step_backward( gh += sd_pre[j] * w_rec[j * n_hid + i]; } grad_h_old[i] = gh; + + // grad_x[k] = sum_i d_pre[i] * W_in[i, k]. Each thread i contributes to + // n_in different output positions — different threads write to different + // grad_x slots if we partition by i, but multiple i's contribute to each k. + // Avoid atomicAdd: rely on thread 0 to do the n_in × n_hid sum + // sequentially (n_in is small — typically 32 or 128, this is one-time per + // backward call). Discriminated by nullptr → skip path for back-compat. + if (grad_x != nullptr && i == 0) { + for (int k = 0; k < n_in; ++k) { + float gx = 0.0f; + for (int j = 0; j < n_hid; ++j) { + gx += sd_pre[j] * w_in[j * n_in + k]; + } + grad_x[k] = gx; + } + } } diff --git a/crates/ml-alpha/src/cfc/step.rs b/crates/ml-alpha/src/cfc/step.rs index 9cb05b9c2..beb605b48 100644 --- a/crates/ml-alpha/src/cfc/step.rs +++ b/crates/ml-alpha/src/cfc/step.rs @@ -22,7 +22,11 @@ pub struct CfcWeights { /// Backward through ONE cfc_step (truncated BPTT, K=1). /// -/// Returns (grad_w_in, grad_w_rec, grad_b, grad_h_old). +/// Returns (grad_w_in, grad_w_rec, grad_b, grad_h_old, grad_x). +/// grad_x is the gradient w.r.t. the input x (needed when chaining +/// into a prior encoder like Mamba2). For the original CfC-only +/// trainer where x comes from snap_feature_assemble (no upstream +/// learnable layer), grad_x can be discarded. pub fn cfc_step_backward_gpu( dev: &MlDevice, w: &CfcWeights, @@ -30,7 +34,7 @@ pub fn cfc_step_backward_gpu( h_old: &[f32], grad_h_new: &[f32], dt_s: f32, -) -> Result<(Vec, Vec, Vec, Vec)> { +) -> Result<(Vec, Vec, Vec, Vec, Vec)> { assert_eq!(x.len(), w.n_in); assert_eq!(h_old.len(), w.n_hid); assert_eq!(grad_h_new.len(), w.n_hid); @@ -51,6 +55,7 @@ pub fn cfc_step_backward_gpu( let mut grad_w_rec_d = stream.alloc_zeros::(w.n_hid * w.n_hid).context("grad_w_rec alloc")?; let mut grad_b_d = stream.alloc_zeros::(w.n_hid).context("grad_b alloc")?; let mut grad_h_old_d = stream.alloc_zeros::(w.n_hid).context("grad_h_old alloc")?; + let mut grad_x_d = stream.alloc_zeros::(w.n_in).context("grad_x alloc")?; let n_in_i = w.n_in as i32; let n_hid_i = w.n_hid as i32; @@ -69,7 +74,8 @@ pub fn cfc_step_backward_gpu( .arg(&x_d).arg(&h_old_d).arg(&grad_h_new_d) .arg(&dt_s).arg(&n_in_i).arg(&n_hid_i) .arg(&mut grad_w_in_d).arg(&mut grad_w_rec_d) - .arg(&mut grad_b_d).arg(&mut grad_h_old_d); + .arg(&mut grad_b_d).arg(&mut grad_h_old_d) + .arg(&mut grad_x_d); unsafe { launch.launch(cfg).context("cfc_bwd launch")?; } Ok(( @@ -77,6 +83,7 @@ pub fn cfc_step_backward_gpu( download(stream, &grad_w_rec_d)?, download(stream, &grad_b_d)?, download(stream, &grad_h_old_d)?, + download(stream, &grad_x_d)?, )) } diff --git a/crates/ml-alpha/src/trainer/perception.rs b/crates/ml-alpha/src/trainer/perception.rs index dc07108ca..73d3f1f2a 100644 --- a/crates/ml-alpha/src/trainer/perception.rs +++ b/crates/ml-alpha/src/trainer/perception.rs @@ -100,6 +100,7 @@ pub struct PerceptionTrainer { grad_heads_b_d: CudaSlice, grad_h_old_d: CudaSlice, grad_h_new_d: CudaSlice, + grad_x_d: CudaSlice, // discarded — CfC input is snap_features (no upstream gradient consumer) // Optimizers opt_w_in: AdamW, @@ -179,6 +180,7 @@ impl PerceptionTrainer { grad_heads_b_d: stream.alloc_zeros::(N_HORIZONS)?, grad_h_old_d: stream.alloc_zeros::(cfg.n_hid)?, grad_h_new_d: stream.alloc_zeros::(cfg.n_hid)?, + grad_x_d: stream.alloc_zeros::(cfg.n_in)?, stg_bid_px: unsafe { MappedF32Buffer::new(10) }.map_err(|e| anyhow::anyhow!("stg bid_px: {e}"))?, stg_bid_sz: unsafe { MappedF32Buffer::new(10) }.map_err(|e| anyhow::anyhow!("stg bid_sz: {e}"))?, stg_ask_px: unsafe { MappedF32Buffer::new(10) }.map_err(|e| anyhow::anyhow!("stg ask_px: {e}"))?, @@ -292,13 +294,19 @@ impl PerceptionTrainer { shared_mem_bytes: shared_mem, }; { + // PerceptionTrainer drives CfC from raw snap_features (no + // upstream learnable layer), so grad_x is discarded. Pass + // grad_h_old's buffer as a dummy non-null sink — the kernel + // only writes grad_x when the pointer is non-null AND only + // from thread 0, so any allocated buffer works. let mut launch = self.stream.launch_builder(&self.step_bwd_fn); launch .arg(&self.w_in_d).arg(&self.w_rec_d).arg(&self.b_d).arg(&self.tau_d) .arg(&self.snap_feat_d).arg(&self.h_old_d).arg(&self.grad_h_new_d) .arg(&dt_s).arg(&n_in_i).arg(&n_hid_i) .arg(&mut self.grad_w_in_d).arg(&mut self.grad_w_rec_d) - .arg(&mut self.grad_b_d).arg(&mut self.grad_h_old_d); + .arg(&mut self.grad_b_d).arg(&mut self.grad_h_old_d) + .arg(&mut self.grad_x_d); unsafe { launch.launch(cfg_bwd).context("cfc bwd")?; } } self.stream.synchronize().context("bwd sync")?; diff --git a/crates/ml-alpha/tests/backward_finite_diff.rs b/crates/ml-alpha/tests/backward_finite_diff.rs index 730685b59..9fe445b1a 100644 --- a/crates/ml-alpha/tests/backward_finite_diff.rs +++ b/crates/ml-alpha/tests/backward_finite_diff.rs @@ -108,7 +108,7 @@ fn cfc_backward_grad_b_matches_finite_diff() { // Surrogate loss L = sum_i h_new[i]. let grad_h_new = vec![1.0f32; n_hid]; - let (_, _, grad_b_analytic, _) = + let (_, _, grad_b_analytic, _, _) = cfc_step_backward_gpu(&dev, &w, &x, &h_old, &grad_h_new, dt_s).unwrap(); let eps = 1e-3_f32; @@ -137,7 +137,7 @@ fn cfc_backward_grad_h_old_matches_finite_diff() { let h_old: Vec = (0..n_hid).map(|_| r.gen_range(-1.0..1.0)).collect(); let dt_s = 0.02_f32; let grad_h_new = vec![1.0f32; n_hid]; - let (_, _, _, grad_h_old_analytic) = + let (_, _, _, grad_h_old_analytic, _) = cfc_step_backward_gpu(&dev, &w, &x, &h_old, &grad_h_new, dt_s).unwrap(); let eps = 1e-3_f32;