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

250 lines
7.4 KiB
Rust

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