#[test]
fn test_forward_with_cache_f32_multi_token() {
let apr = build_f32_apr(8, 32, 16);
let mut cache = AprKVCache::new(&apr.config);
let _ = apr.forward_with_cache(1, &mut cache, 0).expect("test value should be present");
let result = apr.forward_with_cache(2, &mut cache, 1);
assert!(
result.is_ok(),
"F32 cache second token: {}",
result.unwrap_err()
);
let logits = result.expect("test value should be present");
assert_eq!(logits.len(), 16);
assert!(logits.iter().all(|x| x.is_finite()));
}
#[test]
fn test_forward_with_cache_f32_three_tokens() {
let apr = build_f32_apr(8, 32, 16);
let mut cache = AprKVCache::new(&apr.config);
for pos in 0..3 {
let result = apr.forward_with_cache(pos as u32, &mut cache, pos);
assert!(
result.is_ok(),
"F32 cache token {pos}: {}",
result.unwrap_err()
);
}
}
#[test]
fn test_forward_with_cache_q4k_first_token() {
let apr = build_q4k_apr(256, 256, 4);
let mut cache = AprKVCache::new(&apr.config);
let result = apr.forward_with_cache(1, &mut cache, 0);
assert!(
result.is_ok(),
"Q4K cache first token: {}",
result.unwrap_err()
);
let logits = result.expect("test value should be present");
assert_eq!(logits.len(), 4);
}
#[test]
fn test_forward_with_cache_q4k_multi_token() {
let apr = build_q4k_apr(256, 256, 4);
let mut cache = AprKVCache::new(&apr.config);
let _ = apr.forward_with_cache(1, &mut cache, 0).expect("test value should be present");
let result = apr.forward_with_cache(2, &mut cache, 1);
assert!(
result.is_ok(),
"Q4K cache second token: {}",
result.unwrap_err()
);
}
#[test]
fn test_forward_with_cache_q6k_weights() {
let hidden = 256;
let intermediate = 256;
let vocab = 4;
let meta = minimal_metadata(hidden, 1, 4, 4, vocab, intermediate);
let mut tensors = make_hf_tensors(
hidden,
intermediate,
4,
4,
vocab,
DTYPE_F32,
DTYPE_Q4_K,
DTYPE_F32,
);
for t in &mut tensors {
if t.name.contains("v_proj") || t.name.contains("down_proj") || t.name.contains("up_proj") {
let num_elements: usize = t.dims.iter().map(|d| *d as usize).product();
t.dtype = DTYPE_Q6_K;
t.data = make_q6k_data(num_elements);
}
}
let data = build_apr_v2(&meta, &tensors);
let apr = AprTransformer::from_apr_bytes(&data).expect("Q6K APR build failed");
let mut cache = AprKVCache::new(&apr.config);
let result = apr.forward_with_cache(1, &mut cache, 0);
assert!(
result.is_ok(),
"Q6K cache first token: {}",
result.unwrap_err()
);
}
#[test]
fn test_forward_with_cache_with_qkv_bias() {
let hidden = 8;
let intermediate = 32;
let vocab = 16;
let kv_dim = hidden; let meta = minimal_metadata(hidden, 1, 4, 4, vocab, intermediate);
let mut tensors = make_hf_tensors(
hidden,
intermediate,
4,
4,
vocab,
DTYPE_F32,
DTYPE_F32,
DTYPE_F32,
);
let qkv_out = hidden + 2 * kv_dim;
tensors.push(TensorDef {
name: "model.layers.0.self_attn.qkv_proj.bias".into(),
dtype: DTYPE_F32,
dims: vec![qkv_out as u64],
data: make_f32_data(qkv_out, 0.01),
});
let data = build_apr_v2(&meta, &tensors);
let apr = AprTransformer::from_apr_bytes(&data).expect("biased APR build failed");
let mut cache = AprKVCache::new(&apr.config);
let result = apr.forward_with_cache(1, &mut cache, 0);
assert!(result.is_ok(), "Biased cache: {}", result.unwrap_err());
}
#[test]
fn test_forward_f32_single_token() {
let apr = build_f32_apr(8, 32, 16);
let result = apr.forward(&[1]);
assert!(result.is_ok(), "Forward single: {}", result.unwrap_err());
let logits = result.expect("test value should be present");
assert_eq!(logits.len(), 16);
}
#[test]
fn test_forward_f32_multi_token() {
let apr = build_f32_apr(8, 32, 16);
let result = apr.forward(&[1, 2, 3]);
assert!(result.is_ok(), "Forward multi: {}", result.unwrap_err());
let logits = result.expect("test value should be present");
assert_eq!(logits.len(), 16);
}
#[test]
fn test_forward_q4k_single_token() {
let apr = build_q4k_apr(256, 256, 4);
let result = apr.forward(&[1]);
assert!(result.is_ok(), "Q4K forward: {}", result.unwrap_err());
let logits = result.expect("test value should be present");
assert_eq!(logits.len(), 4);
}
#[test]
fn test_generate_f32_greedy() {
let apr = build_f32_apr(8, 32, 16);
let result = apr.generate(&[1], 3);
assert!(result.is_ok(), "Generate: {}", result.unwrap_err());
let tokens = result.expect("test value should be present");
assert!(tokens.len() >= 2); assert!(tokens.len() <= 4); assert_eq!(tokens[0], 1); }
#[test]
fn test_apr_transformer_config_accessor() {
let apr = build_f32_apr(8, 32, 16);
let config = apr.config();
assert_eq!(config.hidden_dim, 8);
assert_eq!(config.intermediate_dim, 32);
assert_eq!(config.vocab_size, 16);
}
#[test]
fn test_apr_transformer_num_parameters() {
let apr = build_f32_apr(8, 32, 16);
let params = apr.num_parameters();
assert!(params > 0);
assert!(params >= 16 * 8 * 2);
}
fn build_gelu_apr(hidden: usize, intermediate: usize, vocab: usize) -> AprTransformer {
let meta = minimal_metadata(hidden, 1, 4, 4, vocab, intermediate);
let mut tensors = make_hf_tensors(
hidden,
intermediate,
4,
4,
vocab,
DTYPE_F32,
DTYPE_F32,
DTYPE_F32,
);
tensors.retain(|t| !t.name.contains("gate_proj"));
let data = build_apr_v2(&meta, &tensors);
AprTransformer::from_apr_bytes(&data).expect("GELU APR build failed")
}
#[test]
fn test_forward_gelu_model_single_token() {
let apr = build_gelu_apr(8, 32, 16);
assert!(apr.layers[0].ffn_gate_weight.is_none());
let result = apr.forward(&[1]);
assert!(result.is_ok(), "GELU forward: {}", result.unwrap_err());
let logits = result.expect("test value should be present");
assert_eq!(logits.len(), 16);
}
#[test]
fn test_forward_with_cache_gelu_model() {
let apr = build_gelu_apr(8, 32, 16);
let mut cache = AprKVCache::new(&apr.config);
let result = apr.forward_with_cache(1, &mut cache, 0);
assert!(result.is_ok(), "GELU cache first: {}", result.unwrap_err());
let result = apr.forward_with_cache(2, &mut cache, 1);
assert!(result.is_ok(), "GELU cache second: {}", result.unwrap_err());
}
#[test]
fn test_forward_with_cache_no_ffn_norm() {
let hidden = 8;
let intermediate = 32;
let vocab = 16;
let meta = minimal_metadata(hidden, 1, 4, 4, vocab, intermediate);
let mut tensors = make_hf_tensors(
hidden,
intermediate,
4,
4,
vocab,
DTYPE_F32,
DTYPE_F32,
DTYPE_F32,
);
tensors.retain(|t| !t.name.contains("post_attention_layernorm"));
let data = build_apr_v2(&meta, &tensors);
let apr = AprTransformer::from_apr_bytes(&data).expect("no-ffn-norm APR build failed");
assert!(apr.layers[0].ffn_norm_weight.is_none());
let mut cache = AprKVCache::new(&apr.config);
let result = apr.forward_with_cache(1, &mut cache, 0);
assert!(result.is_ok(), "No FFN norm: {}", result.unwrap_err());
}
#[test]
fn test_forward_with_cache_q4k_with_bias() {
let hidden = 256;
let intermediate = 256;
let vocab = 4;
let kv_dim = hidden; let qkv_out = hidden + 2 * kv_dim;
let meta = minimal_metadata(hidden, 1, 4, 4, vocab, intermediate);
let mut tensors = make_hf_tensors(
hidden,
intermediate,
4,
4,
vocab,
DTYPE_F32,
DTYPE_Q4_K,
DTYPE_F32,
);
tensors.push(TensorDef {
name: "model.layers.0.self_attn.qkv_proj.bias".into(),
dtype: DTYPE_F32,
dims: vec![qkv_out as u64],
data: make_f32_data(qkv_out, 0.01),
});
let data = build_apr_v2(&meta, &tensors);
let apr = AprTransformer::from_apr_bytes(&data).expect("Q4K+bias APR build failed");
assert!(apr.q4k_layers.is_some());
assert!(apr.layers[0].qkv_bias.is_some());
let mut cache = AprKVCache::new(&apr.config);
let result = apr.forward_with_cache(1, &mut cache, 0);
assert!(result.is_ok(), "Q4K+bias cache: {}", result.unwrap_err());
}
#[test]
fn test_forward_with_cache_q4k_lm_head() {
let hidden = 256;
let intermediate = 256;
let vocab = 256;
let meta = minimal_metadata(hidden, 1, 4, 4, vocab, intermediate);
let tensors = make_hf_tensors(
hidden,
intermediate,
4,
4,
vocab,
DTYPE_F32,
DTYPE_F32,
DTYPE_Q4_K,
);
let data = build_apr_v2(&meta, &tensors);
let apr = AprTransformer::from_apr_bytes(&data).expect("Q4K lm_head APR build failed");
assert!(apr.lm_head_weight_q4k.is_some());
let mut cache = AprKVCache::new(&apr.config);
let result = apr.forward_with_cache(1, &mut cache, 0);
assert!(result.is_ok(), "Q4K lm_head cache: {}", result.unwrap_err());
}
#[test]
fn test_forward_with_cache_q6k_lm_head() {
let hidden = 256;
let intermediate = 256;
let vocab = 256;
let meta = minimal_metadata(hidden, 1, 4, 4, vocab, intermediate);
let mut tensors = make_hf_tensors(
hidden,
intermediate,
4,
4,
vocab,
DTYPE_F32,
DTYPE_F32,
DTYPE_F32,
);
for t in &mut tensors {
if t.name == "lm_head.weight" {
let num_elements: usize = t.dims.iter().map(|d| *d as usize).product();
t.dtype = DTYPE_Q6_K;
t.data = make_q6k_data(num_elements);
}
}
let data = build_apr_v2(&meta, &tensors);
let apr = AprTransformer::from_apr_bytes(&data).expect("Q6K lm_head APR build failed");
assert!(apr.lm_head_weight_q6k.is_some());
let mut cache = AprKVCache::new(&apr.config);
let result = apr.forward_with_cache(1, &mut cache, 0);
assert!(result.is_ok(), "Q6K lm_head cache: {}", result.unwrap_err());
}