#![allow(clippy::many_single_char_names)]
use crate::layers::*;
#[test]
fn test_attention_forward_empty_query_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test");
let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.forward(&q, &k, &v);
assert!(result.is_err(), "Should error on Q head_dim mismatch");
}
#[test]
fn test_attention_forward_empty_key_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test"); let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.forward(&q, &k, &v);
assert!(result.is_err(), "Should error on K head_dim mismatch");
}
#[test]
fn test_attention_forward_empty_value_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test");
let result = attn.forward(&q, &k, &v);
assert!(result.is_err(), "Should error on V head_dim mismatch");
}
#[test]
fn test_attention_forward_single_dim_tensors() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![4], vec![1.0, 2.0, 3.0, 4.0]).expect("test");
let k = Tensor::from_vec(vec![4], vec![1.0, 0.0, 0.0, 0.0]).expect("test");
let v = Tensor::from_vec(vec![4], vec![5.0, 6.0, 7.0, 8.0]).expect("test");
let output = attn.forward(&q, &k, &v).expect("test");
assert_eq!(output.shape(), &[1, 4]);
for i in 0..4 {
assert!((output.data()[i] - v.data()[i]).abs() < 1e-6);
}
}
#[test]
fn test_attention_scale_computation() {
for head_dim in [1, 4, 16, 64, 128] {
let attn = Attention::new(head_dim).expect("test");
let expected_scale = 1.0 / (head_dim as f32).sqrt();
assert!(
(attn.scale() - expected_scale).abs() < 1e-6,
"Scale for head_dim={} should be {}",
head_dim,
expected_scale
);
}
}
#[test]
fn test_flash_forward_empty_q_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test"); let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.flash_forward(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward should error on Q head_dim mismatch"
);
}
#[test]
fn test_flash_forward_empty_k_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test"); let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.flash_forward(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward should error on K head_dim mismatch"
);
}
#[test]
fn test_flash_forward_empty_v_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test");
let result = attn.flash_forward(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward should error on V head_dim mismatch"
);
}
#[test]
fn test_flash_forward_kv_seq_len_mismatch() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![3, 4], vec![1.0; 12]).expect("test"); let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.flash_forward(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward should error on K/V seq_len mismatch"
);
}
#[test]
fn test_flash_forward_single_position() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![4], vec![1.0, 0.0, 0.0, 0.0]).expect("test");
let k = Tensor::from_vec(vec![4], vec![1.0, 0.0, 0.0, 0.0]).expect("test");
let v = Tensor::from_vec(vec![4], vec![2.0, 4.0, 6.0, 8.0]).expect("test");
let output = attn.flash_forward(&q, &k, &v, 1).expect("test");
assert_eq!(output.shape(), &[1, 4]);
for i in 0..4 {
assert!((output.data()[i] - v.data()[i]).abs() < 1e-6);
}
}
#[test]
fn test_flash_forward_block_size_larger_than_seq() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v =
Tensor::from_vec(vec![2, 4], vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]).expect("test");
let output = attn.flash_forward(&q, &k, &v, 10).expect("test");
assert_eq!(output.shape(), &[2, 4]);
}
#[test]
fn test_flash_forward_v2_empty_q_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test"); let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.flash_forward_v2(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward_v2 should error on Q head_dim mismatch"
);
}
#[test]
fn test_flash_forward_v2_empty_k_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test"); let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.flash_forward_v2(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward_v2 should error on K head_dim mismatch"
);
}
#[test]
fn test_flash_forward_v2_empty_v_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test");
let result = attn.flash_forward_v2(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward_v2 should error on V head_dim mismatch"
);
}
#[test]
fn test_flash_forward_v2_kv_seq_len_mismatch() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![3, 4], vec![1.0; 12]).expect("test"); let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.flash_forward_v2(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward_v2 should error on K/V seq_len mismatch"
);
}
#[test]
fn test_flash_forward_v2_single_position() {
let attn = Attention::new(8).expect("test");
let q = Tensor::from_vec(vec![8], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![8], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![8], vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]).expect("test");
let output = attn.flash_forward_v2(&q, &k, &v, 1).expect("test");
assert_eq!(output.shape(), &[1, 8]);
for i in 0..8 {
assert!((output.data()[i] - v.data()[i]).abs() < 1e-5);
}
}
#[test]
fn test_flash_forward_v2_simd_aligned_dimensions() {
let attn = Attention::new(16).expect("test");
let q = Tensor::from_vec(vec![4, 16], vec![0.1; 64]).expect("test");
let k = Tensor::from_vec(vec![4, 16], vec![0.2; 64]).expect("test");
let v = Tensor::from_vec(vec![4, 16], vec![0.3; 64]).expect("test");
let standard = attn.forward(&q, &k, &v).expect("test");
let v2 = attn.flash_forward_v2(&q, &k, &v, 2).expect("test");
assert_eq!(standard.shape(), v2.shape());
for i in 0..standard.data().len() {
assert!(
(standard.data()[i] - v2.data()[i]).abs() < 1e-4,
"SIMD aligned mismatch at {}: {} vs {}",
i,
standard.data()[i],
v2.data()[i]
);
}
}
#[test]
fn test_flash_forward_v2_simd_unaligned_dimensions() {
let attn = Attention::new(7).expect("test");
let q = Tensor::from_vec(vec![3, 7], vec![0.1; 21]).expect("test");
let k = Tensor::from_vec(vec![3, 7], vec![0.2; 21]).expect("test");
let v = Tensor::from_vec(vec![3, 7], vec![0.3; 21]).expect("test");
let standard = attn.forward(&q, &k, &v).expect("test");
let v2 = attn.flash_forward_v2(&q, &k, &v, 2).expect("test");
assert_eq!(standard.shape(), v2.shape());
for i in 0..standard.data().len() {
assert!(
(standard.data()[i] - v2.data()[i]).abs() < 1e-4,
"SIMD unaligned mismatch at {}: {} vs {}",
i,
standard.data()[i],
v2.data()[i]
);
}
}
#[test]
fn test_flash_forward_parallel_empty_q_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test"); let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.flash_forward_parallel(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward_parallel should error on Q head_dim mismatch"
);
}
#[test]
fn test_flash_forward_parallel_empty_k_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test"); let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.flash_forward_parallel(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward_parallel should error on K head_dim mismatch"
);
}
#[test]
fn test_flash_forward_parallel_empty_v_shape_error() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test");
let result = attn.flash_forward_parallel(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward_parallel should error on V head_dim mismatch"
);
}
#[test]
fn test_flash_forward_parallel_kv_seq_len_mismatch() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![3, 4], vec![1.0; 12]).expect("test"); let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = attn.flash_forward_parallel(&q, &k, &v, 2);
assert!(
result.is_err(),
"flash_forward_parallel should error on K/V seq_len mismatch"
);
}
#[test]
fn test_flash_forward_parallel_single_position() {
let attn = Attention::new(4).expect("test");
let q = Tensor::from_vec(vec![4], vec![1.0, 0.0, 0.0, 0.0]).expect("test");
let k = Tensor::from_vec(vec![4], vec![1.0, 0.0, 0.0, 0.0]).expect("test");
let v = Tensor::from_vec(vec![4], vec![2.0, 4.0, 6.0, 8.0]).expect("test");
let output = attn.flash_forward_parallel(&q, &k, &v, 1).expect("test");
assert_eq!(output.shape(), &[1, 4]);
for i in 0..4 {
assert!((output.data()[i] - v.data()[i]).abs() < 1e-6);
}
}
#[test]
fn test_sliding_window_bidirectional_shape_errors() {
let swa = SlidingWindowAttention::new(4, 3).expect("test");
let q = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test");
let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = swa.forward_with_mask(&q, &k, &v, false);
assert!(
result.is_err(),
"Bidirectional should error on Q head_dim mismatch"
);
}
#[test]
fn test_sliding_window_bidirectional_k_shape_error() {
let swa = SlidingWindowAttention::new(4, 3).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test");
let v = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = swa.forward_with_mask(&q, &k, &v, false);
assert!(
result.is_err(),
"Bidirectional should error on K head_dim mismatch"
);
}
#[test]
fn test_sliding_window_bidirectional_v_shape_error() {
let swa = SlidingWindowAttention::new(4, 3).expect("test");
let q = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let k = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let v = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test");
let result = swa.forward_with_mask(&q, &k, &v, false);
assert!(
result.is_err(),
"Bidirectional should error on V head_dim mismatch"
);
}
include!("sliding_window_02.rs");