//! Comprehensive tests for cross-entropy loss implementation //! //! NOTE: Disabled until loss API is fully implemented #![cfg(feature = "disabled_tests")] use rtx_tensor::{Device, Tensor}; use rtx_transformers::{ TransformerError, losses::{CrossEntropyConfig, CrossEntropyLoss, Reduction}, }; #[cfg(test)] mod cross_entropy_tests { use super::*; fn create_test_device() -> Device { Device::cpu() } fn create_test_logits_and_targets() -> (Tensor, Tensor) { // Logits: [batch_size=2, seq_len=3, vocab_size=4] let logits = Tensor::from_data( vec![ // Batch 0 1.0, 2.0, 3.0, 4.0, // Token 0 2.0, 3.0, 1.0, 0.0, // Token 1 0.0, 1.0, 2.0, 3.0, // Token 2 // Batch 1 4.0, 3.0, 2.0, 1.0, // Token 0 0.0, 1.0, 2.0, 3.0, // Token 1 3.0, 2.0, 1.0, 0.0, // Token 2 ], vec![2, 3, 4], &create_test_device(), ) .unwrap(); // Targets: [batch_size=2, seq_len=3] let targets = Tensor::from_data( vec![ 3.0, 1.0, 2.0, // Batch 0: target indices 0.0, 2.0, 1.0, // Batch 1: target indices ], vec![2, 3], &create_test_device(), ) .unwrap(); (logits, targets) } #[test] fn test_cross_entropy_creation() { let loss = CrossEntropyLoss::new(); assert_eq!(loss.config().reduction, Reduction::Mean); assert!(loss.config().ignore_index.is_none()); assert_eq!(loss.config().label_smoothing, 0.0); } #[test] fn test_cross_entropy_with_config() { let config = CrossEntropyConfig { reduction: Reduction::Sum, ignore_index: Some(-100), label_smoothing: 0.1, }; let loss = CrossEntropyLoss::with_config(config.clone()); assert_eq!(loss.config().reduction, Reduction::Sum); assert_eq!(loss.config().ignore_index, Some(-100)); assert_eq!(loss.config().label_smoothing, 0.1); } #[test] fn test_sparse_cross_entropy_forward() { let loss = CrossEntropyLoss::new(); let (logits, targets) = create_test_logits_and_targets(); let result = loss.forward(&logits, &targets); assert!(result.is_ok()); let loss_value = result.unwrap(); assert_eq!(loss_value.shape().dims(), &[1]); // Mean reduces to scalar } #[test] fn test_cross_entropy_with_ignore_index() { let config = CrossEntropyConfig { reduction: Reduction::Mean, ignore_index: Some(-100), label_smoothing: 0.0, }; let loss = CrossEntropyLoss::with_config(config); let logits = Tensor::from_data( vec![1.0, 2.0, 3.0, 2.0, 1.0, 3.0], vec![2, 3], &create_test_device(), ) .unwrap(); let targets = Tensor::from_data( vec![1.0, -100.0], // Second target should be ignored vec![2], &create_test_device(), ) .unwrap(); let result = loss.forward(&logits, &targets); assert!(result.is_ok()); } #[test] fn test_cross_entropy_with_label_smoothing() { let config = CrossEntropyConfig { reduction: Reduction::Mean, ignore_index: None, label_smoothing: 0.1, }; let loss = CrossEntropyLoss::with_config(config); let (logits, targets) = create_test_logits_and_targets(); let result = loss.forward(&logits, &targets); assert!(result.is_ok()); } #[test] fn test_cross_entropy_sum_reduction() { let config = CrossEntropyConfig { reduction: Reduction::Sum, ignore_index: None, label_smoothing: 0.0, }; let loss = CrossEntropyLoss::with_config(config); let (logits, targets) = create_test_logits_and_targets(); let result = loss.forward(&logits, &targets); assert!(result.is_ok()); let loss_value = result.unwrap(); assert_eq!(loss_value.shape().dims(), &[1]); // Sum reduces to scalar } #[test] fn test_cross_entropy_no_reduction() { let config = CrossEntropyConfig { reduction: Reduction::None, ignore_index: None, label_smoothing: 0.0, }; let loss = CrossEntropyLoss::with_config(config); let (logits, targets) = create_test_logits_and_targets(); let result = loss.forward(&logits, &targets); assert!(result.is_ok()); let loss_value = result.unwrap(); assert_eq!(loss_value.shape().dims(), &[2, 3]); // No reduction keeps shape } #[test] fn test_cross_entropy_backward() { let loss = CrossEntropyLoss::new(); let (logits, targets) = create_test_logits_and_targets(); // Forward pass let loss_value = loss.forward(&logits, &targets).unwrap(); // Backward pass let gradients = loss.backward(&logits, &targets); assert!(gradients.is_ok()); let grad_tensor = gradients.unwrap(); assert_eq!(grad_tensor.shape(), logits.shape()); } #[test] fn test_log_softmax() { let loss = CrossEntropyLoss::new(); let logits = Tensor::from_data(vec![1.0, 2.0, 3.0, 4.0], vec![1, 4], &create_test_device()).unwrap(); let log_probs = loss.log_softmax(&logits, -1); assert!(log_probs.is_ok()); let result = log_probs.unwrap(); assert_eq!(result.shape(), logits.shape()); // Log probabilities should be negative or zero let data = result.to_cpu().unwrap(); for val in data { assert!(val <= 0.0); } } #[test] fn test_numerical_stability() { let loss = CrossEntropyLoss::new(); // Test with large logits that could cause overflow let logits = Tensor::from_data( vec![1000.0, -1000.0, 0.0], vec![1, 3], &create_test_device(), ) .unwrap(); let targets = Tensor::from_data(vec![0.0], vec![1], &create_test_device()).unwrap(); let result = loss.forward(&logits, &targets); assert!(result.is_ok()); let loss_value = result.unwrap(); let data = loss_value.to_cpu().unwrap(); // Check that value is finite assert!(data[0].is_finite()); } #[test] fn test_mismatched_shapes() { let loss = CrossEntropyLoss::new(); let logits = Tensor::from_data(vec![1.0, 2.0, 3.0], vec![1, 3], &create_test_device()).unwrap(); let targets = Tensor::from_data( vec![0.0, 1.0], // Wrong shape vec![2], &create_test_device(), ) .unwrap(); let result = loss.forward(&logits, &targets); assert!(result.is_err()); } #[test] fn test_dense_cross_entropy() { let loss = CrossEntropyLoss::new(); // Dense targets (one-hot encoded): [batch_size, seq_len, vocab_size] let logits = Tensor::from_data( vec![1.0, 2.0, 3.0, 2.0, 3.0, 1.0], vec![2, 3], &create_test_device(), ) .unwrap(); let dense_targets = Tensor::from_data( vec![ 0.0, 1.0, 0.0, // One-hot for class 1 1.0, 0.0, 0.0, // One-hot for class 0 ], vec![2, 3], &create_test_device(), ) .unwrap(); let result = loss.dense_cross_entropy(&logits, &dense_targets); assert!(result.is_ok()); } }