//! Comprehensive tests for LogCosh loss implementation //! //! NOTE: Disabled until loss API is fully implemented #![cfg(feature = "disabled_tests")] use rtx_tensor::{Device, Tensor}; use rtx_transformers::{ TransformerError, losses::{LogCoshLoss, Reduction}, }; #[cfg(test)] mod logcosh_loss_tests { use super::*; fn create_test_device() -> Device { Device::cpu() } fn create_test_tensors() -> (Tensor, Tensor) { let predictions = Tensor::from_data(vec![1.0, 2.0, 3.0, 4.0], vec![4], &create_test_device()).unwrap(); let targets = Tensor::from_data(vec![1.5, 2.5, 2.5, 3.5], vec![4], &create_test_device()).unwrap(); (predictions, targets) } #[test] fn test_logcosh_loss_creation() { let loss = LogCoshLoss::new(); assert_eq!(loss.reduction(), Reduction::Mean); assert_eq!(loss.beta(), 1.0); } #[test] fn test_logcosh_with_custom_beta() { let loss = LogCoshLoss::with_beta(2.0); assert_eq!(loss.beta(), 2.0); assert_eq!(loss.reduction(), Reduction::Mean); } #[test] fn test_logcosh_with_reduction() { let loss = LogCoshLoss::with_reduction(Reduction::Sum); assert_eq!(loss.reduction(), Reduction::Sum); assert_eq!(loss.beta(), 1.0); } #[test] fn test_logcosh_forward_mean_reduction() { let loss = LogCoshLoss::new(); let (predictions, targets) = create_test_tensors(); let result = loss.forward(&predictions, &targets, None); assert!(result.is_ok()); let loss_value = result.unwrap(); assert_eq!(loss_value.shape().dims(), &[1]); // Mean reduces to scalar } #[test] fn test_logcosh_forward_sum_reduction() { let loss = LogCoshLoss::with_reduction(Reduction::Sum); let (predictions, targets) = create_test_tensors(); let result = loss.forward(&predictions, &targets, None); assert!(result.is_ok()); let loss_value = result.unwrap(); assert_eq!(loss_value.shape().dims(), &[1]); // Sum reduces to scalar } #[test] fn test_logcosh_forward_no_reduction() { let loss = LogCoshLoss::with_reduction(Reduction::None); let (predictions, targets) = create_test_tensors(); let result = loss.forward(&predictions, &targets, None); assert!(result.is_ok()); let loss_value = result.unwrap(); assert_eq!(loss_value.shape().dims(), &[4]); // No reduction keeps element-wise } #[test] fn test_logcosh_with_sample_weights() { let loss = LogCoshLoss::new(); let (predictions, targets) = create_test_tensors(); let weights = Tensor::from_data(vec![1.0, 2.0, 0.5, 1.5], vec![4], &create_test_device()).unwrap(); let result = loss.forward(&predictions, &targets, Some(&weights)); assert!(result.is_ok()); } #[test] fn test_logcosh_backward() { let loss = LogCoshLoss::new(); let (predictions, targets) = create_test_tensors(); // First forward pass let loss_value = loss.forward(&predictions, &targets, None).unwrap(); // Backward pass let gradients = loss.backward(&predictions, &targets, None); assert!(gradients.is_ok()); let grad_tensor = gradients.unwrap(); assert_eq!(grad_tensor.shape(), predictions.shape()); } #[test] fn test_logcosh_backward_with_weights() { let loss = LogCoshLoss::new(); let (predictions, targets) = create_test_tensors(); let weights = Tensor::from_data(vec![2.0, 1.0, 0.5, 1.5], vec![4], &create_test_device()).unwrap(); let gradients = loss.backward(&predictions, &targets, Some(&weights)); assert!(gradients.is_ok()); let grad_tensor = gradients.unwrap(); assert_eq!(grad_tensor.shape(), predictions.shape()); } #[test] fn test_logcosh_numerical_stability() { let loss = LogCoshLoss::new(); // Test with large differences (should handle overflow) let predictions = Tensor::from_data(vec![1000.0, -1000.0, 0.0], vec![3], &create_test_device()).unwrap(); let targets = Tensor::from_data(vec![0.0, 0.0, 1000.0], vec![3], &create_test_device()).unwrap(); let result = loss.forward(&predictions, &targets, None); assert!(result.is_ok()); let loss_value = result.unwrap(); let data = loss_value.to_cpu().unwrap(); // Check that values are finite (no inf or nan) for &val in &data { assert!(val.is_finite()); } } #[test] fn test_logcosh_with_different_beta_values() { let (predictions, targets) = create_test_tensors(); // Test different beta values for beta in [0.5, 1.0, 2.0, 5.0] { let loss = LogCoshLoss::with_beta(beta); let result = loss.forward(&predictions, &targets, None); assert!(result.is_ok()); } } #[test] fn test_logcosh_zero_error() { let loss = LogCoshLoss::new(); // When predictions == targets, loss should be 0 let tensor = Tensor::from_data(vec![1.0, 2.0, 3.0], vec![3], &create_test_device()).unwrap(); let result = loss.forward(&tensor, &tensor, None).unwrap(); let data = result.to_cpu().unwrap(); // LogCosh(0) = 0 assert!((data[0]).abs() < 1e-6); } #[test] fn test_logcosh_batch_processing() { let loss = LogCoshLoss::new(); // 2D batch: [batch_size, features] let predictions = Tensor::from_data( vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], &create_test_device(), ) .unwrap(); let targets = Tensor::from_data( vec![1.5, 2.5, 3.5, 4.5, 5.5, 6.5], vec![2, 3], &create_test_device(), ) .unwrap(); let result = loss.forward(&predictions, &targets, None); assert!(result.is_ok()); } #[test] fn test_logcosh_gradient_computation() { let loss = LogCoshLoss::new(); // For LogCosh, gradient = beta * tanh(beta * error) let predictions = Tensor::from_data(vec![2.0], vec![1], &create_test_device()).unwrap(); let targets = Tensor::from_data(vec![1.0], vec![1], &create_test_device()).unwrap(); let gradients = loss.backward(&predictions, &targets, None).unwrap(); let grad_data = gradients.to_cpu().unwrap(); // Error = 2.0 - 1.0 = 1.0 // Gradient ≈ tanh(1.0) ≈ 0.7616 assert!((grad_data[0] - 0.7616).abs() < 0.01); } #[test] fn test_logcosh_mismatched_shapes() { let loss = LogCoshLoss::new(); let predictions = Tensor::from_data(vec![1.0, 2.0, 3.0], vec![3], &create_test_device()).unwrap(); let targets = Tensor::from_data(vec![1.0, 2.0], vec![2], &create_test_device()).unwrap(); let result = loss.forward(&predictions, &targets, None); assert!(result.is_err()); } #[test] fn test_logcosh_weight_shape_mismatch() { let loss = LogCoshLoss::new(); let (predictions, targets) = create_test_tensors(); let weights = Tensor::from_data( vec![1.0, 2.0], // Wrong size vec![2], &create_test_device(), ) .unwrap(); let result = loss.forward(&predictions, &targets, Some(&weights)); assert!(result.is_err()); } }