P0 CRITICAL BUG IDENTIFIED: MAMBA-2 SSM matrices never train ROOT CAUSE (98% confidence): - SSM matrices initialized as raw Tensors (NOT in VarMap) - Gradients stored with generic keys (varmap_param_X) - Optimizer searches for non-existent keys (A_0, B_0, C_0) - Result: Optimizer lookups ALWAYS fail → SSM frozen at random init IMPACT: - Only projection layers learn, SSM core frozen - Model capacity severely limited (cannot learn temporal dynamics) - Validation loss ~44M vs expected ~38-40M (10-15% worse) SOLUTION (4-Phase Fix): 1. Register SSM matrices in VarMap during model creation 2. Remove special-case gradient extraction (rely on VarMap) 3. Simplify optimizer to unified VarMap loop 4. Update projection logic to query VarMap DOCUMENTS: - CRITICAL_SSM_TRAINING_BUG_ANALYSIS.md (8,500 words, complete analysis) - SSM_TRAINING_FIX_IMPLEMENTATION_GUIDE.md (2,800 words, step-by-step) - EXECUTIVE_SUMMARY_SSM_TRAINING_BUG.md (1,200 words, high-level) EVIDENCE: - Line 337-414: Tensor::from_vec() bypasses VarMap - Line 1650: Gradients stored as "varmap_param_X" - Lines 1792-1795: Optimizer searches "A_0", "B_0" (NEVER found) VERIFICATION TESTS: 1. Gradient flow: Assert SSM matrices change >1e-4 after training 2. Gradient presence: Assert gradient keys exist in HashMap 3. Spectral radius: Assert projection works with VarMap EFFORT: 6-8 hours (implementation + testing) RISK: LOW (leveraging battle-tested Candle VarMap) EXPECTED: +10-15% validation performance, smooth convergence ANALYSIS METHOD: zen thinkdeep (30 steps, 3 files, expert validation) STATUS: Ready for implementation Generated with Claude Code Co-Authored-By: Claude <noreply@anthropic.com>
15 KiB
SSM Training Fix - Implementation Guide
Quick Reference: Step-by-step instructions for implementing the SSM training fix.
See CRITICAL_SSM_TRAINING_BUG_ANALYSIS.md for complete analysis.
Phase 1: Register SSM Matrices in VarMap
File: ml/src/mamba/mod.rs
Step 1.1: Add helper function to generate SSM initialization vectors
Location: After line 310 (before impl Mamba2State)
/// Generate random initialization vector for SSM matrices
fn generate_ssm_init_vec(num_elements: usize) -> Vec<f64> {
(0..num_elements)
.map(|_| {
use rand::Rng;
let mut rng = rand::thread_rng();
rng.gen_range(-1.0..1.0) * 0.02
})
.collect()
}
Step 1.2: Create new Mamba2State constructor
Location: After line 466 (after existing zeros function)
/// Create state with SSM matrices from VarBuilder (TRAINABLE)
///
/// This constructor is used during model initialization to register
/// SSM matrices in VarMap so they can be trained via backpropagation.
///
/// # Arguments
/// * `config` - Model configuration
/// * `device` - Device for tensor allocation
/// * `vb` - VarBuilder for parameter registration
///
/// # Returns
/// State with trainable SSM matrices registered in VarMap
pub fn from_varbuilder(
config: &Mamba2Config,
device: &Device,
vb: &VarBuilder,
) -> Result<Self, MLError> {
let mut hidden_states = Vec::new();
let mut ssm_states = Vec::new();
let d_inner = config.d_model * config.expand;
for layer_idx in 0..config.num_layers {
// Hidden state (NOT trainable)
let hidden = Tensor::zeros((config.batch_size, config.d_model), DType::F64, device)
.map_err(|e| MLError::TensorCreationError {
operation: format!("hidden state creation for layer {}", layer_idx),
reason: e.to_string(),
})?;
hidden_states.push(hidden);
// === TRAINABLE SSM MATRICES ===
// A matrix: [d_state, d_state] - State transition matrix
let a_init_vec = generate_ssm_init_vec(config.d_state * config.d_state);
let a_init_tensor = Tensor::from_vec(
a_init_vec,
(config.d_state, config.d_state),
vb.device()
)?;
let A = vb.var_copy(a_init_tensor, &format!("ssm_{}.A", layer_idx))?;
// B matrix: [d_state, d_inner] - Input matrix
let b_init_vec = generate_ssm_init_vec(config.d_state * d_inner);
let b_init_tensor = Tensor::from_vec(
b_init_vec,
(config.d_state, d_inner),
vb.device()
)?;
let B = vb.var_copy(b_init_tensor, &format!("ssm_{}.B", layer_idx))?;
// C matrix: [d_inner, d_state] - Output matrix
let c_init_vec = generate_ssm_init_vec(d_inner * config.d_state);
let c_init_tensor = Tensor::from_vec(
c_init_vec,
(d_inner, config.d_state),
vb.device()
)?;
let C = vb.var_copy(c_init_tensor, &format!("ssm_{}.C", layer_idx))?;
// Delta: [d_model] - Discretization parameter
let delta = Tensor::ones((config.d_model,), DType::F64, device)
.map_err(|e| MLError::TensorCreationError {
operation: format!("delta tensor creation for layer {}", layer_idx),
reason: e.to_string(),
})?;
let delta_var = vb.var_copy(delta, &format!("ssm_{}.delta", layer_idx))?;
// SSM hidden state (NOT trainable)
let ssm_hidden = Tensor::zeros((config.batch_size, config.d_state), DType::F64, device)
.map_err(|e| MLError::TensorCreationError {
operation: format!("SSM hidden state creation for layer {}", layer_idx),
reason: e.to_string(),
})?;
ssm_states.push(SSMState {
A: A.as_tensor().clone(),
B: B.as_tensor().clone(),
C: C.as_tensor().clone(),
delta: delta_var.as_tensor().clone(),
hidden_state: ssm_hidden,
});
}
Ok(Self {
hidden_states,
ssm_states,
})
}
Step 1.3: Update Mamba2Model::new to use new constructor
Location: Replace line 667
BEFORE:
let state = Mamba2State::zeros(&config, device)?;
AFTER:
let state = Mamba2State::from_varbuilder(&config, device, &vb)?;
Phase 2: Simplify Gradient Extraction
File: ml/src/mamba/mod.rs
Location: Replace lines 1621-1673
BEFORE (Special-case VarMap loop):
// Extract gradients from all VarMap parameters
let all_vars = self.varmap.all_vars();
// ... 50 lines of special-case logic ...
AFTER (Simplified):
// Extract gradients from all VarMap parameters
// VarMap now includes SSM matrices with proper keys
for var in self.varmap.all_vars() {
if let Some(grad) = grads.get(var) {
// Get variable name from VarMap
// Note: Candle 0.4+ provides Var::name() method
let var_name = /* TODO: Extract variable name from Var */;
// Store gradient with proper key
self.gradients.insert(var_name.to_string(), grad.clone());
// Optional: Compute gradient norm for monitoring
let grad_norm: f64 = grad.flatten_all()?
.to_vec1::<f64>()?
.iter()
.map(|&g| g.powi(2))
.sum::<f64>()
.sqrt();
if grad_norm > 1e-12 {
trace!("Gradient for {}: norm={:.6e}", var_name, grad_norm);
}
}
}
// Verify we got non-zero gradients
let total_grad_norm: f64 = self.gradients.values()
.map(|g| {
g.sqr().unwrap()
.sum_all().unwrap()
.to_scalar::<f64>().unwrap()
})
.sum();
if total_grad_norm < 1e-12 {
return Err(MLError::TrainingError(
"No gradients computed - check computational graph".to_string()
));
}
trace!("Total gradient norm: {:.6e}", total_grad_norm);
Note: You'll need to determine how to extract variable names from Candle's Var type. Check Candle documentation or source code for the correct API.
Phase 3: Simplify Optimizer
File: ml/src/mamba/mod.rs
Location: Replace lines 1787-1870
BEFORE (SSM-specific update logic):
// Apply Adam updates to all SSM parameters per layer
let num_layers = self.state.ssm_states.len();
for layer_idx in 0..num_layers {
let a_grad = self.gradients.get(&format!("A_{}", layer_idx)).cloned();
// ... 80 lines of SSM-specific logic ...
}
AFTER (Unified VarMap loop):
// Unified Adam update for ALL VarMap parameters (including SSM matrices)
for var in self.varmap.all_vars() {
let var_name = /* TODO: Extract variable name */;
if let Some(grad) = self.gradients.get(&var_name) {
// Get or initialize momentum buffers
let m_key = format!("{}_momentum", var_name);
let v_key = format!("{}_variance", var_name);
let m = self.optimizer_state
.entry(m_key.clone())
.or_insert_with(|| Tensor::zeros_like(var.as_tensor()).unwrap());
let v = self.optimizer_state
.entry(v_key.clone())
.or_insert_with(|| Tensor::zeros_like(var.as_tensor()).unwrap());
// Adam update equations
let m_new = ((m * beta1)? + (grad * (1.0 - beta1))?)?;
let v_new = ((v * beta2)? + (grad.sqr()? * (1.0 - beta2))?)?;
let m_hat = (&m_new / bias_correction1)?;
let v_hat = (&v_new / bias_correction2)?;
let update = (m_hat / (v_hat.sqrt()? + eps)?)?;
let new_param = (var.as_tensor() - (&update * lr)?)?;
// Update VarMap parameter
var.set(&new_param)?;
// Store updated momentum/variance
self.optimizer_state.insert(m_key, m_new);
self.optimizer_state.insert(v_key, v_new);
trace!("Updated parameter: {}", var_name);
}
}
// Apply spectral radius projection to A matrices AFTER optimizer step
self.project_ssm_matrices()?;
Phase 4: Update Projection Logic
File: ml/src/mamba/mod.rs
Location: Replace lines 2412-2461
BEFORE (Direct tensor access):
fn project_ssm_matrices(&mut self) -> Result<()> {
for layer_idx in 0..num_layers {
let a_tensor = &self.state.ssm_states[layer_idx].A;
// ...
}
}
AFTER (VarMap query):
fn project_ssm_matrices(&mut self) -> Result<(), MLError> {
let num_layers = self.config.num_layers;
for layer_idx in 0..num_layers {
let a_name = format!("ssm_{}.A", layer_idx);
// Query VarMap for A matrix
if let Some(a_var) = self.varmap.get(&a_name) {
let a_tensor = a_var.as_tensor();
// Compute spectral radius (EXISTING LOGIC - UNCHANGED)
// TODO: Use existing eigenvalue computation
let spectral_radius = self.compute_spectral_radius(a_tensor)?;
if spectral_radius >= 1.0 {
// Project to unit ball
let projected_a = (a_tensor * (0.99 / spectral_radius))?;
// Update VarMap with projected tensor
a_var.set(&projected_a)?;
trace!(
"Layer {} A matrix projected: spectral_radius={:.6} → 0.99",
layer_idx,
spectral_radius
);
}
} else {
warn!("A matrix not found in VarMap for layer {}", layer_idx);
}
}
Ok(())
}
Verification Tests
Test 1: Gradient Flow
File: Create ml/tests/ssm_training_test.rs
#[test]
fn test_ssm_matrices_update_during_training() {
let device = Device::cuda_if_available(0).unwrap();
let config = Mamba2Config {
d_model: 64,
d_state: 16,
num_layers: 2,
batch_size: 16,
expand: 2,
..Default::default()
};
let mut model = Mamba2SSM::new(config.clone(), &device).unwrap();
// Save initial A matrix
let a_init = model.varmap.get("ssm_0.A").unwrap()
.as_tensor()
.to_vec2::<f64>()
.unwrap();
// Train for 10 epochs
for epoch in 0..10 {
let x = Tensor::randn(
0f64,
1.0,
(config.batch_size, 128, config.d_model),
&device
).unwrap();
let target = Tensor::randn(
0f64,
1.0,
(config.batch_size, 128, 1),
&device
).unwrap();
let output = model.forward(&x).unwrap();
let loss = model.compute_loss(&output, &target).unwrap();
model.backward_pass(&loss, &x, &target).unwrap();
model.optimizer_step_adam(0.001).unwrap();
println!("Epoch {}: loss={:.6}", epoch, loss.to_scalar::<f64>().unwrap());
}
// Check final A matrix
let a_final = model.varmap.get("ssm_0.A").unwrap()
.as_tensor()
.to_vec2::<f64>()
.unwrap();
// Compute mean absolute difference
let mut diff_sum = 0.0;
let mut count = 0;
for (init_row, final_row) in a_init.iter().zip(a_final.iter()) {
for (a, b) in init_row.iter().zip(final_row.iter()) {
diff_sum += (a - b).abs();
count += 1;
}
}
let mean_diff = diff_sum / count as f64;
println!("A matrix mean absolute change: {:.6e}", mean_diff);
// Assert matrix changed significantly
assert!(
mean_diff > 1e-4,
"A matrix did NOT update during training (change: {:.6e})",
mean_diff
);
println!("✅ SSM matrices are being trained!");
}
Test 2: Gradient Presence
File: Same file as Test 1
#[test]
fn test_ssm_gradients_exist() {
let device = Device::cuda_if_available(0).unwrap();
let config = Mamba2Config {
d_model: 64,
d_state: 16,
num_layers: 1,
batch_size: 8,
..Default::default()
};
let mut model = Mamba2SSM::new(config.clone(), &device).unwrap();
// Forward + backward pass
let x = Tensor::randn(0f64, 1.0, (8, 64, 64), &device).unwrap();
let output = model.forward(&x).unwrap();
let loss = output.mean_all().unwrap();
model.backward_pass(&loss, &x, &x).unwrap();
// Check SSM gradient keys exist
let expected_keys = vec!["ssm_0.A", "ssm_0.B", "ssm_0.C", "ssm_0.delta"];
for key in expected_keys {
assert!(
model.gradients.contains_key(key),
"Gradient missing for: {}",
key
);
}
println!("✅ All SSM gradients present in gradient map");
}
Troubleshooting
Issue 1: Candle API for Variable Names
Problem: How to extract variable name from Var?
Solution: Check Candle documentation or inspect VarMap internals. Possible approaches:
- Use
VarMap::all_vars()which may return(name, var)tuples - Maintain a separate HashMap mapping
Varto names - Use reflection/metadata if Candle provides it
Issue 2: Tensor Cloning Overhead
Problem: Cloning large tensors in SSM state may cause memory issues.
Solution: Store Var directly in SSMState instead of Tensor:
pub struct SSMState {
pub A: Var, // Changed from Tensor
pub B: Var,
pub C: Var,
pub delta: Var,
pub hidden_state: Tensor, // Keep as Tensor (not trainable)
}
Issue 3: Backward Compatibility
Problem: Existing checkpoints won't load with new VarMap structure.
Solution: Add checkpoint version detection:
fn load_checkpoint(&mut self, path: &str) -> Result<(), MLError> {
let checkpoint_version = detect_checkpoint_version(path)?;
match checkpoint_version {
Version::V2_0 => self.load_checkpoint_v2_0(path),
Version::V2_1 => self.load_checkpoint_v2_1(path),
_ => Err(MLError::CheckpointError("Unsupported version".into())),
}
}
Success Criteria
✅ All tests pass
- Test 1: SSM matrices change by >1e-4 after training
- Test 2: All SSM gradient keys present in gradient map
- Test 3: Spectral radius projection works
✅ Training converges smoothly
- No spikes at epoch boundaries
- Validation loss decreases monotonically
- Expected final loss: ~38-40M (10-15% improvement)
✅ Code is clean
- Removed all special-case SSM logic
- Unified parameter management through VarMap
- <200 lines of code changes
Estimated Effort
- Phase 1: 2 hours (SSM initialization refactor)
- Phase 2: 1 hour (gradient extraction simplification)
- Phase 3: 2 hours (optimizer unification)
- Phase 4: 1 hour (projection adaptation)
- Testing: 2 hours (verification tests)
Total: 8 hours
Next Steps
- Implement Phase 1 (SSM initialization)
- Verify compilation:
cargo check -p ml - Implement Phase 2 (gradient extraction)
- Implement Phase 3 (optimizer)
- Implement Phase 4 (projection)
- Write verification tests
- Run tests:
cargo test -p ml --test ssm_training_test - Train for 30 epochs, verify smooth convergence
- Update CLAUDE.md with new status
- Commit changes
Good luck! 🚀