//! Test for Rainbow DQN Loss Computation Shape Mismatch //! //! This test reproduces the shape mismatch bug in compute_rainbow_loss: //! `shape mismatch in mul, lhs: [32, 1], rhs: [32]` #![allow(unused_crate_dependencies)] use candle_core::{Device, Tensor}; /// Test: Reproduce shape mismatch in target Q-value computation /// /// This test simulates the exact tensor operations in `compute_rainbow_loss` /// that cause the shape mismatch between target_q [32, 1] and gamma_tensor [32] #[test] fn test_target_q_value_shape_mismatch() { let device = Device::Cpu; let batch_size = 32; // Simulate next_q_values shape [32, 4, 1] (batch, actions, 1) from get_q_values // This happens because to_scalar uses sum_keepdim which keeps the last dim as 1 let next_q_values = Tensor::randn(0.0_f32, 1.0, (batch_size, 4, 1), &device).unwrap(); // Simulate next_actions [32] from argmax let next_actions = Tensor::from_vec( (0..batch_size).map(|_| 0_u32).collect::>(), batch_size, &device, ) .unwrap(); // Gather operation: extract Q-values for selected actions // gather produces [32, 1, 1] // Need to unsqueeze twice: once for gather dimension, once for the trailing 1 let gathered = next_q_values .gather(&next_actions.unsqueeze(1).unwrap().unsqueeze(2).unwrap(), 1) .unwrap(); // squeeze(1) produces [32, 1] - THIS IS THE BUG let target_q = gathered.squeeze(1).unwrap(); // Verify shape is [32, 1] (this is the problematic shape) assert_eq!(target_q.shape().dims(), &[32, 1]); // Create gamma_tensor [32] let gamma_tensor = Tensor::from_vec(vec![0.99_f32; batch_size], batch_size, &device).unwrap(); // Verify shape is [32] assert_eq!(gamma_tensor.shape().dims(), &[32]); // This multiplication SHOULD FAIL with shape mismatch [32, 1] vs [32] let result = target_q.mul(&gamma_tensor); // The test should fail here showing the shape mismatch match result { Ok(_) => panic!("Expected shape mismatch error but operation succeeded!"), Err(e) => { let error_msg = format!("{:?}", e); assert!( error_msg.contains("shape mismatch") || error_msg.contains("incompatible"), "Expected shape mismatch error, got: {}", error_msg ); }, } } /// Test: Correct shape handling with squeeze /// /// This test shows the FIX - we need to squeeze both dimensions after gather #[test] fn test_target_q_value_shape_fix() { let device = Device::Cpu; let batch_size = 32; // Simulate next_q_values shape [32, 4, 1] (batch, actions, 1) from get_q_values let next_q_values = Tensor::randn(0.0_f32, 1.0, (batch_size, 4, 1), &device).unwrap(); // Simulate next_actions [32] from argmax let next_actions = Tensor::from_vec( (0..batch_size).map(|_| 0_u32).collect::>(), batch_size, &device, ) .unwrap(); // Gather operation: extract Q-values for selected actions [32, 1, 1] let gathered = next_q_values .gather(&next_actions.unsqueeze(1).unwrap().unsqueeze(2).unwrap(), 1) .unwrap(); // FIX: squeeze BOTH dimensions to get [32] let target_q = gathered.squeeze(1).unwrap().squeeze(1).unwrap(); // Verify shape is [32] (fixed!) assert_eq!(target_q.shape().dims(), &[32]); // Create gamma_tensor [32] let gamma_tensor = Tensor::from_vec(vec![0.99_f32; batch_size], batch_size, &device).unwrap(); // Verify shape is [32] assert_eq!(gamma_tensor.shape().dims(), &[32]); // This multiplication should now work! let result = target_q.mul(&gamma_tensor); assert!( result.is_ok(), "Multiplication should succeed with matching shapes" ); let product = result.unwrap(); assert_eq!(product.shape().dims(), &[32]); } /// Test: Current action Q-values shape handling /// /// Verifies the same issue exists for current_action_q computation #[test] fn test_current_action_q_shape_mismatch() { let device = Device::Cpu; let batch_size = 32; // Simulate current_q_values shape [32, 4, 1] from get_q_values let current_q_values = Tensor::randn(0.0_f32, 1.0, (batch_size, 4, 1), &device).unwrap(); // Simulate actions [32] let actions = Tensor::from_vec( (0..batch_size).map(|i| (i % 4) as u32).collect::>(), batch_size, &device, ) .unwrap(); // Gather operation [32, 1, 1] let gathered = current_q_values .gather(&actions.unsqueeze(1).unwrap().unsqueeze(2).unwrap(), 1) .unwrap(); // squeeze(1) produces [32, 1] - same bug let current_action_q = gathered.squeeze(1).unwrap(); // Verify shape is [32, 1] assert_eq!(current_action_q.shape().dims(), &[32, 1]); // Create target_values [32] let target_values = Tensor::randn(0.0_f32, 1.0, batch_size, &device).unwrap(); // Verify shape is [32] assert_eq!(target_values.shape().dims(), &[32]); // Subtraction should fail with shape mismatch let result = current_action_q.sub(&target_values); match result { Ok(_) => panic!("Expected shape mismatch error but operation succeeded!"), Err(e) => { let error_msg = format!("{:?}", e); assert!( error_msg.contains("shape mismatch") || error_msg.contains("incompatible"), "Expected shape mismatch error, got: {}", error_msg ); }, } } /// Test: Current action Q-values shape fix /// /// Verifies the fix works for current_action_q computation #[test] fn test_current_action_q_shape_fix() { let device = Device::Cpu; let batch_size = 32; // Simulate current_q_values shape [32, 4, 1] from get_q_values let current_q_values = Tensor::randn(0.0_f32, 1.0, (batch_size, 4, 1), &device).unwrap(); // Simulate actions [32] let actions = Tensor::from_vec( (0..batch_size).map(|i| (i % 4) as u32).collect::>(), batch_size, &device, ) .unwrap(); // Gather operation [32, 1, 1] let gathered = current_q_values .gather(&actions.unsqueeze(1).unwrap().unsqueeze(2).unwrap(), 1) .unwrap(); // FIX: squeeze BOTH dimensions to get [32] let current_action_q = gathered.squeeze(1).unwrap().squeeze(1).unwrap(); // Verify shape is [32] assert_eq!(current_action_q.shape().dims(), &[32]); // Create target_values [32] let target_values = Tensor::randn(0.0_f32, 1.0, batch_size, &device).unwrap(); // Verify shape is [32] assert_eq!(target_values.shape().dims(), &[32]); // Subtraction should now work! let result = current_action_q.sub(&target_values); assert!( result.is_ok(), "Subtraction should succeed with matching shapes" ); let diff = result.unwrap(); assert_eq!(diff.shape().dims(), &[32]); }