265 lines
7.6 KiB
Rust
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());
|
|
}
|
|
}
|