use std::io::Write;
use crate::gguf::model::OwnedQuantizedModel;
use crate::gguf::test_factory::build_executable_pygmy_gguf;
use crate::gguf::transformer::QuantizedGGUFTransformer;
use crate::gguf::types::GGUFModel;
use crate::gguf::MappedGGUFModel;
use crate::gguf::OwnedQuantizedKVCache;
fn create_pygmy_temp_file() -> (tempfile::NamedTempFile, Vec<u8>) {
let gguf_data = build_executable_pygmy_gguf();
let mut temp_file = tempfile::NamedTempFile::new().expect("Failed to create temp file");
temp_file
.write_all(&gguf_data)
.expect("Failed to write temp file");
temp_file.flush().expect("Failed to flush temp file");
(temp_file, gguf_data)
}
#[test]
fn test_active_pygmy_load_quantized_transformer() {
let gguf_data = build_executable_pygmy_gguf();
let gguf_model = GGUFModel::from_bytes(&gguf_data).expect("Should parse GGUF");
let result = QuantizedGGUFTransformer::from_gguf(&gguf_model, &gguf_data);
assert!(
result.is_ok(),
"Failed to load QuantizedGGUFTransformer: {:?}",
result.err()
);
let transformer = result.expect("test value should be present");
assert_eq!(transformer.config.hidden_dim, 32);
assert_eq!(transformer.config.num_layers, 1);
assert_eq!(transformer.config.num_heads, 4);
assert_eq!(transformer.config.vocab_size, 32);
}
#[test]
fn test_active_pygmy_load_owned_quantized_model() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let result = OwnedQuantizedModel::from_mapped(&mapped);
assert!(
result.is_ok(),
"Failed to load OwnedQuantizedModel: {:?}",
result.err()
);
let model = result.expect("test value should be present");
assert_eq!(model.config.hidden_dim, 32);
assert_eq!(model.config.num_layers, 1);
assert_eq!(model.config.num_heads, 4);
assert_eq!(model.config.vocab_size, 32);
}
#[test]
fn test_active_pygmy_embed() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
let embedding = model.embed(&[0]);
assert_eq!(embedding.len(), 32); assert!(embedding.iter().all(|&v| v.is_finite()));
}
#[test]
fn test_active_pygmy_kv_cache() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
let config = model.config();
let cache = OwnedQuantizedKVCache::from_config(config, 32);
drop(cache);
}
#[test]
fn test_active_pygmy_forward_cached() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
let config = model.config();
let mut cache = OwnedQuantizedKVCache::from_config(config, 32);
let token_id = 1u32; let position = 0usize;
let result = model.forward_cached(token_id, &mut cache, position);
assert!(result.is_ok(), "forward_cached failed: {:?}", result.err());
let logits = result.expect("test value should be present");
assert_eq!(logits.len(), 32);
assert!(
logits.iter().all(|&v| v.is_finite()),
"Logits contain NaN/Inf: {:?}",
logits
);
}
#[test]
fn test_active_pygmy_multi_token_generation() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
let config = model.config();
let mut cache = OwnedQuantizedKVCache::from_config(config, 32);
let prompt_tokens = [0u32, 1, 2]; let mut all_logits = Vec::new();
for (pos, &token) in prompt_tokens.iter().enumerate() {
let logits = model
.forward_cached(token, &mut cache, pos)
.expect("Prefill forward should succeed");
all_logits.push(logits);
}
for gen_idx in 0..2 {
let pos = prompt_tokens.len() + gen_idx;
let last_logits = all_logits.last().expect("test value should be present");
let next_token = last_logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap_or(std::cmp::Ordering::Equal))
.map_or(0, |(idx, _)| idx as u32);
let logits = model
.forward_cached(next_token, &mut cache, pos)
.expect("Generation forward should succeed");
all_logits.push(logits);
}
assert_eq!(all_logits.len(), 5);
for (i, logits) in all_logits.iter().enumerate() {
assert_eq!(logits.len(), 32, "Logits {} wrong size", i);
assert!(
logits.iter().all(|&v| v.is_finite()),
"Logits {} contain NaN/Inf",
i
);
}
}
#[test]
fn test_active_pygmy_edge_tokens() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
let config = model.config();
let mut cache = OwnedQuantizedKVCache::from_config(config, 32);
let result = model.forward_cached(0, &mut cache, 0);
assert!(result.is_ok(), "Token 0 should work");
let mut cache2 = OwnedQuantizedKVCache::from_config(config, 32);
let result = model.forward_cached(31, &mut cache2, 0); assert!(result.is_ok(), "Token 31 should work");
}
#[test]
fn test_active_pygmy_config() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
let config = model.config();
assert_eq!(config.architecture, "llama");
assert_eq!(config.hidden_dim, 32);
assert_eq!(config.num_layers, 1);
assert_eq!(config.num_heads, 4);
assert_eq!(config.num_kv_heads, 4);
assert_eq!(config.vocab_size, 32);
assert!(
config.intermediate_dim > 0,
"intermediate_dim should be positive"
);
assert!((config.rope_theta - 10000.0).abs() < 1.0);
assert!((config.eps - 1e-5).abs() < 1e-6);
}
#[test]
fn test_active_pygmy_layer_weights() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
assert_eq!(model.layers.len(), 1);
let layer = &model.layers[0];
assert_eq!(layer.attn_norm_weight.len(), 32);
assert!(layer
.attn_norm_weight
.iter()
.all(|&v| (v - 1.0).abs() < 0.01));
}
#[test]
fn test_active_pygmy_output_weights() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
assert_eq!(model.output_norm_weight.len(), 32);
assert_eq!(model.lm_head_weight.in_dim, 32); assert_eq!(model.lm_head_weight.out_dim, 32); }
#[test]
fn test_active_pygmy_cache_isolation() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
let config = model.config();
let mut cache1 = OwnedQuantizedKVCache::from_config(config, 32);
let mut cache2 = OwnedQuantizedKVCache::from_config(config, 32);
let logits1 = model
.forward_cached(1, &mut cache1, 0)
.expect("test value should be present");
let logits2 = model
.forward_cached(1, &mut cache2, 0)
.expect("test value should be present");
assert_eq!(logits1.len(), logits2.len());
for (i, (l1, l2)) in logits1.iter().zip(logits2.iter()).enumerate() {
assert!(
(l1 - l2).abs() < 1e-6,
"Logit {} differs: {} vs {}",
i,
l1,
l2
);
}
}
#[test]
fn test_active_pygmy_cache_accumulation() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
let config = model.config();
let mut cache = OwnedQuantizedKVCache::from_config(config, 32);
let logits0 = model
.forward_cached(5, &mut cache, 0)
.expect("test value should be present");
let logits1 = model
.forward_cached(5, &mut cache, 1)
.expect("test value should be present");
let logits2 = model
.forward_cached(5, &mut cache, 2)
.expect("test value should be present");
let _diff_01: f32 = logits0
.iter()
.zip(&logits1)
.map(|(a, b)| (a - b).abs())
.sum();
let _diff_12: f32 = logits1
.iter()
.zip(&logits2)
.map(|(a, b)| (a - b).abs())
.sum();
for logits in [&logits0, &logits1, &logits2] {
assert!(
logits.iter().all(|&v| v.is_finite()),
"Logits contain NaN/Inf"
);
}
}
#[test]
fn test_active_pygmy_all_tokens_embed() {
let (temp_file, _) = create_pygmy_temp_file();
let mapped = MappedGGUFModel::from_path(temp_file.path()).expect("Should load MappedGGUFModel");
let model = OwnedQuantizedModel::from_mapped(&mapped).expect("Should load model");
for token_id in 0..32u32 {
let embedding = model.embed(&[token_id]);
assert_eq!(
embedding.len(),
32,
"Token {} embedding wrong size",
token_id
);
assert!(
embedding.iter().all(|&v| v.is_finite()),
"Token {} embedding contains NaN/Inf",
token_id
);
}
}
#[test]
fn test_active_pygmy_transformer_layers() {
let gguf_data = build_executable_pygmy_gguf();
let gguf_model = GGUFModel::from_bytes(&gguf_data).expect("Should parse GGUF");
let transformer = QuantizedGGUFTransformer::from_gguf(&gguf_model, &gguf_data)
.expect("Should load transformer");
assert_eq!(transformer.layers.len(), 1);
let layer = &transformer.layers[0];
assert_eq!(layer.attn_norm_weight.len(), 32);
match &layer.qkv_weight {
crate::gguf::quantized::QKVWeights::Separate { q, k, v } => {
assert!(q.num_elements > 0);
assert!(k.num_elements > 0);
assert!(v.num_elements > 0);
},
crate::gguf::quantized::QKVWeights::Fused(fused) => {
assert!(fused.num_elements > 0);
},
}
assert!(layer.ffn_up_weight.num_elements > 0);
assert!(layer.ffn_down_weight.num_elements > 0);
}