Files
rustytorch/crates/training/rtx-transformers/tests/cross_entropy_test.rs
T
2026-03-04 00:08:42 +00:00

265 lines
7.6 KiB
Rust

//! 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());
}
}