diff --git a/crates/ml-alpha/cuda/mamba2_alpha_kernel.cu b/crates/ml-alpha/cuda/mamba2_alpha_kernel.cu index af48c55c8..45af591a2 100644 --- a/crates/ml-alpha/cuda/mamba2_alpha_kernel.cu +++ b/crates/ml-alpha/cuda/mamba2_alpha_kernel.cu @@ -1,34 +1,48 @@ /* ===================================================================== - * Mamba2 selective scan — ml-alpha simplified kernels (Phase 1d.1) + * Mamba2 selective scan — ml-alpha (Phase 1d.1) * * Forward + analytical backward for the SSM scan used by the supervised - * snapshot-stream pipeline. The original `mamba2_temporal_kernel.cu` in - * crates/ml/src/cuda_pipeline/ carries DQN-specific scaffolding (ISV - * adaptive signals, per-position temporal_weight) that the alpha - * supervised path doesn't use. This file is a clean rewrite: no optional - * args, no NULL-pointer branches, no DQN coupling. + * snapshot-stream pipeline. Purpose-built for the supervised path — no + * ISV adaptive signals, no per-position temporal_weight, no NULL-pointer + * branches. * - * Tensor shapes (row-major): - * a_proj : [N, K, state_d] per-step state-update gates (raw, pre-sigmoid) - * b_proj : [N, K, state_d] per-step state-input contributions - * w_c : [sh2, state_d] output mix weights (per-hidden-channel context) - * h_s2 : [N, sh2] residual input added to final hidden output - * h_enriched : [N, sh2] OUTPUT: h_s2 + (W_c @ final_state) + * Layout conventions (row-major): + * a_proj : [N, K, state_d] per-step state-update gates (pre-sigmoid) + * b_proj : [N, K, state_d] per-step state-input contributions + * w_c : [sh2, state_d] output mix weights (per-hidden-channel) + * h_s2 : [N, sh2] residual input added to scan output + * h_enriched : [N, sh2] OUTPUT of forward: h_s2 + (W_c · final_state) + * d_h_enriched : [N, sh2] upstream gradient flowing into h_enriched * - * State dim is hardcoded at ≤ 16 via the `float x[16]` register array; the - * Rust constructor rejects configs with state_d > 16. + * Backward scratch (per-channel, atomicAdd-free): + * d_a_per_channel : [N, sh2, K, state_d] each (i, j) thread writes its + * unique slot; reduced across j by + * `mamba2_alpha_reduce_d_proj` + * d_b_per_channel : [N, sh2, K, state_d] symmetric + * d_w_c_per_sample: [N, sh2, state_d] reduced across N by + * `mamba2_alpha_reduce_d_w_c` * - * Grid layout (forward): (N, ceil(sh2/32)) - * Block layout (forward): 32 threads + * Backward outputs (after reductions): + * d_a_proj : [N, K, state_d] sum over j of d_a_per_channel + * d_b_proj : [N, K, state_d] sum over j of d_b_per_channel + * d_w_c : [sh2, state_d] sum over N of d_w_c_per_sample + * d_h_s2 : [N, sh2] identity passthrough of d_h_enriched * - * One thread handles one (batch, hidden-channel) pair: scans through K - * timesteps maintaining a length-state_d state vector in registers, then - * computes the dot product with w_c[j] to produce the single output - * element at (i, j). + * Kernel constraints (must hold at every launch): + * state_d ≤ 16 (register array `float x[16]`) + * K ≤ 32 (local-memory array `float x_hist[32 * 16]` ≤ 2 KiB / thread) + * + * Rust constructor enforces both via Mamba2BlockConfig::validate. * ===================================================================== */ #include +#define MAMBA2_ALPHA_MAX_STATE_D 16 +#define MAMBA2_ALPHA_MAX_K 32 + +/* --------------------------------------------------------------------- + * Forward scan: one thread per (sample i, hidden-channel j). + * --------------------------------------------------------------------- */ extern "C" __global__ void mamba2_alpha_scan_fwd( const float* __restrict__ a_proj, const float* __restrict__ b_proj, @@ -44,9 +58,9 @@ extern "C" __global__ void mamba2_alpha_scan_fwd( int j = blockIdx.y * blockDim.x + threadIdx.x; if (i >= N || j >= sh2) return; - float x[16]; + float x[MAMBA2_ALPHA_MAX_STATE_D]; #pragma unroll - for (int s = 0; s < 16; s++) x[s] = 0.0f; + for (int s = 0; s < MAMBA2_ALPHA_MAX_STATE_D; s++) x[s] = 0.0f; for (int t = 0; t < K; t++) { long long base = ((long long)i * K + t) * state_d; @@ -64,49 +78,26 @@ extern "C" __global__ void mamba2_alpha_scan_fwd( } -/* ===================================================================== - * Backward kernel — analytical gradients through the selective scan. +/* --------------------------------------------------------------------- + * Backward scan: one thread per (sample i, hidden-channel j). * - * Inputs from the forward pass: - * a_proj, b_proj : same as forward (replayed to recover the state path) - * d_h_enriched : [N, sh2] upstream gradient flowing into h_enriched - * w_c : [sh2, state_d] + * Each thread writes to its UNIQUE slot in the per-channel scratch buffers + * `d_a_per_channel[i, j, t, s]` and `d_b_per_channel[i, j, t, s]`. No + * atomicAdd; cross-channel reduction is done by `mamba2_alpha_reduce_d_proj` + * below. The `d_w_c_per_sample[i, j, s]` slot is also (i, j)-unique; + * reduced across N by `mamba2_alpha_reduce_d_w_c`. * - * Outputs: - * d_a_proj : [N, K, state_d] ∂L/∂a_proj - * d_b_proj : [N, K, state_d] ∂L/∂b_proj - * d_w_c : [sh2, state_d] ∂L/∂w_c (atomicAdd-free via outer reduction; - * each thread writes to a private accumulator, - * block-tree-reduced by the host before commit) - * d_h_s2 : [N, sh2] ∂L/∂h_s2 (identity passthrough of d_h_enriched) - * - * Algorithm per (sample i, hidden-channel j): - * 1. Forward replay to populate x_history[t][s] for t in 0..K, s in 0..state_d - * (state_d ≤ 16, K is typically 16-64 → register/local memory fits) - * 2. d_state_K[s] = d_h_enriched[i, j] * w_c[j, s] (gradient at the final state) - * 3. Reverse scan: at each t, accumulate d_a_proj, d_b_proj, then propagate - * d_state[s] = gate * d_state_next[s] - * - * One thread per (i, j) pair (same grid layout as fwd). - * - * d_w_c gradient is more subtle: each (i, j) pair contributes - * d_w_c[j, s] += d_h_enriched[i, j] * x_K[s] - * Across i, this is a sum over the batch — which would normally be an atomicAdd. - * Per project memory (feedback_no_atomicadd), we forbid atomics. The solution - * used here: compute the per-(i, j, s) contribution and let the host orchestrate - * a block-tree reduction across i via a separate reduction kernel call. To keep - * the bwd kernel self-contained, we expose a `d_w_c_per_sample[N, sh2, state_d]` - * scratch buffer and a separate reduction kernel below. - * ===================================================================== */ + * `d_h_s2[i, j]` is the identity passthrough of `d_h_enriched[i, j]`. + * --------------------------------------------------------------------- */ extern "C" __global__ void mamba2_alpha_scan_bwd( const float* __restrict__ a_proj, const float* __restrict__ b_proj, const float* __restrict__ d_h_enriched, const float* __restrict__ w_c, - float* __restrict__ d_a_proj, - float* __restrict__ d_b_proj, - float* __restrict__ d_w_c_per_sample, - float* __restrict__ d_h_s2, + float* __restrict__ d_a_per_channel, // [N, sh2, K, state_d] + float* __restrict__ d_b_per_channel, // [N, sh2, K, state_d] + float* __restrict__ d_w_c_per_sample, // [N, sh2, state_d] + float* __restrict__ d_h_s2, // [N, sh2] int N, int K, int sh2, @@ -116,106 +107,112 @@ extern "C" __global__ void mamba2_alpha_scan_bwd( int j = blockIdx.y * blockDim.x + threadIdx.x; if (i >= N || j >= sh2) return; - /* Forward replay: build x_history (state at each t for this batch). - * x_history[t][s] = state s AFTER timestep t. Index 0 = initial state = 0. - * We need x_history[K-1] for the final state, and x_history[t-1] for the - * pre-gate value used in d_b at step t. - * - * To bound memory, we store the state-d vector at each timestep using - * a flat array x_hist[K * state_d] in local memory. With K=32, state_d=16, - * this is 32*16*4 = 2048 bytes per thread — fits in fast local mem. - */ - float x_hist[32 * 16]; // sized for max(K)=32, max(state_d)=16 - /* If K * state_d exceeds 512, kernel must be invoked with smaller K */ - - float x[16]; + /* Replay the forward state path; cache x[t][s] for the reverse scan. */ + float x_hist[MAMBA2_ALPHA_MAX_K * MAMBA2_ALPHA_MAX_STATE_D]; + float x[MAMBA2_ALPHA_MAX_STATE_D]; #pragma unroll - for (int s = 0; s < 16; s++) x[s] = 0.0f; + for (int s = 0; s < MAMBA2_ALPHA_MAX_STATE_D; s++) x[s] = 0.0f; for (int t = 0; t < K; t++) { long long base = ((long long)i * K + t) * state_d; for (int s = 0; s < state_d; s++) { float gate = 1.0f / (1.0f + expf(-a_proj[base + s])); x[s] = gate * x[s] + b_proj[base + s]; - x_hist[t * state_d + s] = x[s]; + x_hist[t * MAMBA2_ALPHA_MAX_STATE_D + s] = x[s]; } } - /* d_state at final timestep K-1: from d_h_enriched flowing back through - * the ctx dot-product. d_state[s] = d_h_enriched[i, j] * w_c[j, s]. */ - float d_state[16]; + /* Initial d_state at t = K-1: gradient flowing back through the dot + * product `ctx[i,j] = sum_s w_c[j,s] * x[s]` is d_h_ij * w_c[j,s]. */ + float d_state[MAMBA2_ALPHA_MAX_STATE_D]; float d_h_ij = d_h_enriched[(long long)i * sh2 + j]; for (int s = 0; s < state_d; s++) { d_state[s] = d_h_ij * w_c[(long long)j * state_d + s]; } - /* Reverse scan. At each t, compute: - * d_b_proj[t, s] = d_state[s] (b is added linearly) - * d_gate[t, s] = d_state[s] * x_prev[s] (gate multiplies x_prev) - * d_a_proj[t, s] = d_gate[t, s] * sigmoid'(a_proj[t, s]) - * d_state[s] = d_state[s] * gate[t, s] (propagate to t-1) + /* W_c gradient contribution from this (i, j): + * d_w_c[j, s] += d_h_ij * x_final[s] + * Each (i, j) thread writes its UNIQUE slot in d_w_c_per_sample[i, j, s]; + * reduction across i happens in mamba2_alpha_reduce_d_w_c. */ + long long w_c_slot = ((long long)i * sh2 + j) * state_d; + for (int s = 0; s < state_d; s++) { + d_w_c_per_sample[w_c_slot + s] = + d_h_ij * x_hist[(K - 1) * MAMBA2_ALPHA_MAX_STATE_D + s]; + } + + /* Reverse scan. At each step t (going K-1 → 0): + * d_b[t, s] = d_state[s] (b is added linearly) + * d_a[t, s] = d_state[s] * x_prev[s] * σ'(a[t, s]) (a feeds the gate) + * d_state[s] = d_state[s] * gate[t, s] (propagate to t-1) * - * Many (i, j) threads contribute to the same d_a_proj / d_b_proj slot. - * Per project memory (feedback_no_atomicadd), we forbid atomics. The - * solution: each (i, j) pair writes to a PRIVATE per-channel slot: - * d_a_proj[i*K*state_d*sh2 + j*K*state_d + t*state_d + s] (4-D layout) - * That's not how the kernel is currently parameterized. For Phase 1d.1 - * forward-only, we DON'T need the bwd kernel yet — this body is - * scaffolded but the host-side d_a_proj allocation pattern is deferred - * to session 3. For now: write the per-(i, j) contributions back to - * the standard 3-D layout using a single-thread-per-(i, j) gate; safe - * because each (i, j) thread updates a UNIQUE slot pattern only when - * `j == 0` is the producer — see session 3 design notes. - * - * SHIPPING SCAFFOLD: write only the j==0 thread's contributions so - * we have non-zero gradient output without atomic conflicts. Full - * d_a / d_b sum-over-j lands in session 3 with proper per-channel - * scratch + reduction kernel. + * Per-channel scratch slot for (i, j, t, s): + * d_a_per_channel[((i * sh2 + j) * K + t) * state_d + s] */ - if (j == 0) { - for (int t = K - 1; t >= 0; t--) { - long long base = ((long long)i * K + t) * state_d; - for (int s = 0; s < state_d; s++) { - float a_raw = a_proj[base + s]; - float gate = 1.0f / (1.0f + expf(-a_raw)); - float sig_deriv = gate * (1.0f - gate); + long long per_chan_base = ((long long)i * sh2 + j) * (long long)K * state_d; + for (int t = K - 1; t >= 0; t--) { + long long fwd_base = ((long long)i * K + t) * state_d; + for (int s = 0; s < state_d; s++) { + float a_raw = a_proj[fwd_base + s]; + float gate = 1.0f / (1.0f + expf(-a_raw)); + float sig_deriv = gate * (1.0f - gate); - /* x_prev[s] is the state BEFORE this step's gate-multiply. */ - float x_prev = (t == 0) ? 0.0f : x_hist[(t - 1) * state_d + s]; + float x_prev = (t == 0) ? 0.0f : x_hist[(t - 1) * MAMBA2_ALPHA_MAX_STATE_D + s]; - d_b_proj[base + s] = d_state[s]; - d_a_proj[base + s] = d_state[s] * x_prev * sig_deriv; + long long slot = per_chan_base + (long long)t * state_d + s; + d_b_per_channel[slot] = d_state[s]; + d_a_per_channel[slot] = d_state[s] * x_prev * sig_deriv; - /* Propagate to previous timestep. */ - d_state[s] = d_state[s] * gate; - } + /* Propagate to t-1. */ + d_state[s] = d_state[s] * gate; } } - /* d_w_c contribution from this (i, j): d_w_c[j, s] += d_h_ij * x_hist[K-1][s]. - * Write to per-sample scratch [N, sh2, state_d]; host reduces across N. - */ - long long w_c_off = ((long long)i * sh2 + j) * state_d; - for (int s = 0; s < state_d; s++) { - d_w_c_per_sample[w_c_off + s] = d_h_ij * x_hist[(K - 1) * state_d + s]; - } - - /* d_h_s2 is identity passthrough (h_s2 is added linearly to h_enriched). */ + /* d_h_s2 is identity passthrough (h_s2 added linearly to h_enriched). */ d_h_s2[(long long)i * sh2 + j] = d_h_ij; } -/* ===================================================================== - * Reduction kernel: sum d_w_c_per_sample[N, sh2, state_d] across N → d_w_c[sh2, state_d]. +/* --------------------------------------------------------------------- + * Reduction kernel: sum d_a_per_channel[N, sh2, K, state_d] across j → + * d_a_proj[N, K, state_d]. * - * One thread per (j, s) pair; block tree-reduces the N samples in shared memory. - * No atomicAdd — uses warp shuffle inside the block, then a single thread - * writes the final result. For batches up to N=1024, a single block handles - * the full reduction; larger N would need a two-stage launch (deferred). - * ===================================================================== */ + * Same kernel handles d_b reduction (call twice with different + * input/output pointers). Avoids code duplication. + * + * Grid: (N, K, ceil(state_d / 32)), Block: 32 + * --------------------------------------------------------------------- */ +extern "C" __global__ void mamba2_alpha_reduce_d_proj( + const float* __restrict__ d_per_channel, // [N, sh2, K, state_d] + float* __restrict__ d_proj, // [N, K, state_d] + int N, + int K, + int sh2, + int state_d +) { + int i = blockIdx.x; + int t = blockIdx.y; + int s = blockIdx.z * blockDim.x + threadIdx.x; + if (i >= N || t >= K || s >= state_d) return; + + /* Sum over j: d_proj[i, t, s] = sum_j d_per_channel[i, j, t, s] */ + float sum = 0.0f; + for (int j = 0; j < sh2; j++) { + long long slot = (((long long)i * sh2 + j) * K + t) * state_d + s; + sum += d_per_channel[slot]; + } + d_proj[((long long)i * K + t) * state_d + s] = sum; +} + + +/* --------------------------------------------------------------------- + * Reduction kernel: sum d_w_c_per_sample[N, sh2, state_d] across N → + * d_w_c[sh2, state_d]. + * + * Grid: (sh2, ceil(state_d / 32)), Block: 32 + * --------------------------------------------------------------------- */ extern "C" __global__ void mamba2_alpha_reduce_d_w_c( - const float* __restrict__ d_w_c_per_sample, - float* __restrict__ d_w_c, + const float* __restrict__ d_w_c_per_sample, // [N, sh2, state_d] + float* __restrict__ d_w_c, // [sh2, state_d] int N, int sh2, int state_d @@ -224,6 +221,7 @@ extern "C" __global__ void mamba2_alpha_reduce_d_w_c( int s = blockIdx.y * blockDim.x + threadIdx.x; if (j >= sh2 || s >= state_d) return; + /* Sum over i: d_w_c[j, s] = sum_i d_w_c_per_sample[i, j, s] */ float sum = 0.0f; for (int i = 0; i < N; i++) { sum += d_w_c_per_sample[((long long)i * sh2 + j) * state_d + s]; diff --git a/crates/ml-alpha/src/mamba2_block.rs b/crates/ml-alpha/src/mamba2_block.rs index d4f23b0c6..b88210c8b 100644 --- a/crates/ml-alpha/src/mamba2_block.rs +++ b/crates/ml-alpha/src/mamba2_block.rs @@ -42,9 +42,12 @@ use cudarc::driver::{CudaFunction, CudaModule, CudaSlice, CudaStream, LaunchConf use ml_core::cuda_autograd::gpu_tensor::GpuTensor; use ml_core::cuda_autograd::linear::OwnedGpuLinear; -/// Mamba2 hidden state dim is hardcoded at 16 floats per (sample, batch) in -/// the kernel (`float x[16]`). Configurations exceeding this are rejected. +/// Hard kernel limits. The forward and backward kernels use register/local +/// arrays sized at compile-time: `float x[16]` (state register) and +/// `float x_hist[32 * 16]` (backward replay cache, 2 KiB per thread). +/// Exceeding either silently corrupts; the Rust constructor enforces. pub const MAMBA2_KERNEL_STATE_MAX: usize = 16; +pub const MAMBA2_KERNEL_SEQ_MAX: usize = 32; /// Configuration for a Mamba2 sequence block. #[derive(Debug, Clone)] @@ -69,15 +72,43 @@ impl Mamba2BlockConfig { } if self.state_dim > MAMBA2_KERNEL_STATE_MAX { return Err(anyhow!( - "Mamba2BlockConfig: state_dim {} exceeds kernel max {} (mamba2_temporal_kernel.cu \ - hardcodes `float x[{}]` per thread)", + "Mamba2BlockConfig: state_dim {} exceeds kernel max {} (mamba2_alpha_kernel.cu \ + hardcodes `float x[{}]` register array per thread)", self.state_dim, MAMBA2_KERNEL_STATE_MAX, MAMBA2_KERNEL_STATE_MAX )); } + if self.seq_len > MAMBA2_KERNEL_SEQ_MAX { + return Err(anyhow!( + "Mamba2BlockConfig: seq_len {} exceeds kernel max {} (backward kernel's \ + x_hist replay cache is `float x_hist[{}*{}]` per thread)", + self.seq_len, MAMBA2_KERNEL_SEQ_MAX, + MAMBA2_KERNEL_SEQ_MAX, MAMBA2_KERNEL_STATE_MAX + )); + } Ok(()) } } +/// Saved activations + per-stage tensors that the backward pass needs. Returned +/// by [`Mamba2Block::forward_train`]; consumed by the backward path (lands in +/// session 3). For inference-only ([`Mamba2Block::forward`]) the cache is +/// computed but immediately dropped. +pub struct Mamba2ForwardCache { + /// Input as reshaped 2-D `[N*K, in_dim]` — needed to backprop into W_in. + pub input_2d: GpuTensor, + /// Output of W_in projection, `[N*K, hidden_dim]`. Needed to backprop into + /// W_a / W_b and to compose dx/dx_prev for the W_in gradient. + pub x: GpuTensor, + /// A projection output `[N*K, state_dim]`. Backward kernel reads this to + /// replay the forward state path. + pub a_proj: GpuTensor, + /// B projection output `[N*K, state_dim]`. Same role as a_proj. + pub b_proj: GpuTensor, + /// Output of the scan kernel, `[N, hidden_dim]`. Needed to backprop into + /// W_out (and is also the input to the W_out forward GEMM). + pub h_enriched: GpuTensor, +} + /// GPU-resident Mamba2 sequence block. /// /// Owns all parameter tensors (`W_in`, `W_a`, `W_b`, `W_c`, `W_out`) on the @@ -110,6 +141,10 @@ pub struct Mamba2Block { _module: Arc, pub kernel_fwd: CudaFunction, pub kernel_bwd: CudaFunction, + /// Reduces `d_a_per_channel` or `d_b_per_channel` (same kernel, two + /// call sites with different I/O pointers). + pub kernel_reduce_d_proj: CudaFunction, + /// Reduces `d_w_c_per_sample[N, sh2, state_d]` across N. pub kernel_reduce_d_w_c: CudaFunction, /// cuBLAS handle for the four linear projections (input / A / B / output). @@ -136,6 +171,9 @@ impl Mamba2Block { let kernel_bwd = module .load_function("mamba2_alpha_scan_bwd") .map_err(|e| anyhow!("Mamba2Block: backward kernel symbol resolve: {e}"))?; + let kernel_reduce_d_proj = module + .load_function("mamba2_alpha_reduce_d_proj") + .map_err(|e| anyhow!("Mamba2Block: d_proj reduction kernel resolve: {e}"))?; let kernel_reduce_d_w_c = module .load_function("mamba2_alpha_reduce_d_w_c") .map_err(|e| anyhow!("Mamba2Block: d_w_c reduction kernel resolve: {e}"))?; @@ -180,6 +218,7 @@ impl Mamba2Block { _module: module, kernel_fwd, kernel_bwd, + kernel_reduce_d_proj, kernel_reduce_d_w_c, cublas, }) @@ -187,29 +226,23 @@ impl Mamba2Block { // ── Forward pass ─────────────────────────────────────────────────── - /// GPU-native forward inference. `input` is a flat `[n_batch × seq_len × in_dim]` - /// `GpuTensor` (row-major); returns a flat `[n_batch]` tensor of raw logits - /// (one per sequence, predicted from the final-position state). + /// GPU-native forward pass with full activation cache returned alongside + /// the output logits. Use this when training (backward needs the cache). /// - /// Steps: + /// Steps (all on GPU; no host roundtrip): /// 1. `x = input @ Wᵢₙ.T + bᵢₙ` (cuBLAS sgemm + bias add) /// 2. `a = x @ Wₐ.T + bₐ` (cuBLAS sgemm + bias add) /// 3. `b = x @ Wᵦ.T + bᵦ` (cuBLAS sgemm + bias add) /// 4. zero-init `h_s2[B, hidden_dim]`, scratch `h_enriched[B, hidden_dim]` /// 5. launch `mamba2_alpha_scan_fwd(a, b, w_c, h_s2, h_enriched, B, K, H, S)` /// 6. `logit = h_enriched @ Wₒᵤₜ.T + bₒᵤₜ` (cuBLAS sgemm + bias add) - /// - /// All intermediate tensors live on GPU; no host roundtrip. The caller is - /// responsible for shuttling input/output between host and GPU at the - /// pipeline boundary (pinned-memory transfers handled by ml-core helpers - /// elsewhere in the stack). - pub fn forward(&self, input: &GpuTensor) -> Result { + pub fn forward_train(&self, input: &GpuTensor) -> Result<(GpuTensor, Mamba2ForwardCache)> { let c = &self.config; let n_batch = match input.shape() { [b, k, d] if *k == c.seq_len && *d == c.in_dim => *b, shape => { return Err(anyhow!( - "Mamba2Block::forward: expected input shape [B, {}, {}], got {:?}", + "Mamba2Block::forward_train: expected input shape [B, {}, {}], got {:?}", c.seq_len, c.in_dim, shape )); } @@ -224,24 +257,22 @@ impl Mamba2Block { .map_err(|e| anyhow!("reshape input → 2D: {e}"))?; // (1) input projection: x = input @ W_in.T + b_in → [N*K, hidden_dim] - let (x, _act_in) = self.w_in.inner.forward_with_slices( + let (x, _) = self.w_in.inner.forward_with_slices( &input_2d, &self.w_in.weight, &self.w_in.bias, &self.cublas, &self.stream, ).map_err(|e| anyhow!("w_in forward: {e}"))?; // (2) a_proj = x @ W_a.T + b_a → [N*K, state_dim] - let (a_proj, _act_a) = self.w_a.inner.forward_with_slices( + let (a_proj, _) = self.w_a.inner.forward_with_slices( &x, &self.w_a.weight, &self.w_a.bias, &self.cublas, &self.stream, ).map_err(|e| anyhow!("w_a forward: {e}"))?; // (3) b_proj = x @ W_b.T + b_b → [N*K, state_dim] - let (b_proj, _act_b) = self.w_b.inner.forward_with_slices( + let (b_proj, _) = self.w_b.inner.forward_with_slices( &x, &self.w_b.weight, &self.w_b.bias, &self.cublas, &self.stream, ).map_err(|e| anyhow!("w_b forward: {e}"))?; - // (4) initial hidden h_s2 and output buffer h_enriched, both [N, hidden_dim]. - // h_s2 is the residual the kernel adds to the scan context; for the - // pure-supervised pipeline we start with zero residual (no carry from - // a previous chunk). + // (4) initial hidden h_s2 (zero residual; no carry from previous chunk) + // and output buffer h_enriched, both [N, hidden_dim]. let h_s2 = GpuTensor::zeros(&[n_batch, c.hidden_dim], &self.stream) .map_err(|e| anyhow!("alloc h_s2: {e}"))?; let mut h_enriched = GpuTensor::zeros(&[n_batch, c.hidden_dim], &self.stream) @@ -250,7 +281,8 @@ impl Mamba2Block { // (5) launch the scan kernel. Grid: (N, ceil(H/32)), Block: 32. let block_threads: u32 = 32; let grid_x: u32 = n_batch as u32; - let grid_y: u32 = ((c.hidden_dim + block_threads as usize - 1) / block_threads as usize) as u32; + let grid_y: u32 = + ((c.hidden_dim + block_threads as usize - 1) / block_threads as usize) as u32; let cfg = LaunchConfig { grid_dim: (grid_x, grid_y, 1), block_dim: (block_threads, 1, 1), @@ -277,11 +309,19 @@ impl Mamba2Block { } // (6) output projection: logit = h_enriched @ W_out.T + b_out → [N, 1] - let (logit, _act_out) = self.w_out.inner.forward_with_slices( + let (logit, _) = self.w_out.inner.forward_with_slices( &h_enriched, &self.w_out.weight, &self.w_out.bias, &self.cublas, &self.stream, ).map_err(|e| anyhow!("w_out forward: {e}"))?; - Ok(logit) + let cache = Mamba2ForwardCache { input_2d, x, a_proj, b_proj, h_enriched }; + Ok((logit, cache)) + } + + /// Forward inference only — runs `forward_train` and discards the cache. + /// Use when you only need predictions (eval, calibration set scoring). + #[inline] + pub fn forward(&self, input: &GpuTensor) -> Result { + self.forward_train(input).map(|(logit, _cache)| logit) } /// Total trainable parameter count (sum of all projections + W_c). @@ -319,6 +359,38 @@ mod tests { assert!(cfg.validate().is_err(), "state_dim > 16 must be rejected"); } + #[test] + fn test_mamba2_config_rejects_seq_len_over_32() { + let cfg = Mamba2BlockConfig { + in_dim: 81, hidden_dim: 64, state_dim: 16, seq_len: 33, + }; + assert!(cfg.validate().is_err(), "seq_len > 32 must be rejected (backward x_hist limit)"); + } + + #[test] + fn test_mamba2_block_forward_train_returns_cache() { + let stream = match cuda_stream_or_skip() { + Some(s) => s, + None => return, + }; + let cfg = Mamba2BlockConfig { + in_dim: 81, hidden_dim: 32, state_dim: 8, seq_len: 16, + }; + let n_batch = 4; + let block = Mamba2Block::new(cfg.clone(), Arc::clone(&stream)).expect("init"); + let n = n_batch * cfg.seq_len * cfg.in_dim; + let h: Vec = (0..n).map(|i| (i as f32 * 1e-3).cos()).collect(); + let d = stream.clone_htod(&h).expect("htod"); + let input = GpuTensor::new(d, vec![n_batch, cfg.seq_len, cfg.in_dim]).expect("input"); + let (logit, cache) = block.forward_train(&input).expect("forward_train"); + assert_eq!(logit.shape(), &[n_batch, 1]); + assert_eq!(cache.input_2d.shape(), &[n_batch * cfg.seq_len, cfg.in_dim]); + assert_eq!(cache.x.shape(), &[n_batch * cfg.seq_len, cfg.hidden_dim]); + assert_eq!(cache.a_proj.shape(), &[n_batch * cfg.seq_len, cfg.state_dim]); + assert_eq!(cache.b_proj.shape(), &[n_batch * cfg.seq_len, cfg.state_dim]); + assert_eq!(cache.h_enriched.shape(), &[n_batch, cfg.hidden_dim]); + } + #[test] fn test_mamba2_config_rejects_zero_dims() { let cfg = Mamba2BlockConfig {