use candle_core::{Device, Tensor}; use candle_nn::{VarBuilder, VarMap}; use ml::MLError; #[test] fn test_scatter_add_gradient_flow() -> Result<(), MLError> { println!("\n=== Testing scatter_add gradient flow ===\n"); let device = Device::cuda_if_available(0)?; // Create learnable source values let varmap = VarMap::new(); let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device); let source = vb.get((2, 3), "source")?; println!("Source values: {:?}", source.to_vec2::()?); // Create indices (fixed, not learnable) let indices = Tensor::new(&[[0i64, 1, 2], [2, 1, 0]], &device)?; println!("Indices: {:?}", indices.to_vec2::()?); // Create base tensor let base = Tensor::zeros((2, 5), candle_core::DType::F32, &device)?; // Scatter add let result = base.scatter_add(&indices, &source, 1)?; println!("Result: {:?}", result.to_vec2::()?); // Compute loss let loss = result.sum_all()?; println!("Loss: {:.6}", loss.to_scalar::()?); // Backward let grads = loss.backward()?; // Check gradients let all_vars = varmap.all_vars(); if let Some(grad) = grads.get(&all_vars[0]) { println!("Gradients: {:?}", grad.to_vec2::()?); let grad_norm = grad.sqr()?.sum_all()?.to_scalar::()?.sqrt(); println!("Gradient norm: {:.6}", grad_norm); assert!( grad_norm > 1e-6, "scatter_add breaks gradient flow! Got norm={:.9}", grad_norm ); println!("\n✅ SUCCESS: scatter_add preserves gradient flow!\n"); } else { panic!("No gradients computed!"); } Ok(()) }