#[test]
fn test_model_generate_respects_max_tokens() {
let config = ModelConfig {
vocab_size: 10,
hidden_dim: 4,
num_heads: 1,
num_layers: 1,
intermediate_dim: 16,
eps: 1e-5,
};
let model = Model::new(config).expect("test");
let gen_config = GenerationConfig::greedy().with_max_tokens(3);
let tokens = model.generate(&[0, 1], &gen_config).expect("test");
assert!(tokens.len() <= 5);
}
#[test]
fn test_model_generate_with_eos() {
let config = ModelConfig {
vocab_size: 10,
hidden_dim: 4,
num_heads: 1,
num_layers: 1,
intermediate_dim: 16,
eps: 1e-5,
};
let model = Model::new(config).expect("test");
let gen_config = GenerationConfig::greedy()
.with_max_tokens(100)
.with_eos_token_id(5);
let tokens = model.generate(&[0], &gen_config).expect("test");
assert!(tokens.len() <= 101);
}
#[test]
fn test_model_generate_empty_prompt_error() {
let config = ModelConfig {
vocab_size: 10,
hidden_dim: 4,
num_heads: 1,
num_layers: 1,
intermediate_dim: 16,
eps: 1e-5,
};
let model = Model::new(config).expect("test");
let gen_config = GenerationConfig::greedy();
let result = model.generate(&[], &gen_config);
assert!(result.is_err());
}
#[test]
fn test_model_generate_deterministic_with_seed() {
let config = ModelConfig {
vocab_size: 20,
hidden_dim: 4,
num_heads: 1,
num_layers: 1,
intermediate_dim: 16,
eps: 1e-5,
};
let model = Model::new(config).expect("test");
let gen_config = GenerationConfig::greedy()
.with_max_tokens(5)
.with_seed(12345);
let tokens1 = model.generate(&[0], &gen_config).expect("test");
let tokens2 = model.generate(&[0], &gen_config).expect("test");
assert_eq!(tokens1, tokens2);
}
#[test]
fn test_model_generate_top_k() {
let config = ModelConfig {
vocab_size: 20,
hidden_dim: 4,
num_heads: 1,
num_layers: 1,
intermediate_dim: 16,
eps: 1e-5,
};
let model = Model::new(config).expect("test");
let gen_config = GenerationConfig::top_k(5).with_max_tokens(3).with_seed(42);
let tokens = model.generate(&[0], &gen_config).expect("test");
assert!(tokens.len() <= 4);
for &token in &tokens {
assert!(token < 20);
}
}
#[test]
fn test_multi_head_attention_creation_mha() {
let mha = MultiHeadAttention::mha(64, 8).expect("test");
assert_eq!(mha.num_heads(), 8);
assert_eq!(mha.num_kv_heads(), 8);
assert_eq!(mha.head_dim(), 8); assert_eq!(mha.hidden_dim(), 64);
assert!(mha.is_mha());
assert!(!mha.is_mqa());
assert!(!mha.is_gqa());
}
#[test]
fn test_multi_head_attention_creation_mqa() {
let mqa = MultiHeadAttention::mqa(64, 8).expect("test");
assert_eq!(mqa.num_heads(), 8);
assert_eq!(mqa.num_kv_heads(), 1);
assert_eq!(mqa.head_dim(), 8);
assert_eq!(mqa.hidden_dim(), 64);
assert!(mqa.is_mqa());
assert!(!mqa.is_mha());
assert!(!mqa.is_gqa());
}
#[test]
fn test_multi_head_attention_creation_gqa() {
let gqa = MultiHeadAttention::gqa(64, 8, 2).expect("test");
assert_eq!(gqa.num_heads(), 8);
assert_eq!(gqa.num_kv_heads(), 2);
assert_eq!(gqa.head_dim(), 8);
assert_eq!(gqa.hidden_dim(), 64);
assert!(gqa.is_gqa());
assert!(!gqa.is_mha());
assert!(!gqa.is_mqa());
}
#[test]
fn test_multi_head_attention_zero_hidden_dim_error() {
let result = MultiHeadAttention::new(0, 8, 8);
assert!(result.is_err());
}
#[test]
fn test_multi_head_attention_zero_num_heads_error() {
let result = MultiHeadAttention::new(64, 0, 1);
assert!(result.is_err());
}
#[test]
fn test_multi_head_attention_zero_num_kv_heads_error() {
let result = MultiHeadAttention::new(64, 8, 0);
assert!(result.is_err());
}
#[test]
fn test_multi_head_attention_kv_heads_too_large_error() {
let result = MultiHeadAttention::new(64, 8, 16);
assert!(result.is_err());
}
#[test]
fn test_multi_head_attention_indivisible_error() {
let result = MultiHeadAttention::new(65, 8, 8);
assert!(result.is_err());
}
#[test]
fn test_multi_head_attention_heads_not_divisible_error() {
let result = MultiHeadAttention::new(64, 8, 3);
assert!(result.is_err());
}
#[test]
fn test_multi_head_attention_mha_forward() {
let mha = MultiHeadAttention::mha(8, 2).expect("test");
let input = Tensor::from_vec(
vec![2, 8],
vec![
1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, ],
)
.expect("test");
let output = mha.forward(&input).expect("test");
assert_eq!(output.shape(), &[2, 8]);
}
#[test]
fn test_multi_head_attention_mqa_forward() {
let mqa = MultiHeadAttention::mqa(8, 2).expect("test");
let input = Tensor::from_vec(
vec![2, 8],
vec![
1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, ],
)
.expect("test");
let output = mqa.forward(&input).expect("test");
assert_eq!(output.shape(), &[2, 8]);
}
#[test]
fn test_multi_head_attention_shape_validation() {
let mha = MultiHeadAttention::mha(8, 2).expect("test");
let input_1d = Tensor::from_vec(vec![8], vec![1.0; 8]).expect("test");
let result = mha.forward(&input_1d);
assert!(result.is_err());
let input_wrong_dim = Tensor::from_vec(vec![2, 16], vec![1.0; 32]).expect("test");
let result = mha.forward(&input_wrong_dim);
assert!(result.is_err());
}
#[test]
fn test_multi_head_attention_mha_vs_mqa_shape_consistency() {
let mha = MultiHeadAttention::mha(16, 4).expect("test");
let mqa = MultiHeadAttention::mqa(16, 4).expect("test");
let input = Tensor::from_vec(vec![3, 16], vec![0.5; 48]).expect("test");
let multi_head_output = mha.forward(&input).expect("test");
let multi_query_output = mqa.forward(&input).expect("test");
assert_eq!(multi_head_output.shape(), &[3, 16]);
assert_eq!(multi_query_output.shape(), &[3, 16]);
assert_eq!(multi_head_output.shape(), multi_query_output.shape());
}
#[test]
fn test_multi_head_attention_single_head() {
let mha = MultiHeadAttention::mha(8, 1).expect("test");
let input = Tensor::from_vec(vec![2, 8], vec![0.5; 16]).expect("test");
let output = mha.forward(&input).expect("test");
assert_eq!(output.shape(), &[2, 8]);
}
#[test]
fn test_multi_head_attention_mqa_kv_sharing() {
let mqa = MultiHeadAttention::mqa(32, 8).expect("test");
let input = Tensor::from_vec(vec![4, 32], vec![0.1; 128]).expect("test");
let output = mqa.forward(&input).expect("test");
assert_eq!(output.shape(), &[4, 32]);
}
#[test]
fn test_multi_head_attention_long_sequence() {
let mha = MultiHeadAttention::mha(16, 4).expect("test");
let input = Tensor::from_vec(vec![10, 16], vec![0.3; 160]).expect("test");
let output = mha.forward(&input).expect("test");
assert_eq!(output.shape(), &[10, 16]);
}
#[test]
fn test_multi_head_attention_mqa_memory_efficiency() {
let mqa = MultiHeadAttention::mqa(64, 16).expect("test");
let input = Tensor::from_vec(vec![2, 64], vec![0.2; 128]).expect("test");
let output = mqa.forward(&input).expect("test");
assert_eq!(output.shape(), &[2, 64]);
assert_eq!(output.data().len(), 128); }
#[test]
fn test_multi_head_attention_gqa_forward() {
let gqa = MultiHeadAttention::gqa(32, 8, 2).expect("test");
let input = Tensor::from_vec(vec![3, 32], vec![0.1; 96]).expect("test");
let output = gqa.forward(&input).expect("test");
assert_eq!(output.shape(), &[3, 32]);
}
#[test]
fn test_multi_head_attention_gqa_shape_consistency() {
let mha = MultiHeadAttention::mha(64, 8).expect("test");
let mqa = MultiHeadAttention::mqa(64, 8).expect("test");
let gqa = MultiHeadAttention::gqa(64, 8, 2).expect("test");
let input = Tensor::from_vec(vec![4, 64], vec![0.5; 256]).expect("test");
let multi_head_out = mha.forward(&input).expect("test");
let multi_query_out = mqa.forward(&input).expect("test");
let grouped_query_out = gqa.forward(&input).expect("test");
assert_eq!(multi_head_out.shape(), &[4, 64]);
assert_eq!(multi_query_out.shape(), &[4, 64]);
assert_eq!(grouped_query_out.shape(), &[4, 64]);
assert_eq!(multi_head_out.shape(), multi_query_out.shape());
assert_eq!(multi_head_out.shape(), grouped_query_out.shape());
}
#[test]
fn test_multi_head_attention_gqa_different_group_sizes() {
let gqa1 = MultiHeadAttention::gqa(128, 16, 4).expect("test");
let input = Tensor::from_vec(vec![2, 128], vec![0.3; 256]).expect("test");
let output1 = gqa1.forward(&input).expect("test");
assert_eq!(output1.shape(), &[2, 128]);
let gqa2 = MultiHeadAttention::gqa(128, 16, 8).expect("test");
let output2 = gqa2.forward(&input).expect("test");
assert_eq!(output2.shape(), &[2, 128]);
}
#[test]
#[ignore = "performance benchmark - run explicitly with --include-ignored"]
fn test_phase3_acceptance_tokens_per_second() {
use crate::generate::GenerationConfig;
use std::time::Instant;
let config = ModelConfig {
vocab_size: 100, hidden_dim: 64, num_heads: 4, num_layers: 2, intermediate_dim: 128,
eps: 1e-5,
};
let model = Model::new(config).expect("test");
let prompt = vec![1, 5, 10];
let gen_config = GenerationConfig::greedy().with_max_tokens(5);
let _ = model.generate(&prompt, &gen_config).expect("test");
let tokens_per_run = 20;
let num_runs = 10;
let gen_config = GenerationConfig::greedy().with_max_tokens(tokens_per_run);
let start = Instant::now();
for _ in 0..num_runs {
let _ = model.generate(&prompt, &gen_config).expect("test");
}
let elapsed = start.elapsed();
let total_tokens = tokens_per_run * num_runs;
let tok_per_sec = total_tokens as f64 / elapsed.as_secs_f64();
assert!(
tok_per_sec >= 25.0,
"Phase 3 acceptance FAILED: {:.1} tok/s < 25.0 tok/s target. \
Note: Full optimization requires integrating Flash Attention v2 \
and FusedLayerNormLinear into Model::forward()",
tok_per_sec
);
eprintln!(
"Phase 3 acceptance PASSED: {:.1} tok/s (target: ≥25.0 tok/s)",
tok_per_sec
);
}