#[test]
fn test_phase3_flash_attention_v2_performance() {
use std::time::Instant;
let head_dim = 64;
let seq_len = 32;
let attn = Attention::new(head_dim).expect("test");
let q = Tensor::from_vec(vec![seq_len, head_dim], vec![0.1; seq_len * head_dim]).expect("test");
let k = Tensor::from_vec(vec![seq_len, head_dim], vec![0.2; seq_len * head_dim]).expect("test");
let v = Tensor::from_vec(vec![seq_len, head_dim], vec![0.3; seq_len * head_dim]).expect("test");
let _ = attn.flash_forward_v2(&q, &k, &v, 8).expect("test");
let _ = attn.flash_forward_parallel(&q, &k, &v, 8).expect("test");
let iterations = 100;
let start = Instant::now();
for _ in 0..iterations {
let _ = attn.flash_forward_v2(&q, &k, &v, 8).expect("test");
}
let v2_time = start.elapsed();
let start = Instant::now();
for _ in 0..iterations {
let _ = attn.flash_forward_parallel(&q, &k, &v, 8).expect("test");
}
let parallel_time = start.elapsed();
let v2_us = v2_time.as_micros() as f64 / iterations as f64;
let parallel_us = parallel_time.as_micros() as f64 / iterations as f64;
eprintln!(
"Flash Attention v2: {:.2}us/iter, Parallel: {:.2}us/iter",
v2_us, parallel_us
);
assert!(v2_us > 0.0, "v2 should have measurable time");
assert!(parallel_us > 0.0, "parallel should have measurable time");
}
#[test]
fn test_phase3_fused_layernorm_linear_performance() {
use std::time::Instant;
let feature_dim = 256;
let out_features = 512;
let batch_size = 32;
let fused = FusedLayerNormLinear::new(feature_dim, out_features, 1e-5).expect("test");
let input = Tensor::from_vec(
vec![batch_size, feature_dim],
vec![0.5; batch_size * feature_dim],
)
.expect("test");
let _ = fused.forward(&input).expect("test");
let _ = fused.forward_parallel(&input).expect("test");
let iterations = 100;
let start = Instant::now();
for _ in 0..iterations {
let _ = fused.forward(&input).expect("test");
}
let fused_time = start.elapsed();
let start = Instant::now();
for _ in 0..iterations {
let _ = fused.forward_parallel(&input).expect("test");
}
let parallel_time = start.elapsed();
let fused_us = fused_time.as_micros() as f64 / iterations as f64;
let parallel_us = parallel_time.as_micros() as f64 / iterations as f64;
eprintln!(
"FusedLayerNormLinear: {:.2}us/iter, Parallel: {:.2}us/iter",
fused_us, parallel_us
);
assert!(fused_us > 0.0, "fused should have measurable time");
assert!(parallel_us > 0.0, "parallel should have measurable time");
}
#[test]
fn test_quantized_linear_creation() {
let in_features = 256;
let out_features = 4;
let bytes_per_row = 144; let weight_bytes = vec![0u8; out_features * bytes_per_row];
let bias = vec![0.0f32; out_features];
let layer = QuantizedLinear::new(in_features, out_features, weight_bytes, bias);
assert!(
layer.is_ok(),
"Should create QuantizedLinear from Q4_K bytes"
);
let layer = layer.expect("test");
assert_eq!(layer.in_features(), in_features);
assert_eq!(layer.out_features(), out_features);
}
#[test]
fn test_quantized_linear_forward() {
let in_features = 256;
let out_features = 4;
let bytes_per_row = 144;
let weight_bytes = vec![0u8; out_features * bytes_per_row];
let bias = vec![1.0f32; out_features];
let layer = QuantizedLinear::new(in_features, out_features, weight_bytes, bias)
.expect("Should create layer");
let input = Tensor::from_vec(vec![in_features], vec![1.0f32; in_features])
.expect("Should create input");
let output = layer.forward(&input).expect("Forward should work");
assert_eq!(output.shape(), &[out_features]);
for &val in output.data() {
assert!(
(val - 1.0).abs() < 1e-5,
"Output should equal bias with zero weights"
);
}
}
#[test]
fn test_quantized_linear_batch_forward() {
let in_features = 256;
let out_features = 4;
let batch_size = 8;
let bytes_per_row = 144;
let weight_bytes = vec![0u8; out_features * bytes_per_row];
let bias = vec![2.0f32; out_features];
let layer = QuantizedLinear::new(in_features, out_features, weight_bytes, bias)
.expect("Should create layer");
let input = Tensor::from_vec(
vec![batch_size, in_features],
vec![1.0f32; batch_size * in_features],
)
.expect("Should create batch input");
let output = layer.forward(&input).expect("Batch forward should work");
assert_eq!(output.shape(), &[batch_size, out_features]);
}
#[test]
fn test_quantized_linear_memory_efficiency() {
let in_features = 4096; let out_features = 4096;
let f32_bytes = in_features * out_features * std::mem::size_of::<f32>();
let super_blocks_per_row = in_features.div_ceil(256);
let q4k_bytes = out_features * super_blocks_per_row * 144;
let ratio = f32_bytes as f64 / q4k_bytes as f64;
assert!(
ratio > 6.0,
"Q4_K should be >6x smaller than f32: ratio={}",
ratio
);
eprintln!(
"Memory efficiency: f32={} bytes, Q4_K={} bytes, ratio={:.2}x",
f32_bytes, q4k_bytes, ratio
);
}
#[test]
fn test_sliding_window_attention_new() {
let swa = SlidingWindowAttention::new(64, 4096).expect("test");
assert_eq!(swa.head_dim(), 64);
assert_eq!(swa.window_size(), 4096);
assert!((swa.scale() - 0.125).abs() < 1e-6); }
#[test]
fn test_sliding_window_attention_new_errors() {
assert!(SlidingWindowAttention::new(0, 4096).is_err());
assert!(SlidingWindowAttention::new(64, 0).is_err());
}
#[test]
fn test_sliding_window_attention_forward_basic() {
let swa = SlidingWindowAttention::new(4, 3).expect("test");
let query_data: Vec<f32> = (0..20).map(|i| i as f32 * 0.1).collect();
let key_data: Vec<f32> = (0..20).map(|i| i as f32 * 0.1).collect();
let value_data: Vec<f32> = (0..20).map(|i| (i % 4) as f32).collect();
let query = Tensor::from_vec(vec![5, 4], query_data).expect("test");
let key = Tensor::from_vec(vec![5, 4], key_data).expect("test");
let value = Tensor::from_vec(vec![5, 4], value_data).expect("test");
let output = swa.forward(&query, &key, &value).expect("test");
assert_eq!(output.size(), 20); }
#[test]
fn test_sliding_window_attention_causal_masking() {
let swa = SlidingWindowAttention::new(2, 10).expect("test"); let query = Tensor::from_vec(vec![3, 2], vec![1.0, 0.0, 1.0, 0.0, 1.0, 0.0]).expect("test");
let key = Tensor::from_vec(vec![3, 2], vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0]).expect("test");
let value = Tensor::from_vec(vec![3, 2], vec![1.0, 0.0, 0.0, 1.0, 0.5, 0.5]).expect("test");
let output = swa.forward(&query, &key, &value).expect("test");
assert_eq!(output.size(), 6);
let data = output.data();
assert!(data[0].abs() > 0.0 || data[1].abs() > 0.0);
}
#[test]
fn test_sliding_window_attention_window_boundary() {
let swa = SlidingWindowAttention::new(2, 2).expect("test");
let query = Tensor::from_vec(vec![5, 2], vec![1.0; 10]).expect("test");
let key = Tensor::from_vec(vec![5, 2], vec![1.0; 10]).expect("test");
let value_data: Vec<f32> = (0..10).map(|i| i as f32).collect();
let value = Tensor::from_vec(vec![5, 2], value_data).expect("test");
let output = swa.forward(&query, &key, &value).expect("test");
assert_eq!(output.size(), 10);
}
#[test]
fn test_sliding_window_attention_effective_context() {
let swa = SlidingWindowAttention::new(64, 4).expect("test");
assert_eq!(swa.effective_context(0, 10), 1);
assert_eq!(swa.effective_context(3, 10), 4);
assert_eq!(swa.effective_context(7, 10), 4);
assert_eq!(swa.effective_context(2, 3), 3);
}
#[test]
fn test_sliding_window_attention_memory_ratio() {
let swa = SlidingWindowAttention::new(64, 4096).expect("test");
let ratio_short = swa.memory_ratio(1000);
assert!(
ratio_short > 0.9,
"Short sequences should use ~full attention"
);
let ratio_long = swa.memory_ratio(100_000);
let expected = 4096.0 / 100_000.0;
assert!(
(ratio_long - expected).abs() < 0.01,
"Long sequences should use ~window_size/seq_len memory: got {}, expected {}",
ratio_long,
expected
);
}
#[test]
fn test_sliding_window_attention_error_mismatched_kv() {
let swa = SlidingWindowAttention::new(4, 3).expect("test");
let query = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let key = Tensor::from_vec(vec![3, 4], vec![1.0; 12]).expect("test");
let value = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = swa.forward(&query, &key, &value);
assert!(result.is_err());
}
#[test]
fn test_sliding_window_attention_error_bad_head_dim() {
let swa = SlidingWindowAttention::new(4, 3).expect("test");
let query = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let key = Tensor::from_vec(vec![2, 3], vec![1.0; 6]).expect("test");
let value = Tensor::from_vec(vec![2, 4], vec![1.0; 8]).expect("test");
let result = swa.forward(&query, &key, &value);
assert!(result.is_err());
}
#[test]
fn test_sliding_window_attention_bidirectional() {
let swa = SlidingWindowAttention::new(2, 4).expect("test");
let query = Tensor::from_vec(vec![5, 2], vec![1.0; 10]).expect("test");
let key = Tensor::from_vec(vec![5, 2], vec![1.0; 10]).expect("test");
let value_data: Vec<f32> = (0..10).map(|i| i as f32).collect();
let value = Tensor::from_vec(vec![5, 2], value_data).expect("test");
let output_causal = swa.forward(&query, &key, &value).expect("test");
let output_bidir = swa
.forward_with_mask(&query, &key, &value, false)
.expect("test");
assert_eq!(output_causal.size(), output_bidir.size());
assert!(output_causal.data().iter().any(|&x| x.abs() > 0.0));
assert!(output_bidir.data().iter().any(|&x| x.abs() > 0.0));
}
#[test]
fn test_sliding_window_attention_forward_with_mask_causal() {
let swa = SlidingWindowAttention::new(2, 3).expect("test");
let query = Tensor::from_vec(vec![3, 2], vec![1.0; 6]).expect("test");
let key = Tensor::from_vec(vec![3, 2], vec![1.0; 6]).expect("test");
let value = Tensor::from_vec(vec![3, 2], vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]).expect("test");
let output_forward = swa.forward(&query, &key, &value).expect("test");
let output_mask = swa
.forward_with_mask(&query, &key, &value, true)
.expect("test");
for (a, b) in output_forward.data().iter().zip(output_mask.data().iter()) {
assert!(
(a - b).abs() < 1e-6,
"Causal outputs should match: {} vs {}",
a,
b
);
}
}
#[test]
fn test_sliding_window_attention_single_token() {
let swa = SlidingWindowAttention::new(4, 3).expect("test");
let query = Tensor::from_vec(vec![1, 4], vec![1.0, 2.0, 3.0, 4.0]).expect("test");
let key = Tensor::from_vec(vec![1, 4], vec![1.0, 2.0, 3.0, 4.0]).expect("test");
let value = Tensor::from_vec(vec![1, 4], vec![0.5, 0.5, 0.5, 0.5]).expect("test");
let output = swa.forward(&query, &key, &value).expect("test");
assert_eq!(output.size(), 4);
let data = output.data();
for &v in data {
assert!((v - 0.5).abs() < 1e-6);
}
}
#[test]
fn test_fused_qkv_attention_basic() {
let fused = FusedQKVAttention::new(4, 64).expect("test");
let input = Tensor::from_vec(vec![8, 64], vec![0.1; 8 * 64]).expect("test");
let output = fused.forward(&input).expect("test");
assert_eq!(output.shape(), &[8, 64]);
}