#![cfg(feature = "bnb")]
#![allow(
clippy::panic,
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing,
clippy::as_conversions,
clippy::cast_possible_truncation,
clippy::cast_precision_loss,
clippy::similar_names,
clippy::wildcard_enum_match_arm
)]
use std::path::{Path, PathBuf};
use std::time::Instant;
use anamnesis::remember::bnb::{
dequantize_bnb4_double_quant_to_bf16, dequantize_bnb4_to_bf16, dequantize_bnb_int8_to_bf16,
};
use anamnesis::{encode_bnb4, encode_bnb4_double_quant, encode_bnb_int8};
fn read_u32_le(data: &[u8], offset: usize) -> u32 {
let bytes: [u8; 4] = data[offset..offset + 4].try_into().unwrap();
u32::from_le_bytes(bytes)
}
struct Bnb4Fixture {
format_id: u32,
total_elements: usize,
block_size: usize,
nested_offset: f32,
weight_data: Vec<u8>,
absmax_data: Vec<u8>,
quant_map_data: Vec<u8>,
nested_absmax_data: Vec<u8>,
nested_quant_map_data: Vec<u8>,
}
struct BnbInt8Fixture {
out_features: usize,
in_features: usize,
weight_data: Vec<u8>,
scb_data: Vec<u8>,
}
fn parse_bnb4_fixture(data: &[u8]) -> Bnb4Fixture {
let format_id = read_u32_le(data, 0);
let total_elements = read_u32_le(data, 4) as usize;
let block_size = read_u32_le(data, 8) as usize;
let weight_len = read_u32_le(data, 12) as usize;
let absmax_len = read_u32_le(data, 16) as usize;
let quant_map_len = read_u32_le(data, 20) as usize;
let nested_absmax_len = read_u32_le(data, 24) as usize;
let nested_quant_map_len = read_u32_le(data, 28) as usize;
let _expected_len = read_u32_le(data, 32) as usize;
let nested_offset = f32::from_le_bytes(data[36..40].try_into().unwrap());
let header_size = 40;
let mut offset = header_size;
let weight_data = data[offset..offset + weight_len].to_vec();
offset += weight_len;
let absmax_data = data[offset..offset + absmax_len].to_vec();
offset += absmax_len;
let quant_map_data = data[offset..offset + quant_map_len].to_vec();
offset += quant_map_len;
let nested_absmax_data = data[offset..offset + nested_absmax_len].to_vec();
offset += nested_absmax_len;
let nested_quant_map_data = data[offset..offset + nested_quant_map_len].to_vec();
Bnb4Fixture {
format_id,
total_elements,
block_size,
nested_offset,
weight_data,
absmax_data,
quant_map_data,
nested_absmax_data,
nested_quant_map_data,
}
}
fn parse_int8_fixture(data: &[u8]) -> BnbInt8Fixture {
let _format_id = read_u32_le(data, 0);
let out_features = read_u32_le(data, 4) as usize;
let in_features = read_u32_le(data, 8) as usize;
let weight_len = read_u32_le(data, 12) as usize;
let scb_len = read_u32_le(data, 16) as usize;
let _expected_len = read_u32_le(data, 20) as usize;
let header_size = 24;
let mut offset = header_size;
let weight_data = data[offset..offset + weight_len].to_vec();
offset += weight_len;
let scb_data = data[offset..offset + scb_len].to_vec();
BnbInt8Fixture {
out_features,
in_features,
weight_data,
scb_data,
}
}
fn read_pytorch_quantize_us(sidecar_path: &Path) -> Option<f64> {
let bytes = std::fs::read(sidecar_path).ok()?;
let parsed: serde_json::Value = serde_json::from_slice(&bytes).ok()?;
let ns = parsed.get("pytorch_quantize_ns")?.as_u64()?;
Some(ns as f64 / 1000.0)
}
fn sidecar_path_for(fixture_filename: &str) -> PathBuf {
let dir = Path::new(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("fixtures")
.join("bnb_reference");
dir.join(fixture_filename.replace(".bin", ".timing.json"))
}
fn print_runtime_summary(name: &str, anamnesis_us: f64, sidecar: Option<f64>) {
match sidecar {
Some(pytorch_us) => {
let ratio = pytorch_us / anamnesis_us.max(f64::MIN_POSITIVE);
eprintln!(
" {name}: anamnesis encode = {anamnesis_us:.1} \u{00B5}s, \
PyTorch quantize = {pytorch_us:.1} \u{00B5}s ({ratio:.2}x)",
);
}
None => {
eprintln!(
" {name}: anamnesis encode = {anamnesis_us:.1} \u{00B5}s \
(no PyTorch timing sidecar)",
);
}
}
}
fn count_byte_diffs(actual: &[u8], expected: &[u8]) -> usize {
actual
.iter()
.zip(expected.iter())
.filter(|(a, e)| a != e)
.count()
}
fn assert_bytes_equal(actual: &[u8], expected: &[u8], name: &str) {
assert_eq!(
actual.len(),
expected.len(),
"{name}: byte count mismatch ({} vs {})",
actual.len(),
expected.len(),
);
let mut mismatches = 0usize;
for (i, (a, e)) in actual.iter().zip(expected.iter()).enumerate() {
if a != e {
mismatches += 1;
if mismatches <= 5 {
eprintln!(" byte {i}: actual=0x{a:02X}, expected=0x{e:02X}");
}
}
}
assert_eq!(mismatches, 0, "{name}: {mismatches} byte mismatches");
}
fn assert_bf16_equal(actual: &[u8], expected: &[u8], name: &str) {
assert_eq!(
actual.len(),
expected.len(),
"{name}: BF16 byte count mismatch ({} vs {})",
actual.len(),
expected.len(),
);
let mut mismatches = 0usize;
for (i, (a_pair, e_pair)) in actual
.chunks_exact(2)
.zip(expected.chunks_exact(2))
.enumerate()
{
let a_bits = u16::from_le_bytes([a_pair[0], a_pair[1]]);
let e_bits = u16::from_le_bytes([e_pair[0], e_pair[1]]);
let a_is_nan = (a_bits & 0x7F80 == 0x7F80) && (a_bits & 0x007F != 0);
let e_is_nan = (e_bits & 0x7F80 == 0x7F80) && (e_bits & 0x007F != 0);
if a_is_nan && e_is_nan {
continue;
}
if a_bits != e_bits {
mismatches += 1;
if mismatches <= 5 {
eprintln!(" bf16[{i}]: actual=0x{a_bits:04X}, expected=0x{e_bits:04X}");
}
}
}
assert_eq!(
mismatches, 0,
"{name}: {mismatches} BF16 mismatches (decode-equivalence broken)",
);
}
#[derive(Clone, Copy)]
#[allow(dead_code)]
enum ByteContract {
Required,
Diagnostic,
}
fn run_bnb4_encode_cross_validation(
fixture_name: &str,
fixture_filename: &str,
data: &[u8],
byte_contract: ByteContract,
) {
let fixture = parse_bnb4_fixture(data);
assert_eq!(
fixture.format_id, 0,
"this runner only handles plain NF4/FP4 (format_id=0); \
got format_id={} for {fixture_name}",
fixture.format_id,
);
eprintln!(
"{fixture_name}: NF4/FP4 encode, block_size={}, {} elements",
fixture.block_size, fixture.total_elements,
);
let bf16_from_pytorch_bytes = dequantize_bnb4_to_bf16(
&fixture.weight_data,
&fixture.absmax_data,
&fixture.quant_map_data,
fixture.total_elements,
fixture.block_size,
)
.expect("BnB4 decode failed during encode cross-validation");
let start = Instant::now();
let re_encoded = encode_bnb4(
&bf16_from_pytorch_bytes,
&fixture.absmax_data,
&fixture.quant_map_data,
fixture.total_elements,
fixture.block_size,
)
.expect("BnB4 encode failed");
let elapsed = start.elapsed();
let anamnesis_us = elapsed.as_secs_f64() * 1e6;
let bf16_from_re_encoded = dequantize_bnb4_to_bf16(
&re_encoded,
&fixture.absmax_data,
&fixture.quant_map_data,
fixture.total_elements,
fixture.block_size,
)
.expect("BnB4 decode of re-encoded bytes failed");
assert_bf16_equal(
&bf16_from_re_encoded,
&bf16_from_pytorch_bytes,
fixture_name,
);
match byte_contract {
ByteContract::Required => {
assert_bytes_equal(&re_encoded, &fixture.weight_data, fixture_name);
eprintln!(
" {fixture_name}: byte-exact vs PyTorch encoding (0 diffs / {} bytes)",
fixture.weight_data.len(),
);
}
ByteContract::Diagnostic => {
let diffs = count_byte_diffs(&re_encoded, &fixture.weight_data);
eprintln!(
" {fixture_name}: decode-equivalent; {diffs} / {} byte diffs vs PyTorch \
encoding (expected: FP4 quant_map collapses -0 to +0)",
fixture.weight_data.len(),
);
}
}
let sidecar = read_pytorch_quantize_us(&sidecar_path_for(fixture_filename));
print_runtime_summary(fixture_name, anamnesis_us, sidecar);
}
fn run_int8_encode_cross_validation(fixture_name: &str, fixture_filename: &str, data: &[u8]) {
let fixture = parse_int8_fixture(data);
eprintln!(
"{fixture_name}: INT8 encode, {}x{} = {} elements",
fixture.out_features,
fixture.in_features,
fixture.out_features * fixture.in_features,
);
let bf16_from_pytorch_bytes = dequantize_bnb_int8_to_bf16(
&fixture.weight_data,
&fixture.scb_data,
fixture.out_features,
fixture.in_features,
)
.expect("BnB INT8 decode failed during encode cross-validation");
let start = Instant::now();
let re_encoded = encode_bnb_int8(
&bf16_from_pytorch_bytes,
&fixture.scb_data,
fixture.out_features,
fixture.in_features,
)
.expect("BnB INT8 encode failed");
let elapsed = start.elapsed();
let anamnesis_us = elapsed.as_secs_f64() * 1e6;
let bf16_from_re_encoded = dequantize_bnb_int8_to_bf16(
&re_encoded,
&fixture.scb_data,
fixture.out_features,
fixture.in_features,
)
.expect("BnB INT8 decode of re-encoded bytes failed");
assert_bf16_equal(
&bf16_from_re_encoded,
&bf16_from_pytorch_bytes,
fixture_name,
);
assert_bytes_equal(&re_encoded, &fixture.weight_data, fixture_name);
eprintln!(
" {fixture_name}: byte-exact vs PyTorch encoding (0 diffs / {} bytes)",
fixture.weight_data.len(),
);
let sidecar = read_pytorch_quantize_us(&sidecar_path_for(fixture_filename));
print_runtime_summary(fixture_name, anamnesis_us, sidecar);
}
fn run_bnb4_double_quant_encode_cross_validation(
fixture_name: &str,
fixture_filename: &str,
data: &[u8],
) {
let fixture = parse_bnb4_fixture(data);
assert_eq!(
fixture.format_id, 2,
"this runner only handles NF4/FP4 double-quant (format_id=2); \
got format_id={} for {fixture_name}",
fixture.format_id,
);
let absmax_count = fixture.absmax_data.len();
let nested_absmax_count = fixture.nested_absmax_data.len() / 4;
let nested_block_size = if nested_absmax_count > 0 {
absmax_count.div_ceil(nested_absmax_count)
} else {
256
};
eprintln!(
"{fixture_name}: NF4 double-quant encode, block_size={}, nested_block_size={}, \
{} elements",
fixture.block_size, nested_block_size, fixture.total_elements,
);
let bf16_from_pytorch_bytes = dequantize_bnb4_double_quant_to_bf16(
&fixture.weight_data,
&fixture.absmax_data,
&fixture.quant_map_data,
&fixture.nested_absmax_data,
&fixture.nested_quant_map_data,
fixture.nested_offset,
fixture.total_elements,
fixture.block_size,
nested_block_size,
)
.expect("BnB4 double-quant decode failed during encode cross-validation");
let start = Instant::now();
let re_encoded = encode_bnb4_double_quant(
&bf16_from_pytorch_bytes,
&fixture.absmax_data,
&fixture.quant_map_data,
&fixture.nested_absmax_data,
&fixture.nested_quant_map_data,
fixture.nested_offset,
fixture.total_elements,
fixture.block_size,
nested_block_size,
)
.expect("BnB4 double-quant encode failed");
let elapsed = start.elapsed();
let anamnesis_us = elapsed.as_secs_f64() * 1e6;
let bf16_from_re_encoded = dequantize_bnb4_double_quant_to_bf16(
&re_encoded,
&fixture.absmax_data,
&fixture.quant_map_data,
&fixture.nested_absmax_data,
&fixture.nested_quant_map_data,
fixture.nested_offset,
fixture.total_elements,
fixture.block_size,
nested_block_size,
)
.expect("BnB4 double-quant decode of re-encoded bytes failed");
assert_bf16_equal(
&bf16_from_re_encoded,
&bf16_from_pytorch_bytes,
fixture_name,
);
assert_bytes_equal(&re_encoded, &fixture.weight_data, fixture_name);
eprintln!(
" {fixture_name}: byte-exact vs PyTorch encoding (0 diffs / {} bytes)",
fixture.weight_data.len(),
);
let sidecar = read_pytorch_quantize_us(&sidecar_path_for(fixture_filename));
print_runtime_summary(fixture_name, anamnesis_us, sidecar);
}
#[test]
fn cross_validate_encode_llama_1b_nf4() {
let data = include_bytes!("fixtures/bnb_reference/llama_1b_nf4.bin");
run_bnb4_encode_cross_validation(
"Llama-3.2-1B NF4",
"llama_1b_nf4.bin",
data,
ByteContract::Required,
);
}
#[test]
fn cross_validate_encode_llama_1b_fp4() {
let data = include_bytes!("fixtures/bnb_reference/llama_1b_fp4.bin");
run_bnb4_encode_cross_validation(
"Llama-3.2-1B FP4",
"llama_1b_fp4.bin",
data,
ByteContract::Required,
);
}
#[test]
fn cross_validate_encode_llama_1b_int8() {
let data = include_bytes!("fixtures/bnb_reference/llama_1b_int8.bin");
run_int8_encode_cross_validation("Llama-3.2-1B INT8", "llama_1b_int8.bin", data);
}
#[test]
fn cross_validate_encode_llama_1b_nf4_double_quant() {
let data = include_bytes!("fixtures/bnb_reference/llama_1b_nf4_double_quant.bin");
run_bnb4_double_quant_encode_cross_validation(
"Llama-3.2-1B NF4 double-quant",
"llama_1b_nf4_double_quant.bin",
data,
);
}
#[test]
fn cross_validate_encode_qwen2_5_1_5b_nf4_dq() {
let data = include_bytes!("fixtures/bnb_reference/qwen2_5_1_5b_nf4_dq.bin");
run_bnb4_double_quant_encode_cross_validation(
"Qwen2.5-1.5B NF4 double-quant",
"qwen2_5_1_5b_nf4_dq.bin",
data,
);
}
#[test]
fn cross_validate_encode_phi3_5_mini_nf4_dq() {
let data = include_bytes!("fixtures/bnb_reference/phi3_5_mini_nf4_dq.bin");
run_bnb4_double_quant_encode_cross_validation(
"Phi-3.5-mini NF4 double-quant",
"phi3_5_mini_nf4_dq.bin",
data,
);
}
#[test]
fn cross_validate_encode_qwen3_mcqa_fp4() {
let data = include_bytes!("fixtures/bnb_reference/qwen3_mcqa_fp4.bin");
run_bnb4_encode_cross_validation(
"Qwen3 MCQA FP4",
"qwen3_mcqa_fp4.bin",
data,
ByteContract::Required,
);
}