//! # MAMBA-2 Gradient Extraction Test (TDD) //! //! **Test-Driven Development**: This test verifies that gradients are properly extracted //! from VarMap parameters after backward() pass. //! //! **Root Cause**: backward_pass() uses zeros_like() placeholder gradients instead of //! extracting real gradients from VarMap. //! //! **Expected Behavior**: //! 1. Call forward() to compute loss //! 2. Call backward() to compute gradients //! 3. Extract gradients from VarMap parameters (input_proj, output_proj, layer_norms) //! 4. Verify gradients are non-zero and valid (not NaN/Inf) #![allow(unused_crate_dependencies)] use candle_core::{DType, Device, Tensor}; use ml::mamba::Mamba2SSM; use ml::MLError; #[test] fn test_mamba2_gradient_extraction_from_varmap() -> Result<(), MLError> { println!("\n=== MAMBA-2 Gradient Extraction Test ==="); let device = Device::cuda_if_available(0)?; println!("Device: {:?}", device); // Create small MAMBA-2 model let mut model = Mamba2SSM::default_hft(&device)?; println!( "Model created: {} parameters", model.metadata.num_parameters ); // Create dummy input and target let batch_size = model.config.batch_size; let seq_len = model.config.seq_len; let d_model = model.config.d_model; let input_data = vec![0.1f64; batch_size * seq_len * d_model]; let input = Tensor::from_vec(input_data, (batch_size, seq_len, d_model), &device)?; let target_data = vec![0.5f64; batch_size * seq_len]; let target = Tensor::from_vec(target_data, (batch_size, seq_len, 1), &device)?; println!("Input shape: {:?}", input.dims()); println!("Target shape: {:?}", target.dims()); // Forward pass let output = model.forward(&input)?; println!("Output shape: {:?}", output.dims()); // Compute loss (MSE) let diff = output.broadcast_sub(&target)?; let loss = diff.sqr()?.mean_all()?; let loss_value = loss.to_scalar::()?; println!("Loss: {:.6}", loss_value); // Backward pass - THIS SHOULD COMPUTE REAL GRADIENTS let grads = loss.backward()?; // Extract gradients from GradStore println!("\n=== Extracting Gradients from GradStore ==="); let varmap = &model.varmap; let all_vars = varmap.all_vars(); println!("Total VarMap variables: {}", all_vars.len()); let mut vars_with_gradients = 0; let mut total_grad_norm = 0.0f64; for (idx, var) in all_vars.iter().enumerate() { if let Some(grad) = grads.get(var) { // Compute gradient norm let grad_vec = grad.flatten_all()?.to_vec1::()?; let grad_norm: f64 = grad_vec.iter().map(|&g| g.powi(2)).sum::().sqrt(); println!(" Var {}: grad_norm={:.6}", idx, grad_norm); // Verify gradient is valid assert!(!grad_norm.is_nan(), "Gradient {} is NaN", idx); assert!(!grad_norm.is_infinite(), "Gradient {} is Inf", idx); if grad_norm > 1e-9 { vars_with_gradients += 1; total_grad_norm += grad_norm; } } else { println!(" Var {}: NO GRADIENT", idx); } } println!("\n=== Gradient Summary ==="); println!( "Variables with gradients: {}/{}", vars_with_gradients, all_vars.len() ); println!("Total gradient norm: {:.6}", total_grad_norm); // CRITICAL ASSERTION: At least some parameters should have non-zero gradients assert!( vars_with_gradients > 0, "FAIL: No variables have gradients! backward() did not compute gradients." ); assert!( total_grad_norm > 1e-6, "FAIL: Total gradient norm is too small ({:.6}). Gradients may be zeros.", total_grad_norm ); println!("\nāœ… TEST PASSED: Gradients extracted from VarMap"); Ok(()) } #[test] fn test_mamba2_backward_pass_extracts_real_gradients() -> Result<(), MLError> { println!("\n=== MAMBA-2 backward_pass() Real Gradient Test ==="); let device = Device::cuda_if_available(0)?; let mut model = Mamba2SSM::default_hft(&device)?; // Create input/target let batch_size = model.config.batch_size; let seq_len = model.config.seq_len; let d_model = model.config.d_model; let input = Tensor::ones((batch_size, seq_len, d_model), DType::F64, &device)?; let target = Tensor::ones((batch_size, seq_len, 1), DType::F64, &device)?; // Forward + loss let output = model.forward(&input)?; let diff = output.broadcast_sub(&target)?; let loss = diff.sqr()?.mean_all()?; println!("Loss: {:.6}", loss.to_scalar::()?); // Call backward_pass (current implementation uses zeros_like placeholders) model.backward_pass(&loss, &input, &target)?; // Check model.gradients HashMap println!("\n=== Model Gradients HashMap ==="); println!("Total entries: {}", model.gradients.len()); for (key, grad) in model.gradients.iter() { let grad_vec = grad.flatten_all()?.to_vec1::()?; let grad_norm: f64 = grad_vec.iter().map(|&g| g.powi(2)).sum::().sqrt(); println!(" {}: grad_norm={:.6}", key, grad_norm); // CURRENT BUG: All gradients are zeros (zeros_like) // AFTER FIX: Gradients should be non-zero if grad_norm > 1e-9 { println!(" āœ… Non-zero gradient found"); } else { println!(" āŒ ZERO gradient (zeros_like placeholder)"); } } // This test will FAIL until we fix backward_pass() // After fix, gradients should be non-zero let total_grad_norm: f64 = model .gradients .values() .map(|grad| { let grad_vec = grad.flatten_all().unwrap().to_vec1::().unwrap(); grad_vec.iter().map(|&g| g.powi(2)).sum::().sqrt() }) .sum(); println!( "\nTotal gradient norm in model.gradients: {:.6}", total_grad_norm ); // EXPECTED TO FAIL with current zeros_like implementation assert!( total_grad_norm > 1e-6, "FAIL: backward_pass() produced zero gradients. Need to extract from VarMap." ); println!("\nāœ… TEST PASSED: backward_pass() extracts real gradients"); Ok(()) }