//! Test gradient flow when source is derived from a Var use candle_core::{Device, DType, Tensor, Var}; use ml::MLError; #[test] fn test_simple_var_gradient_flow() -> Result<(), MLError> { let device = Device::cuda_if_available(0)?; println!("\n=== Test 1: Simple Var multiplication ==="); { let x_var = Var::new(&[1.0f32, 2.0, 3.0], &device)?; let y = (x_var.as_tensor() * 2.0)?; let loss = y.sum_all()?; let grads = loss.backward()?; let has_grads = grads.get(x_var.as_tensor()).is_some(); println!("Var multiplication: has gradients: {}", has_grads); } println!("\n=== Test 2: Var used in scatter_add source ==="); { let source_var = Var::new(&[[1.0f32, 2.0, 3.0], [4.0, 5.0, 6.0]], &device)?; let base = Tensor::zeros((2, 3), DType::F32, &device)?; let indices = Tensor::new(&[[0i64, 1i64, 2i64], [0i64, 1i64, 2i64]], &device)?; let result = base.scatter_add(&indices, source_var.as_tensor(), 1)?; let loss = result.sum_all()?; let grads = loss.backward()?; let has_grads = grads.get(source_var.as_tensor()).is_some(); println!("scatter_add with Var source: has gradients: {}", has_grads); if let Some(grad) = grads.get(source_var.as_tensor()) { let grad_sum: f32 = grad.sum_all()?.to_scalar()?; println!(" Gradient sum: {}", grad_sum); } } println!("\n=== Test 3: Derived tensor from Var in scatter_add ==="); { let base_var = Var::new(&[[1.0f32, 2.0, 3.0], [4.0, 5.0, 6.0]], &device)?; let source = (base_var.as_tensor() * 2.0)?; // Derive from Var let base = Tensor::zeros((2, 3), DType::F32, &device)?; let indices = Tensor::new(&[[0i64, 1i64, 2i64], [0i64, 1i64, 2i64]], &device)?; let result = base.scatter_add(&indices, &source, 1)?; let loss = result.sum_all()?; let grads = loss.backward()?; let has_grads = grads.get(base_var.as_tensor()).is_some(); println!("scatter_add with derived source: has gradients on base_var: {}", has_grads); if let Some(grad) = grads.get(base_var.as_tensor()) { let grad_sum: f32 = grad.sum_all()?.to_scalar()?; println!(" Gradient sum: {}", grad_sum); } } Ok(()) }