use crate::apr_transformer::{
ActivationStats, AprTransformerConfig, ForwardTrace, LayerActivation,
};
use super::super::{dequant_perrow, dequant_q4k_block, dequant_q6k_block};
#[test]
fn test_dequant_q4k_block_with_dmin() {
let mut block = vec![0u8; 144];
block[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
block[2..4].copy_from_slice(&0x3C00u16.to_le_bytes());
block[4] = 0x01; block[8] = 0x01;
let mut out = vec![0.0f32; 256];
dequant_q4k_block(&block, &mut out);
assert!(
(out[0] - (-1.0)).abs() < 0.01,
"Expected -1.0, got {}",
out[0]
);
}
#[test]
fn test_dequant_q4k_block_nibble_extraction() {
let mut block = vec![0u8; 144];
block[0..2].copy_from_slice(&0x3C00u16.to_le_bytes());
block[2..4].copy_from_slice(&0x0000u16.to_le_bytes());
block[4] = 0x01;
block[16] = 0xAB;
let mut out = vec![0.0f32; 256];
dequant_q4k_block(&block, &mut out);
assert!(
(out[0] - 11.0).abs() < 0.01,
"Expected 11.0 for low nibble, got {}",
out[0]
);
}
#[test]
fn test_dequant_q4k_block_all_ones_qs() {
let mut block = vec![0u8; 144];
block[0..2].copy_from_slice(&0x3800u16.to_le_bytes());
block[2..4].copy_from_slice(&0x0000u16.to_le_bytes());
block[4] = 0x02;
for i in 0..128 {
block[16 + i] = 0xFF;
}
let mut out = vec![0.0f32; 256];
dequant_q4k_block(&block, &mut out);
assert!(
(out[0] - 15.0).abs() < 0.01,
"Expected 15.0, got {}",
out[0]
);
}
#[test]
fn test_dequant_q4k_block_d_negative() {
let mut block = vec![0u8; 144];
block[0..2].copy_from_slice(&0xBC00u16.to_le_bytes());
block[2..4].copy_from_slice(&0x0000u16.to_le_bytes());
block[4] = 0x01; block[16] = 0x01;
let mut out = vec![0.0f32; 256];
dequant_q4k_block(&block, &mut out);
assert!(
(out[0] - (-1.0)).abs() < 0.01,
"Expected -1.0, got {}",
out[0]
);
}
#[test]
fn test_dequant_q6k_block_with_scale_and_ql() {
let mut block = vec![0u8; 210];
block[208..210].copy_from_slice(&0x3800u16.to_le_bytes());
block[192] = 2;
block[0] = 0x11;
block[128] = 0x00;
let mut out = vec![0.0f32; 256];
dequant_q6k_block(&block, &mut out);
assert!(
(out[0] - (-31.0)).abs() < 0.1,
"Expected -31.0, got {}",
out[0]
);
}
#[test]
fn test_dequant_q6k_block_negative_d() {
let mut block = vec![0u8; 210];
block[208..210].copy_from_slice(&0xBC00u16.to_le_bytes());
block[192] = 1;
let mut out = vec![0.0f32; 256];
dequant_q6k_block(&block, &mut out);
assert!((out[0] - 32.0).abs() < 0.1, "Expected 32.0, got {}", out[0]);
}
#[test]
fn test_dequant_q6k_block_high_bits_contribution() {
let mut block = vec![0u8; 210];
block[208..210].copy_from_slice(&0x3C00u16.to_le_bytes());
block[192] = 1;
block[0] = 0x0F;
block[128] = 0x03;
let mut out = vec![0.0f32; 256];
dequant_q6k_block(&block, &mut out);
assert!((out[0] - 31.0).abs() < 0.1, "Expected 31.0, got {}", out[0]);
}
#[test]
fn test_dequant_q6k_block_second_half() {
let mut block = vec![0u8; 210];
block[208..210].copy_from_slice(&0x3C00u16.to_le_bytes());
block[200] = 3;
block[64] = 0x02;
block[160] = 0x00;
let mut out = vec![0.0f32; 256];
dequant_q6k_block(&block, &mut out);
assert!(
(out[128] - (-90.0)).abs() < 0.1,
"Expected -90.0, got {}",
out[128]
);
}
#[test]
fn test_dequant_perrow_multi_block_row() {
let block_bytes = 144;
let block_elems = 256;
let rows = 1;
let cols: usize = 512;
let blocks_per_row = cols.div_ceil(block_elems); let data = vec![0u8; rows * blocks_per_row * block_bytes];
let dims = vec![rows, cols];
let result = dequant_perrow(&data, &dims, block_elems, block_bytes, |_block, out| {
for (i, v) in out.iter_mut().enumerate() {
*v = (i + 1) as f32;
}
});
assert_eq!(result.len(), 512);
assert!((result[0] - 1.0).abs() < 0.001);
assert!((result[256] - 1.0).abs() < 0.001);
}
#[test]
fn test_dequant_perrow_single_row() {
let block_bytes = 210; let block_elems = 256;
let rows = 1;
let cols = 256;
let data = vec![0u8; block_bytes];
let dims = vec![rows, cols];
let result = dequant_perrow(&data, &dims, block_elems, block_bytes, |_block, out| {
for (i, v) in out.iter_mut().enumerate() {
*v = i as f32;
}
});
assert_eq!(result.len(), 256);
assert!((result[0] - 0.0).abs() < 0.001);
assert!((result[255] - 255.0).abs() < 0.001);
}
#[test]
fn test_dequant_perrow_zero_data() {
let block_bytes = 144;
let block_elems = 256;
let rows = 2;
let cols = 256;
let data: Vec<u8> = Vec::new(); let dims = vec![rows, cols];
let result = dequant_perrow(&data, &dims, block_elems, block_bytes, |_block, out| {
for v in out.iter_mut() {
*v = 42.0;
}
});
assert_eq!(result.len(), rows * cols);
for &v in &result {
assert!(
(v - 0.0).abs() < 0.001,
"Expected 0.0 for insufficient data, got {}",
v
);
}
}
#[test]
fn test_dequant_perrow_cols_not_multiple_of_block_elems() {
let block_bytes = 144;
let block_elems = 256;
let rows = 2;
let cols: usize = 100;
let blocks_per_row = cols.div_ceil(block_elems); let data = vec![0u8; rows * blocks_per_row * block_bytes];
let dims = vec![rows, cols];
let result = dequant_perrow(&data, &dims, block_elems, block_bytes, |_block, out| {
for (i, v) in out.iter_mut().enumerate() {
*v = (i + 1) as f32;
}
});
assert_eq!(result.len(), rows * cols);
assert!((result[0] - 1.0).abs() < 0.001);
assert!((result[99] - 100.0).abs() < 0.001);
}
#[test]
fn test_activation_stats_mixed_nan_inf_zero() {
let data = vec![
0.0,
f32::NAN,
f32::INFINITY,
1.0,
f32::NEG_INFINITY,
f32::NAN,
0.0,
2.0,
];
let stats = ActivationStats::from_slice(&data);
assert_eq!(stats.count, 8);
assert_eq!(stats.nan_count, 2);
assert_eq!(stats.inf_count, 2);
assert_eq!(stats.zero_count, 2);
assert!((stats.mean - 0.75).abs() < 0.01);
}
#[test]
fn test_activation_stats_two_elements_variance() {
let data = vec![0.0, 10.0];
let stats = ActivationStats::from_slice(&data);
assert_eq!(stats.count, 2);
assert!((stats.min - 0.0).abs() < 0.001);
assert!((stats.max - 10.0).abs() < 0.001);
assert!((stats.mean - 5.0).abs() < 0.001);
assert!((stats.std_dev - 7.071).abs() < 0.1);
}
#[test]
fn test_activation_stats_all_inf() {
let data = vec![f32::INFINITY, f32::NEG_INFINITY, f32::INFINITY];
let stats = ActivationStats::from_slice(&data);
assert_eq!(stats.count, 3);
assert_eq!(stats.inf_count, 3);
assert_eq!(stats.mean, 0.0); assert_eq!(stats.std_dev, 0.0);
}
#[test]
fn test_activation_stats_negative_only() {
let data = vec![-10.0, -20.0, -30.0];
let stats = ActivationStats::from_slice(&data);
assert!((stats.min - (-30.0)).abs() < 0.001);
assert!((stats.max - (-10.0)).abs() < 0.001);
assert!((stats.mean - (-20.0)).abs() < 0.001);
}
#[test]
fn test_apr_transformer_config_serde_roundtrip() {
let config = AprTransformerConfig {
architecture: "qwen2".to_string(),
hidden_dim: 1536,
num_layers: 28,
num_heads: 12,
num_kv_heads: 2,
vocab_size: 151936,
intermediate_dim: 8960,
context_length: 32768,
rope_theta: 1_000_000.0,
eps: 1e-6,
eos_token_id: None,
..Default::default()
};
let json = serde_json::to_string(&config).expect("serialize failed");
let deserialized: AprTransformerConfig =
serde_json::from_str(&json).expect("deserialize failed");
assert_eq!(deserialized.architecture, "qwen2");
assert_eq!(deserialized.hidden_dim, 1536);
assert_eq!(deserialized.num_layers, 28);
assert_eq!(deserialized.num_heads, 12);
assert_eq!(deserialized.num_kv_heads, 2);
assert_eq!(deserialized.vocab_size, 151936);
assert_eq!(deserialized.intermediate_dim, 8960);
assert_eq!(deserialized.context_length, 32768);
assert!((deserialized.rope_theta - 1_000_000.0).abs() < 1.0);
assert!((deserialized.eps - 1e-6).abs() < 1e-9);
}
#[test]
fn test_apr_transformer_config_small_model() {
let config = AprTransformerConfig {
architecture: "tiny".to_string(),
hidden_dim: 4,
num_layers: 1,
num_heads: 1,
num_kv_heads: 1,
vocab_size: 8,
intermediate_dim: 8,
context_length: 16,
rope_theta: 10000.0,
eps: 1e-5,
eos_token_id: None,
..Default::default()
};
let json = serde_json::to_string(&config).expect("serialize failed");
assert!(json.contains("\"tiny\""));
assert!(json.contains("\"hidden_dim\":4"));
}
#[test]
fn test_from_apr_bytes_too_small() {
use crate::apr_transformer::AprTransformer;
let data = vec![0u8; 32]; let result = AprTransformer::from_apr_bytes(&data);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(
err.contains("too small") || err.contains("64"),
"Expected size error, got: {}",
err
);
}
#[test]
fn test_from_apr_bytes_bad_magic() {
use crate::apr_transformer::AprTransformer;
let mut data = vec![0u8; 128];
data[0] = b'X';
data[1] = b'Y';
data[2] = b'Z';
data[3] = b'0';
let result = AprTransformer::from_apr_bytes(&data);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(
err.contains("magic") || err.contains("Invalid APR") || err.contains("APR"),
"Expected magic error, got: {}",
err
);
}
include!("apr_02.rs");