use candle_core::{DType, Device, Tensor};
use cortex_rust::layers::bit_linear::BitLinear;
use std::time::Instant;
mod common;
#[test]
fn test_u8_preservation() {
let model_path = common::get_test_model_safetensors();
if !model_path.exists() {
eprintln!("โญ๏ธ Skipping: model not found at {:?}", model_path);
return;
}
println!("\n๐งช U8 Preservation Test");
println!("โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ");
let device = Device::Cpu;
let tensors =
candle_core::safetensors::load(model_path, &device).expect("Failed to load safetensors");
let mut u8_count = 0;
let mut f32_count = 0;
let mut other_count = 0;
for (name, tensor) in &tensors {
match tensor.dtype() {
DType::U8 => {
u8_count += 1;
if name.contains("weight_packed") {
println!(" โ
{} โ U8 (preserved!)", name);
}
}
DType::F32 => f32_count += 1,
_ => other_count += 1,
}
}
println!("\n๐ Tensor Stats:");
println!(" - U8: {} tensors", u8_count);
println!(" - F32: {} tensors", f32_count);
println!(" - Other: {} tensors", other_count);
assert!(u8_count > 0, "Expected at least some U8 tensors!");
println!("\nโ
U8 tensors preserved correctly!");
}
#[test]
fn benchmark_single_layer_load() {
let model_path = common::get_test_model_safetensors();
if !model_path.exists() {
eprintln!("โญ๏ธ Skipping: model not found at {:?}", model_path);
return;
}
println!("\nโฑ๏ธ Single Layer Load Benchmark");
println!("โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ");
let device = Device::Cpu;
let iterations = 5;
let model_path_ref = model_path.as_path();
let start = Instant::now();
for _ in 0..iterations {
let tensors = candle_core::safetensors::load(model_path_ref, &device).unwrap();
let layer_prefix = "model.layers.0.mlp.gate_proj";
let packed_key = format!("{}.weight_packed", layer_prefix);
let scales_key = format!("{}.scales", layer_prefix);
if let (Some(packed), Some(scales)) = (tensors.get(&packed_key), tensors.get(&scales_key)) {
let dtype = packed.dtype();
let _ = BitLinear::from_packed_tensors(packed, scales, &device).unwrap();
if dtype != DType::U8 {
println!(" โ ๏ธ Warning: dtype was {:?}, not U8", dtype);
}
}
}
let direct_time = start.elapsed();
let start = Instant::now();
for _ in 0..iterations {
let vb = unsafe {
candle_nn::VarBuilder::from_mmaped_safetensors(&[model_path_ref], DType::F32, &device)
.unwrap()
};
let layer_vb = vb.pp("model.layers.0.mlp.gate_proj");
if let (Ok(packed), Ok(scales)) = (
layer_vb.get(&[2048usize, 1376 / 4], "weight_packed"),
layer_vb.get(&[1usize], "scales"),
) {
let dtype = packed.dtype();
let _ = BitLinear::from_packed_tensors(&packed, &scales, &device);
if dtype != DType::F32 {
println!(" Unexpected: VarBuilder returned {:?}", dtype);
}
}
}
let varbuilder_time = start.elapsed();
println!("\n๐ Results ({} iterations):", iterations);
println!(" - Direct (U8): {:?}", direct_time);
println!(" - VarBuilder (F32): {:?}", varbuilder_time);
let speedup = varbuilder_time.as_secs_f64() / direct_time.as_secs_f64();
println!(" - Speedup: {:.2}x", speedup);
if direct_time < varbuilder_time {
println!("\nโ
Direct load is faster!");
}
}
#[test]
fn test_direct_load_forward() {
let model_path = common::get_test_model_safetensors();
if !model_path.exists() {
eprintln!("โญ๏ธ Skipping: model not found at {:?}", model_path);
return;
}
println!("\n๐งช Forward Pass Test (Direct Load)");
println!("โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ");
let device = Device::Cpu;
let tensors = candle_core::safetensors::load(model_path, &device).unwrap();
let layer_prefix = "model.layers.0.mlp.gate_proj";
let packed_key = format!("{}.weight_packed", layer_prefix);
let scales_key = format!("{}.scales", layer_prefix);
let packed = tensors.get(&packed_key).unwrap();
let scales = tensors.get(&scales_key).unwrap();
println!(" packed dtype: {:?}", packed.dtype());
println!(" packed shape: {:?}", packed.dims());
let layer = BitLinear::from_packed_tensors(packed, scales, &device).unwrap();
let in_dim = layer.in_features;
let x = Tensor::randn(0.0f32, 1.0, (1, in_dim), &device).unwrap();
let output = layer.forward(&x).unwrap();
println!(" output shape: {:?}", output.dims());
let output_vec = output.flatten_all().unwrap().to_vec1::<f32>().unwrap();
let has_nan = output_vec.iter().any(|v| v.is_nan());
let has_inf = output_vec.iter().any(|v| v.is_infinite());
assert!(!has_nan, "Output contains NaN!");
assert!(!has_inf, "Output contains Inf!");
println!("\nโ
Forward pass successful!");
}