use crate::error::RealizarError;
use crate::layers::*;
use crate::tensor::Tensor;
#[test]
fn test_softmax_empty_data_error() {
let input = Tensor::from_vec(vec![1], vec![1.0]).expect("test");
let result = softmax(&input);
assert!(result.is_ok(), "Single element softmax should succeed");
}
#[test]
fn test_softmax_single_element_tensor() {
let input = Tensor::from_vec(vec![1], vec![42.0]).expect("test");
let output = softmax(&input).expect("test");
assert_eq!(output.shape(), &[1]);
assert!(
(output.data()[0] - 1.0).abs() < 1e-6,
"Softmax of single element should be 1.0"
);
}
#[test]
fn test_softmax_very_large_negative_values() {
let input = Tensor::from_vec(vec![3], vec![-1000.0, -1001.0, -1002.0]).expect("test");
let output = softmax(&input).expect("test");
for &val in output.data() {
assert!(
val.is_finite(),
"Softmax should handle large negative values"
);
}
let sum: f32 = output.data().iter().sum();
assert!((sum - 1.0).abs() < 1e-5);
}
#[test]
fn test_softmax_3d_tensor() {
let input = Tensor::from_vec(
vec![2, 2, 3],
vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, ],
)
.expect("test");
let output = softmax(&input).expect("test");
assert_eq!(output.shape(), &[2, 2, 3]);
for row in 0..4 {
let row_sum: f32 = (0..3).map(|i| output.data()[row * 3 + i]).sum();
assert!(
(row_sum - 1.0).abs() < 1e-5,
"Row {} sum should be 1.0, got {}",
row,
row_sum
);
}
}
#[test]
fn test_softmax_identical_values() {
let input = Tensor::from_vec(vec![4], vec![5.0, 5.0, 5.0, 5.0]).expect("test");
let output = softmax(&input).expect("test");
for &val in output.data() {
assert!(
(val - 0.25).abs() < 1e-6,
"Uniform input should give uniform output"
);
}
}
#[test]
fn test_gelu_single_element() {
let input = Tensor::from_vec(vec![1], vec![0.5]).expect("test");
let output = gelu(&input).expect("test");
assert_eq!(output.shape(), &[1]);
assert!(output.data()[0] > 0.3 && output.data()[0] < 0.4);
}
#[test]
fn test_gelu_large_positive() {
let input = Tensor::from_vec(vec![1], vec![10.0]).expect("test");
let output = gelu(&input).expect("test");
assert!((output.data()[0] - 10.0).abs() < 0.01);
}
#[test]
fn test_gelu_large_negative() {
let input = Tensor::from_vec(vec![1], vec![-10.0]).expect("test");
let output = gelu(&input).expect("test");
assert!(output.data()[0].abs() < 0.01);
}
#[test]
fn test_gelu_3d_tensor() {
let input = Tensor::from_vec(
vec![2, 2, 3],
vec![
-1.0, 0.0, 1.0, -2.0, 0.5, 2.0, -3.0, 1.5, 3.0, -0.5, 0.25, 0.75,
],
)
.expect("test");
let output = gelu(&input).expect("test");
assert_eq!(output.shape(), &[2, 2, 3]);
assert!((output.data()[1] - 0.0).abs() < 1e-6);
assert!(output.data()[2] > 0.0); assert!(output.data()[5] > 0.0); }
#[test]
fn test_gelu_symmetry_property() {
let pos_input = Tensor::from_vec(vec![1], vec![0.5]).expect("test");
let neg_input = Tensor::from_vec(vec![1], vec![-0.5]).expect("test");
let pos_output = gelu(&pos_input).expect("test");
let neg_output = gelu(&neg_input).expect("test");
assert!(
(pos_output.data()[0] + neg_output.data()[0]).abs() > 0.1,
"GELU should not be antisymmetric"
);
}
#[test]
fn test_quantized_linear_zero_in_features_error() {
let result = QuantizedLinear::new(0, 256, vec![], vec![0.0; 256]);
assert!(result.is_err(), "Should error on zero in_features");
if let Err(RealizarError::InvalidShape { reason }) = result {
assert!(reason.contains("in_features") || reason.contains("> 0"));
}
}
#[test]
fn test_quantized_linear_zero_out_features_error() {
let result = QuantizedLinear::new(256, 0, vec![], vec![]);
assert!(result.is_err(), "Should error on zero out_features");
if let Err(RealizarError::InvalidShape { reason }) = result {
assert!(reason.contains("out_features") || reason.contains("> 0"));
}
}
#[test]
fn test_quantized_linear_bias_length_mismatch_error() {
let weight_bytes = vec![0u8; 288];
let bias = vec![0.0f32; 3];
let result = QuantizedLinear::new(256, 2, weight_bytes, bias);
assert!(result.is_err(), "Should error on bias length mismatch");
if let Err(RealizarError::InvalidShape { reason }) = result {
assert!(
reason.contains("Bias") || reason.contains("doesn't match"),
"Error should mention bias mismatch: {}",
reason
);
}
}
#[test]
fn test_quantized_linear_weight_bytes_mismatch_error() {
let weight_bytes = vec![0u8; 100]; let bias = vec![0.0f32; 4];
let result = QuantizedLinear::new(256, 4, weight_bytes, bias);
assert!(result.is_err(), "Should error on weight bytes mismatch");
if let Err(RealizarError::InvalidShape { reason }) = result {
assert!(
reason.contains("Weight bytes") || reason.contains("doesn't match"),
"Error should mention weight bytes mismatch: {}",
reason
);
}
}
#[test]
fn test_quantized_linear_getters() {
let weight_bytes = vec![0u8; 144];
let bias = vec![1.0f32];
let layer = QuantizedLinear::new(256, 1, weight_bytes.clone(), bias.clone()).expect("test");
assert_eq!(layer.in_features(), 256);
assert_eq!(layer.out_features(), 1);
assert_eq!(layer.weight_bytes().len(), 144);
assert_eq!(layer.bias().len(), 1);
assert!((layer.bias()[0] - 1.0).abs() < 1e-6);
assert_eq!(layer.memory_bytes(), 144 + 4);
}
#[test]
fn test_quantized_linear_non_256_aligned_in_features() {
let weight_bytes = vec![0u8; 288]; let bias = vec![0.0f32; 1];
let layer = QuantizedLinear::new(300, 1, weight_bytes, bias).expect("test");
assert_eq!(layer.in_features(), 300);
assert_eq!(layer.out_features(), 1);
}
#[test]
fn test_quantized_linear_forward_empty_shape_error() {
let weight_bytes = vec![0u8; 144]; let bias = vec![0.0f32; 1];
let layer = QuantizedLinear::new(256, 1, weight_bytes, bias).expect("test");
let input = Tensor::from_vec(vec![128], vec![0.1; 128]).expect("test");
let result = layer.forward(&input);
assert!(result.is_err(), "Should error on dimension mismatch");
}
#[test]
fn test_quantized_linear_forward_shape_mismatch() {
let weight_bytes = vec![0u8; 144];
let bias = vec![0.0f32; 1];
let layer = QuantizedLinear::new(256, 1, weight_bytes, bias).expect("test");
let input = Tensor::from_vec(vec![512], vec![0.1; 512]).expect("test");
let result = layer.forward(&input);
assert!(result.is_err(), "Should error on in_features mismatch");
if let Err(RealizarError::InvalidShape { reason }) = result {
assert!(reason.contains("doesn't match"));
}
}
#[test]
fn test_fused_layer_norm_linear_parallel_empty_shape_error() {
let fused = FusedLayerNormLinear::new(4, 8, 1e-5).expect("test");
let input = Tensor::from_vec(vec![3], vec![1.0, 2.0, 3.0]).expect("test");
let result = fused.forward_parallel(&input);
assert!(
result.is_err(),
"forward_parallel should error on dimension mismatch"
);
}
#[test]
fn test_fused_layer_norm_linear_parallel_dimension_mismatch() {
let fused = FusedLayerNormLinear::new(8, 4, 1e-5).expect("test");
let input = Tensor::from_vec(vec![16], vec![0.1; 16]).expect("test");
let result = fused.forward_parallel(&input);
assert!(result.is_err(), "Should error on feature_dim mismatch");
if let Err(RealizarError::InvalidShape { reason }) = result {
assert!(reason.contains("doesn't match"));
}
}
#[test]
fn test_fused_layer_norm_linear_parallel_single_row() {
let fused = FusedLayerNormLinear::new(4, 8, 1e-5).expect("test");
let input = Tensor::from_vec(vec![4], vec![1.0, 2.0, 3.0, 4.0]).expect("test");
let serial = fused.forward(&input).expect("test");
let parallel = fused.forward_parallel(&input).expect("test");
assert_eq!(serial.shape(), parallel.shape());
for i in 0..serial.data().len() {
assert!(
(serial.data()[i] - parallel.data()[i]).abs() < 1e-5,
"Single row mismatch at {}: {} vs {}",
i,
serial.data()[i],
parallel.data()[i]
);
}
}
#[test]
fn test_fused_layer_norm_linear_forward_empty_shape_error() {
let fused = FusedLayerNormLinear::new(4, 8, 1e-5).expect("test");
let input = Tensor::from_vec(vec![6], vec![1.0; 6]).expect("test");
let result = fused.forward(&input);
assert!(
result.is_err(),
"forward should error on dimension mismatch"
);
}
#[test]
fn test_fused_layer_norm_linear_weight_mutators() {
let mut fused = FusedLayerNormLinear::new(4, 8, 1e-5).expect("test");
fused.norm_weight_mut()[0] = 2.0;
assert!((fused.norm_weight_mut()[0] - 2.0).abs() < 1e-6);
fused.norm_bias_mut()[0] = 0.5;
assert!((fused.norm_bias_mut()[0] - 0.5).abs() < 1e-6);
fused.linear_weight_mut()[0] = 0.1;
assert!((fused.linear_weight_mut()[0] - 0.1).abs() < 1e-6);
fused.linear_bias_mut()[0] = 0.05;
assert!((fused.linear_bias_mut()[0] - 0.05).abs() < 1e-6);
}
#[test]
fn test_linear_1d_input_output_shape() {
let linear = Linear::new(4, 8).expect("test");
let input = Tensor::from_vec(vec![4], vec![1.0, 2.0, 3.0, 4.0]).expect("test");
let output = linear.forward(&input).expect("test");
assert_eq!(output.shape(), &[8]);
}
#[test]
fn test_linear_2d_input_preserves_batch() {
let linear = Linear::new(4, 8).expect("test");
let input = Tensor::from_vec(vec![3, 4], vec![0.1; 12]).expect("test");
let output = linear.forward(&input).expect("test");
assert_eq!(output.shape(), &[3, 8]);
}
#[test]
fn test_linear_3d_input_preserves_dims() {
let linear = Linear::new(4, 8).expect("test");
let input = Tensor::from_vec(vec![2, 3, 4], vec![0.1; 24]).expect("test");
let output = linear.forward(&input).expect("test");
assert_eq!(output.shape(), &[2, 3, 8]);
}
#[test]
fn test_linear_forward_empty_shape_error() {
let linear = Linear::new(4, 8).expect("test");
let input = Tensor::from_vec(vec![5], vec![1.0; 5]).expect("test");
let result = linear.forward(&input);
assert!(result.is_err(), "Should error on dimension mismatch");
}
include!("layer_norm_02.rs");