# MAMBA-2 Matrix Dimension Bug - Visual Analysis ## Error Visualization ``` ┌──────────────────────────────────────────────────────────────┐ │ MAMBA-2 MATRIX DIMENSION BUG │ └──────────────────────────────────────────────────────────────┘ ERROR: shape mismatch in matmul, lhs: [32, 60, 512], rhs: [512, 16] ┌─────────────────────────────────────────────────────────────┐ │ Current (BROKEN) │ ├─────────────────────────────────────────────────────────────┤ │ │ │ Input (x): B Matrix: │ │ ┌─────────────┐ ┌──────┐ │ │ │ 32 │ │ 16 │ │ │ │ 60 │ @ │ 512 │ ❌ INCOMPATIBLE │ │ │ 512 │ └──────┘ │ │ └─────────────┘ │ │ [batch, seq, 2*d] [n, 2*d] │ │ │ │ Problem: Last dim of x (512) ≠ First dim of B (16) │ │ │ └─────────────────────────────────────────────────────────────┘ ┌─────────────────────────────────────────────────────────────┐ │ Fix 1: TRANSPOSE B │ ├─────────────────────────────────────────────────────────────┤ │ │ │ Input (x): B Matrix (transposed): │ │ ┌─────────────┐ ┌──────┐ │ │ │ 32 │ │ 512 │ │ │ │ 60 │ @ │ 16 │ ✅ COMPATIBLE │ │ │ 512 │ └──────┘ │ │ └─────────────┘ │ │ [batch, seq, 2*d] [2*d, n] │ │ │ │ Result: [32, 60, 16] (batch, seq, state_size) │ │ │ │ CODE: let b_proj = x.matmul(&self.b.t()?)?; │ │ │ └─────────────────────────────────────────────────────────────┘ ┌─────────────────────────────────────────────────────────────┐ │ Fix 2: RESHAPE + TRANSPOSE (if needed) │ ├─────────────────────────────────────────────────────────────┤ │ │ │ Step 1: Flatten batch+seq dimensions │ │ ┌─────────────┐ ┌────────┐ │ │ │ 32 │ │ 1920 │ │ │ │ 60 │ → │ 512 │ │ │ │ 512 │ └────────┘ │ │ └─────────────┘ │ │ [32, 60, 512] [1920, 512] │ │ │ │ Step 2: Matmul with transposed B │ │ ┌────────┐ ┌──────┐ ┌────────┐ │ │ │ 1920 │ │ 512 │ │ 1920 │ │ │ │ 512 │ @ │ 16 │ → │ 16 │ │ │ └────────┘ └──────┘ └────────┘ │ │ [1920, 512] [512, 16] [1920, 16] │ │ │ │ Step 3: Reshape back to 3D │ │ ┌────────┐ ┌─────────────┐ │ │ │ 1920 │ │ 32 │ │ │ │ 16 │ → │ 60 │ │ │ └────────┘ │ 16 │ │ │ └─────────────┘ │ │ [1920, 16] [32, 60, 16] │ │ │ │ CODE: │ │ let (b, s, f) = x.dims3()?; │ │ let x_flat = x.reshape(&[b * s, f])?; │ │ let proj_flat = x_flat.matmul(&self.b.t()?)?; │ │ let proj = proj_flat.reshape(&[b, s, self.n])?; │ │ │ └─────────────────────────────────────────────────────────────┘ ``` ## Dimension Legend ``` batch_size (b) = 32 # Number of samples in batch seq_len (s) = 60 # Sequence length (timesteps) d_model = 256 # Model hidden dimension 2*d_model = 512 # Expanded dimension (2x for selective scan) n (state_size) = 16 # SSM state dimension ``` ## Debug Output Analysis ``` [AGENT 172 DEBUG] Layer 0 B matrix initialized: shape=[16, 512], expected=[16, 512] ^^^^^^^^^^ [n, 2*d_model] This is WRONG shape for matmul! Should be [2*d_model, n] = [512, 16] Expected shapes: Initialization: [n, 2*d_model] = [16, 512] ← Current (wrong for matmul) For matmul: [2*d_model, n] = [512, 16] ← Needs transpose ``` ## Root Cause The B matrix is initialized in the correct shape `[n, 2*d_model] = [16, 512]` for storage, but needs to be transposed to `[2*d_model, n] = [512, 16]` for matmul operations. **Solution**: Add `.t()?` (transpose) to B matrix during matmul ## Files to Fix 1. **Primary**: `/home/jgrusewski/Work/foxhunt/ml/src/mamba/mod.rs` - Method: `Mamba2SSM::forward_with_gradients()` - Line: Search for `x.matmul(&self.b)` - Change: `x.matmul(&self.b.t()?)?` ## Testing Strategy ```bash # 1. Quick compile check cargo check -p ml # 2. Unit test (if exists) cargo test -p ml mamba::tests::test_forward_pass --release # 3. Integration test (1 epoch, ~30 seconds) cargo run -p ml --example train_mamba2_dbn --release -- --epochs 1 # 4. Verify output shapes # Look for these in logs: # ✓ B projection shape: [32, 60, 16] (correct) # ✓ Training loss: 0.XXX (not NaN) # ✓ Gradients flowing (not zero) ``` ## Success Criteria ✅ Compilation succeeds ✅ Shape mismatch error gone ✅ B projection output shape = `[batch, seq, n]` = `[32, 60, 16]` ✅ Training loss is finite (not NaN or Inf) ✅ Gradients are non-zero ✅ First epoch completes successfully ## Expected Timeline - Fix implementation: 2-5 minutes - Compilation: 30-45 seconds - Testing (1 epoch): 30-60 seconds - Validation: 5-10 minutes - **Total**: 10-20 minutes ## Next Steps After Fix 1. ✅ Verify 1 epoch training completes 2. ✅ Check gradient flow (add debug logging) 3. ✅ Run 5 epoch test to verify stability 4. ✅ Add shape validation tests 5. 🚀 Start full 200 epoch training run --- **Created**: Agent 248 (2025-10-15) **Status**: Ready for Agent 249 to implement fix