#![allow(clippy::cast_possible_truncation)]
use crate::error::InferenceError;
#[repr(C)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Q4Block {
pub scale: u16,
pub bias: u16,
pub packed: [u8; 16],
}
const _: () = assert!(std::mem::size_of::<Q4Block>() == 20);
pub const Q4_BLOCK_BYTES: usize = std::mem::size_of::<Q4Block>();
pub(crate) const Q4_BLOCK_WEIGHTS: usize = 32;
#[derive(Debug, Clone)]
pub struct Q4Tensor {
pub blocks: Vec<Q4Block>,
pub shape: Vec<usize>,
pub original_len: usize,
}
#[inline]
pub(crate) fn q4_f32_to_f16(x: f32) -> u16 {
crate::weights::half_bits::f32_to_f16_bits(x)
}
#[inline]
pub(crate) fn q4_f32_to_finite_f16(x: f32) -> Result<u16, u16> {
crate::weights::half_bits::f32_to_finite_f16_bits(x)
}
#[inline]
pub(crate) fn q4_f16_to_f32(bits: u16) -> f32 {
crate::weights::half_bits::f16_bits_to_f32(bits)
}
#[inline]
fn bf16_to_f32(v: u16) -> f32 {
crate::weights::half_bits::bf16_bits_to_f32(v)
}
#[inline]
fn q4_scale_survives_f16(scale: f32) -> bool {
let serialized = q4_f16_to_f32(q4_f32_to_f16(scale));
serialized.is_finite() && serialized > 0.0
}
fn q4_metadata_bits(scale: f32, bias: f32) -> Result<(u16, u16), InferenceError> {
let scale_bits = q4_f32_to_f16(scale);
let bias_bits = q4_f32_to_f16(bias);
let serialized_bias = q4_f16_to_f32(bias_bits);
if !q4_scale_survives_f16(scale) {
return Err(InferenceError::InvalidInput(format!(
"Q4 scale {scale} is not representable as a finite, strictly positive f16 value"
)));
}
if !serialized_bias.is_finite() {
return Err(InferenceError::InvalidInput(format!(
"Q4 bias {bias} is not representable as a finite f16 value"
)));
}
Ok((scale_bits, bias_bits))
}
#[inline]
fn degenerate_safe_scale(candidate: f32) -> f32 {
if candidate < 1.0 && !q4_scale_survives_f16(candidate) {
1.0f32
} else {
candidate
}
}
#[inline]
fn quantize_block_with_mode_len(
vals: &[f32; 32],
valid_len: usize,
symmetric: bool,
) -> Result<Q4Block, InferenceError> {
if !(1..=32).contains(&valid_len) {
return Err(InferenceError::InvalidInput(format!(
"Q4 weight block valid_len {valid_len} must be in 1..=32"
)));
}
for (i, &v) in vals.iter().enumerate() {
if !v.is_finite() {
return Err(InferenceError::InvalidInput(format!(
"Q4 weight block element {i} contains a non-finite value ({v}); \
source weights must be finite"
)));
}
}
if symmetric {
let abs_max = vals.iter().map(|x| x.abs()).fold(0.0f32, f32::max);
let scale = degenerate_safe_scale(abs_max / 7.0);
let inv_scale = 1.0 / scale;
let bias = -8.0 * scale;
let (scale_bits, bias_bits) = q4_metadata_bits(scale, bias)?;
let mut packed = [0u8; 16];
for b in 0..16 {
let q0 = ((vals[2 * b] * inv_scale).round() + 8.0).clamp(0.0, 15.0) as u8;
let q1 = ((vals[2 * b + 1] * inv_scale).round() + 8.0).clamp(0.0, 15.0) as u8;
packed[b] = (q1 << 4) | (q0 & 0x0f);
}
Ok(Q4Block {
scale: scale_bits,
bias: bias_bits,
packed,
})
} else {
let real = &vals[..valid_len];
let min_val = real.iter().copied().fold(f32::INFINITY, f32::min);
let max_val = real.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let range = max_val - min_val;
let scale = degenerate_safe_scale(range / 15.0);
let (scale_bits, bias_bits) = q4_metadata_bits(scale, min_val)?;
let inv_scale = 1.0 / scale;
let mut packed = [0u8; 16];
for b in 0..16 {
let q0 = (((vals[2 * b] - min_val) * inv_scale).round()).clamp(0.0, 15.0) as u8;
let q1 = (((vals[2 * b + 1] - min_val) * inv_scale).round()).clamp(0.0, 15.0) as u8;
packed[b] = (q1 << 4) | (q0 & 0x0f);
}
Ok(Q4Block {
scale: scale_bits,
bias: bias_bits,
packed,
})
}
}
pub fn quantize_row_q4_0(src: &[f32]) -> Result<Vec<u8>, InferenceError> {
let n_blocks = src.len().div_ceil(Q4_BLOCK_WEIGHTS);
let mut out = Vec::with_capacity(n_blocks * 20);
for chunk in src.chunks(32) {
let mut vals = [0.0f32; 32];
vals[..chunk.len()].copy_from_slice(chunk);
let block = quantize_block_with_mode_len(&vals, chunk.len(), false)?;
let bytes: &[u8; 20] = unsafe { &*std::ptr::from_ref(&block).cast() };
out.extend_from_slice(bytes);
}
Ok(out)
}
pub fn dequantize_row_q4_0(data: &[u8], n_weights: usize) -> Vec<f32> {
let mut out = Vec::with_capacity(n_weights);
for chunk in data.chunks_exact(20) {
let scale = q4_f16_to_f32(u16::from_ne_bytes([chunk[0], chunk[1]]));
let bias = q4_f16_to_f32(u16::from_ne_bytes([chunk[2], chunk[3]]));
for b in 0..16 {
let byte_val = chunk[4 + b];
out.push((byte_val & 0x0f) as f32 * scale + bias);
out.push((byte_val >> 4) as f32 * scale + bias);
}
}
out.truncate(n_weights);
out
}
pub fn quantize_tensor_q4_0(
src: &[f32],
rows: usize,
cols: usize,
) -> Result<Vec<u8>, InferenceError> {
let expected_len = rows.checked_mul(cols).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"Q4 tensor shape [{rows}, {cols}] overflows usize element count"
))
})?;
if src.len() != expected_len {
return Err(InferenceError::InvalidInput(format!(
"Q4 tensor source length {} does not match shape [{rows}, {cols}] \
(expected {expected_len})",
src.len()
)));
}
let blocks_per_row = cols.div_ceil(Q4_BLOCK_WEIGHTS);
let mut out = Vec::with_capacity(rows * blocks_per_row * 20);
for row_idx in 0..rows {
let row = &src[row_idx * cols..(row_idx + 1) * cols];
out.extend_from_slice(&quantize_row_q4_0(row)?);
}
Ok(out)
}
fn validate_shape_matches_data_len(shape: &[usize], data_len: usize) -> Result<(), InferenceError> {
let numel = shape
.iter()
.try_fold(1_usize, |acc, &d| acc.checked_mul(d))
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"Q4 tensor shape product overflows usize: shape={shape:?}"
))
})?;
if numel != data_len {
return Err(InferenceError::InvalidInput(format!(
"Q4 tensor shape product {numel} (shape={shape:?}) must equal data length {data_len}"
)));
}
Ok(())
}
pub fn quantize_bf16_to_q4(data: &[u16], shape: &[usize]) -> Result<Q4Tensor, InferenceError> {
validate_shape_matches_data_len(shape, data.len())?;
let original_len = data.len();
let n_blocks = original_len.div_ceil(Q4_BLOCK_WEIGHTS);
let mut blocks = Vec::with_capacity(n_blocks);
for chunk in data.chunks(32) {
let mut vals = [0.0f32; 32];
for (i, &v) in chunk.iter().enumerate() {
vals[i] = bf16_to_f32(v);
}
blocks.push(quantize_block_with_mode_len(&vals, chunk.len(), false)?);
}
Ok(Q4Tensor {
blocks,
shape: shape.to_vec(),
original_len,
})
}
pub fn quantize_f32_to_q4(data: &[f32], shape: &[usize]) -> Result<Q4Tensor, InferenceError> {
validate_shape_matches_data_len(shape, data.len())?;
let original_len = data.len();
let n_blocks = original_len.div_ceil(Q4_BLOCK_WEIGHTS);
let mut blocks = Vec::with_capacity(n_blocks);
for chunk in data.chunks(32) {
let mut vals = [0.0f32; 32];
vals[..chunk.len()].copy_from_slice(chunk);
blocks.push(quantize_block_with_mode_len(&vals, chunk.len(), false)?);
}
Ok(Q4Tensor {
blocks,
shape: shape.to_vec(),
original_len,
})
}
pub fn quantize_f64_to_q4(data: &[f64], shape: &[usize]) -> Result<Q4Tensor, InferenceError> {
quantize_f64_to_q4_mode(data, shape, true) }
pub fn quantize_f64_to_q4_mode(
data: &[f64],
shape: &[usize],
symmetric: bool,
) -> Result<Q4Tensor, InferenceError> {
validate_shape_matches_data_len(shape, data.len())?;
let original_len = data.len();
let n_blocks = original_len.div_ceil(Q4_BLOCK_WEIGHTS);
let mut blocks = Vec::with_capacity(n_blocks);
for chunk in data.chunks(32) {
let mut vals = [0.0f32; 32];
for (i, &v) in chunk.iter().enumerate() {
vals[i] = v as f32;
}
blocks.push(quantize_block_with_mode_len(&vals, chunk.len(), symmetric)?);
}
Ok(Q4Tensor {
blocks,
shape: shape.to_vec(),
original_len,
})
}
pub fn dequantize_q4_to_f32(tensor: &Q4Tensor) -> Vec<f32> {
let mut out = Vec::with_capacity(tensor.original_len);
for block in &tensor.blocks {
let scale = q4_f16_to_f32(block.scale);
let bias = q4_f16_to_f32(block.bias);
for b in 0..16 {
let byte_val = block.packed[b];
out.push((byte_val & 0x0f) as f32 * scale + bias);
out.push((byte_val >> 4) as f32 * scale + bias);
}
}
out.truncate(tensor.original_len);
out
}
pub fn stream_quantize_shard(
bf16_bytes: &[u8],
) -> Result<Vec<Q4Block>, Box<dyn std::error::Error>> {
if !bf16_bytes.len().is_multiple_of(2) {
return Err("bf16_bytes length must be even (2 bytes per BF16 value)".into());
}
let n = bf16_bytes.len() / 2;
let n_blocks = n.div_ceil(Q4_BLOCK_WEIGHTS);
let mut blocks = Vec::with_capacity(n_blocks);
for i in (0..bf16_bytes.len()).step_by(64) {
let end = (i + 64).min(bf16_bytes.len());
let chunk = &bf16_bytes[i..end];
let mut vals = [0.0f32; 32];
for (j, pair) in chunk.chunks_exact(2).enumerate() {
let v = u16::from_ne_bytes([pair[0], pair[1]]);
vals[j] = bf16_to_f32(v);
}
let valid_len = chunk.len() / 2;
blocks.push(
quantize_block_with_mode_len(&vals, valid_len, false)
.map_err(|e| Box::new(e) as Box<dyn std::error::Error>)?,
);
}
Ok(blocks)
}
pub fn save_q4_file(path: &std::path::Path, tensor: &Q4Tensor) -> std::io::Result<()> {
use std::io::Write;
let mut f = std::fs::File::create(path)?;
f.write_all(b"KHQ4")?;
f.write_all(&2u32.to_le_bytes())?;
f.write_all(&(tensor.shape.len() as u32).to_le_bytes())?;
for &dim in &tensor.shape {
f.write_all(&(dim as u64).to_le_bytes())?;
}
f.write_all(&(tensor.original_len as u64).to_le_bytes())?;
let block_bytes: &[u8] = unsafe {
std::slice::from_raw_parts(
tensor.blocks.as_ptr().cast::<u8>(),
tensor.blocks.len() * 20,
)
};
f.write_all(block_bytes)
}
pub struct Q4FileHeader {
pub shape: Vec<usize>,
pub original_len: usize,
pub payload_offset: u64,
}
fn usize_from_u64(value: u64, what: &str) -> Result<usize, Box<dyn std::error::Error>> {
usize::try_from(value).map_err(|_| format!("{what}: value {value} exceeds usize").into())
}
fn checked_alloc_bytes(
count: usize,
elem_size: usize,
file_len: u64,
what: &str,
) -> Result<usize, Box<dyn std::error::Error>> {
let bytes = count
.checked_mul(elem_size)
.ok_or_else(|| format!("{what}: element count {count} × {elem_size} overflows usize"))?;
if bytes as u64 > file_len {
return Err(format!(
"{what}: header claims {bytes} bytes but file is only {file_len} bytes"
)
.into());
}
Ok(bytes)
}
pub fn read_q4_header(
file: &mut std::fs::File,
) -> Result<Q4FileHeader, Box<dyn std::error::Error>> {
use std::io::{Read, Seek, SeekFrom};
let file_len = file.metadata()?.len();
let (shape, original_len, payload_offset) = {
let mut f = std::io::BufReader::new(&mut *file);
let mut magic = [0u8; 4];
f.read_exact(&mut magic)?;
if &magic != b"KHQ4" {
return Err("invalid magic: not a .q4 file".into());
}
let mut b4 = [0u8; 4];
f.read_exact(&mut b4)?;
let ver = u32::from_le_bytes(b4);
if ver == 1 {
return Err("legacy .q4 file (v1 symmetric format) — re-quantize with current quantize_q4 to produce v2 asymmetric blocks".into());
}
if ver != 2 {
return Err(format!("unsupported .q4 file version: {ver}").into());
}
f.read_exact(&mut b4)?;
let ndim = u32::from_le_bytes(b4) as usize;
let shape_bytes = checked_alloc_bytes(ndim, 8, file_len, "shape dims")?;
let mut shape = Vec::with_capacity(ndim);
let mut b8 = [0u8; 8];
for index in 0..ndim {
f.read_exact(&mut b8)?;
shape.push(usize_from_u64(
u64::from_le_bytes(b8),
&format!("shape dimension {index}"),
)?);
}
f.read_exact(&mut b8)?;
let original_len = usize_from_u64(u64::from_le_bytes(b8), "original_len")?;
let payload_offset = 20u64
.checked_add(shape_bytes as u64)
.ok_or("Q4 payload offset overflows u64")?;
(shape, original_len, payload_offset)
};
file.seek(SeekFrom::Start(payload_offset))?;
Ok(Q4FileHeader {
shape,
original_len,
payload_offset,
})
}
pub(crate) fn validate_q4_header_payload_bounds(
header: &Q4FileHeader,
file_len: u64,
path: &std::path::Path,
) -> Result<(), Box<dyn std::error::Error>> {
let payload_bytes = header
.original_len
.div_ceil(Q4_BLOCK_WEIGHTS)
.checked_mul(Q4_BLOCK_BYTES)
.ok_or("Q4 block payload byte count overflows usize")? as u64;
let required_len = header
.payload_offset
.checked_add(payload_bytes)
.ok_or("Q4 payload end offset overflows u64")?;
if file_len < required_len {
return Err(format!(
"{}: file truncated below Q4 block payload ({file_len} bytes < required {required_len})",
path.display()
)
.into());
}
if file_len > required_len {
return Err(format!(
"{}: file has trailing bytes after Q4 block payload ({file_len} bytes > expected \
{required_len})",
path.display()
)
.into());
}
let source = path.display().to_string();
crate::weights::ingress::validate_ingested_tensor(
crate::weights::ingress::IngestedTensor::native_q4(
&source,
"native Q4 tensor",
&header.shape,
header.original_len,
header.original_len.div_ceil(Q4_BLOCK_WEIGHTS),
),
)?;
Ok(())
}
pub(crate) fn validate_q4_file(
file: &mut std::fs::File,
path: &std::path::Path,
expected_shape: Option<&[usize]>,
) -> Result<Q4FileHeader, Box<dyn std::error::Error>> {
use std::io::{Seek, SeekFrom};
let header = read_q4_header(file)?;
let file_len = file.metadata()?.len();
validate_q4_header_payload_bounds(&header, file_len, path)?;
if let Some(expected_shape) = expected_shape {
let source = path.display().to_string();
let tensor_name = path
.file_name()
.and_then(std::ffi::OsStr::to_str)
.unwrap_or("native Q4 tensor");
let block_count = header.original_len.div_ceil(Q4_BLOCK_WEIGHTS);
crate::weights::ingress::validate_ingested_tensor(
crate::weights::ingress::IngestedTensor::native_q4(
&source,
tensor_name,
&header.shape,
header.original_len,
block_count,
)
.with_expected_shape(expected_shape),
)?;
}
file.seek(SeekFrom::Start(header.payload_offset))?;
Ok(header)
}
pub(crate) fn validate_q4_block_metadata(
source: &str,
tensor_name: &str,
index: usize,
scale_bits: u16,
bias_bits: u16,
) -> Result<(), InferenceError> {
crate::weights::ingress::validate_ingested_tensor(
crate::weights::ingress::IngestedTensor::native_q4_block(
source,
tensor_name,
index,
scale_bits,
bias_bits,
),
)
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn validate_q4_block_metadata_scan(
source: &str,
tensor_name: &str,
payload: &[u8],
) -> Result<(), InferenceError> {
for (index, chunk) in payload.chunks_exact(Q4_BLOCK_BYTES).enumerate() {
let scale_bits = u16::from_ne_bytes([chunk[0], chunk[1]]);
let bias_bits = u16::from_ne_bytes([chunk[2], chunk[3]]);
validate_q4_block_metadata(source, tensor_name, index, scale_bits, bias_bits)?;
}
Ok(())
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) enum Q4BlockCheck<'a> {
Now { tensor_name: &'a str },
InCallerTraversal { traversal: &'static str },
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) struct Q4BlocksChecked(());
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn open_and_mmap_q4_file(
path: &std::path::Path,
expected_shape: Option<&[usize]>,
check: Q4BlockCheck<'_>,
) -> Result<(Q4FileHeader, memmap2::Mmap, Option<Q4BlocksChecked>), String> {
let mut file =
std::fs::File::open(path).map_err(|e| format!("failed to open {}: {e}", path.display()))?;
let header = validate_q4_file(&mut file, path, expected_shape)
.map_err(|e| format!("failed to validate Q4 payload {}: {e}", path.display()))?;
let mmap = unsafe { memmap2::MmapOptions::new().map(&file) }
.map_err(|e| format!("failed to mmap {}: {e}", path.display()))?;
let payload = mmap.get(header.payload_offset as usize..).ok_or_else(|| {
format!(
"{}: payload_offset {} beyond mapped length {}",
path.display(),
header.payload_offset,
mmap.len()
)
})?;
let checked = match check {
Q4BlockCheck::Now { tensor_name } => {
let source = path.display().to_string();
validate_q4_block_metadata_scan(&source, tensor_name, payload)
.map_err(|e| e.to_string())?;
Some(Q4BlocksChecked(()))
}
Q4BlockCheck::InCallerTraversal { traversal } => {
debug_assert!(
!traversal.is_empty(),
"a deferred block check must name the traversal that discharges it"
);
None
}
};
Ok((header, mmap, checked))
}
pub fn load_q4_file(path: &std::path::Path) -> Result<Q4Tensor, Box<dyn std::error::Error>> {
let f = std::fs::File::open(path)?;
load_q4_from_open_file(f, path, None)
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn load_q4_file_checked(
path: &std::path::Path,
expected_shape: &[usize],
) -> Result<Q4Tensor, Box<dyn std::error::Error>> {
let f = std::fs::File::open(path)?;
load_q4_from_open_file(f, path, Some(expected_shape))
}
pub(crate) fn load_q4_from_open_file(
mut f: std::fs::File,
path: &std::path::Path,
expected_shape: Option<&[usize]>,
) -> Result<Q4Tensor, Box<dyn std::error::Error>> {
use std::io::Read;
let file_len = f.metadata()?.len();
let header = validate_q4_file(&mut f, path, expected_shape)?;
let n_blocks = header.original_len.div_ceil(Q4_BLOCK_WEIGHTS);
let raw_len = checked_alloc_bytes(n_blocks, Q4_BLOCK_BYTES, file_len, "block payload")?;
let mut raw = vec![0u8; raw_len];
f.read_exact(&mut raw)?;
let source = path.display().to_string();
let tensor_name = path
.file_name()
.and_then(std::ffi::OsStr::to_str)
.unwrap_or("native Q4 tensor");
let mut blocks: Vec<Q4Block> = Vec::with_capacity(n_blocks);
for (index, c) in raw.chunks_exact(Q4_BLOCK_BYTES).enumerate() {
let scale = u16::from_ne_bytes([c[0], c[1]]);
let bias = u16::from_ne_bytes([c[2], c[3]]);
validate_q4_block_metadata(&source, tensor_name, index, scale, bias)?;
let mut packed = [0u8; 16];
packed.copy_from_slice(&c[4..20]);
blocks.push(Q4Block {
scale,
bias,
packed,
});
}
Ok(Q4Tensor {
blocks,
shape: header.shape,
original_len: header.original_len,
})
}
fn read_f16_header(
f: &mut std::fs::File,
display_path: &str,
file_len: u64,
) -> Result<(Vec<usize>, usize), Box<dyn std::error::Error>> {
use std::io::Read;
let mut magic = [0u8; 4];
f.read_exact(&mut magic)?;
if &magic != b"KHF1" {
return Err(
format!("invalid magic at {display_path}: expected KHF1, got {magic:?}").into(),
);
}
let mut b4 = [0u8; 4];
f.read_exact(&mut b4)?;
if u32::from_le_bytes(b4) != 1 {
return Err("unsupported .f16 file version".into());
}
f.read_exact(&mut b4)?;
let ndim = u32::from_le_bytes(b4) as usize;
checked_alloc_bytes(ndim, 8, file_len, "shape dims")?;
let mut shape = Vec::with_capacity(ndim);
let mut b8 = [0u8; 8];
for index in 0..ndim {
f.read_exact(&mut b8)?;
shape.push(usize_from_u64(
u64::from_le_bytes(b8),
&format!("shape dimension {index}"),
)?);
}
f.read_exact(&mut b8)?;
let numel = usize_from_u64(u64::from_le_bytes(b8), "numel")?;
let shape_product = shape
.iter()
.try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
.ok_or("shape dims overflow usize")?;
if shape_product != numel {
return Err(format!(
"{display_path}: shape product {shape_product} (shape={shape:?}) != numel {numel}"
)
.into());
}
Ok((shape, numel))
}
#[derive(Debug)]
pub enum F16LoadError {
ShapeMismatch {
declared: Vec<usize>,
},
Other(Box<dyn std::error::Error>),
}
impl std::fmt::Display for F16LoadError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::ShapeMismatch { declared } => {
write!(f, "f16 header declares shape {declared:?}")
}
Self::Other(_) => write!(f, "f16 tensor could not be read"),
}
}
}
impl std::error::Error for F16LoadError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Other(e) => Some(e.as_ref()),
Self::ShapeMismatch { .. } => None,
}
}
}
pub fn load_f16_tensor_file_expecting(
path: &std::path::Path,
expected: &[usize],
) -> Result<(Vec<f32>, Vec<usize>), F16LoadError> {
let f = std::fs::File::open(path).map_err(|e| F16LoadError::Other(Box::new(e)))?;
load_f16_tensor_from_open_file_expecting(f, &path.display().to_string(), expected)
}
pub(crate) fn load_f16_tensor_from_open_file_expecting(
mut f: std::fs::File,
display_path: &str,
expected: &[usize],
) -> Result<(Vec<f32>, Vec<usize>), F16LoadError> {
let file_len = f
.metadata()
.map_err(|e| F16LoadError::Other(Box::new(e)))?
.len();
let (shape, numel) =
read_f16_header(&mut f, display_path, file_len).map_err(F16LoadError::Other)?;
if shape != expected {
return Err(F16LoadError::ShapeMismatch { declared: shape });
}
read_f16_payload(&mut f, shape, numel, file_len, display_path, Some(expected))
.map_err(F16LoadError::Other)
}
pub fn load_f16_tensor_file(
path: &std::path::Path,
) -> Result<(Vec<f32>, Vec<usize>), Box<dyn std::error::Error>> {
let f = std::fs::File::open(path)?;
load_f16_tensor_from_open_file(f, &path.display().to_string(), None)
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn load_f16_tensor_file_checked(
path: &std::path::Path,
expected_shape: &[usize],
) -> Result<(Vec<f32>, Vec<usize>), Box<dyn std::error::Error>> {
let f = std::fs::File::open(path)?;
load_f16_tensor_from_open_file(f, &path.display().to_string(), Some(expected_shape))
}
pub(crate) fn load_f16_tensor_from_open_file(
mut f: std::fs::File,
display_path: &str,
expected_shape: Option<&[usize]>,
) -> Result<(Vec<f32>, Vec<usize>), Box<dyn std::error::Error>> {
let file_len = f.metadata()?.len();
let (shape, numel) = read_f16_header(&mut f, display_path, file_len)?;
read_f16_payload(&mut f, shape, numel, file_len, display_path, expected_shape)
}
fn read_f16_payload(
f: &mut std::fs::File,
shape: Vec<usize>,
numel: usize,
file_len: u64,
display_path: &str,
expected_shape: Option<&[usize]>,
) -> Result<(Vec<f32>, Vec<usize>), Box<dyn std::error::Error>> {
use std::io::Read;
let raw_len = checked_alloc_bytes(numel, 2, file_len, "f16 data")?;
let shape_bytes = shape
.len()
.checked_mul(8)
.ok_or("KHF1 shape byte count overflows usize")?;
let payload_offset = 20u64
.checked_add(shape_bytes as u64)
.ok_or("KHF1 payload offset overflows u64")?;
let required_len = payload_offset
.checked_add(raw_len as u64)
.ok_or("KHF1 payload end offset overflows u64")?;
if file_len < required_len {
return Err(format!(
"{display_path}: file truncated below KHF1 payload ({file_len} bytes < required \
{required_len})"
)
.into());
}
if file_len > required_len {
return Err(format!(
"{display_path}: file has trailing bytes after KHF1 payload ({file_len} bytes > \
expected {required_len})"
)
.into());
}
let mut raw = vec![0u8; raw_len];
f.read_exact(&mut raw)?;
let values: Vec<f32> = raw
.chunks_exact(2)
.map(|c| {
let bits = u16::from_le_bytes([c[0], c[1]]);
q4_f16_to_f32(bits)
})
.collect();
let tensor_name = std::path::Path::new(display_path)
.file_name()
.and_then(std::ffi::OsStr::to_str)
.unwrap_or("native KHF1 tensor");
let tensor = crate::weights::ingress::IngestedTensor::decoded_f32(
display_path,
tensor_name,
&shape,
"KHF1/F16",
&values,
);
let tensor = if let Some(expected_shape) = expected_shape {
tensor.with_expected_shape(expected_shape)
} else {
tensor
};
crate::weights::ingress::validate_ingested_tensor(tensor)?;
Ok((values, shape))
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) const MAX_Q4_MERGE_PAYLOAD_LEN: u64 = 1 << 31;
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn read_q4_payload_bounded(
path: &std::path::Path,
max_len: u64,
) -> Result<(Q4FileHeader, Vec<u8>), Box<dyn std::error::Error>> {
use std::io::{Read, Seek, SeekFrom};
let mut file = std::fs::File::open(path)?;
let header = read_q4_header(&mut file)?;
let file_len = file.metadata()?.len();
if file_len < header.payload_offset {
return Err(format!(
"{}: file truncated below header ({file_len} bytes < payload_offset {})",
path.display(),
header.payload_offset
)
.into());
}
let payload_len = file_len - header.payload_offset;
if payload_len > max_len {
return Err(format!(
"{}: payload too large: {payload_len} bytes exceeds cap of {max_len} bytes",
path.display()
)
.into());
}
file.seek(SeekFrom::Start(0))?;
let header = validate_q4_file(&mut file, path, None)?;
let mut buf = Vec::new();
file.take(max_len.saturating_add(1)).read_to_end(&mut buf)?;
if buf.len() as u64 > max_len {
return Err(format!(
"{}: payload too large: read exceeds cap of {max_len} bytes",
path.display()
)
.into());
}
Ok((header, buf))
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn q4_sha256_hex(bytes: &[u8]) -> String {
use sha2::{Digest, Sha256};
let mut hasher = Sha256::new();
hasher.update(bytes);
let digest = hasher.finalize();
let mut hex = String::with_capacity(digest.len() * 2);
for byte in digest.as_slice() {
use std::fmt::Write as _;
let _ = write!(&mut hex, "{byte:02x}");
}
hex
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn merged_qkvz_expected_size(qkv_file_len: u64, z_file_len: u64) -> Result<u64, String> {
const HEADER_LEN: u64 = 36;
let qkv_payload_len = qkv_file_len.checked_sub(HEADER_LEN).ok_or_else(|| {
format!("qkv source file too small: {qkv_file_len} bytes < {HEADER_LEN}-byte header")
})?;
let z_payload_len = z_file_len.checked_sub(HEADER_LEN).ok_or_else(|| {
format!("z source file too small: {z_file_len} bytes < {HEADER_LEN}-byte header")
})?;
Ok(HEADER_LEN + qkv_payload_len + z_payload_len)
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn merged_qkvz_source_fingerprint(
qkv_path: &std::path::Path,
z_path: &std::path::Path,
) -> Result<String, Box<dyn std::error::Error>> {
let (_qkv_hdr, mut qkv_payload) = read_q4_payload_bounded(qkv_path, MAX_Q4_MERGE_PAYLOAD_LEN)?;
let (_z_hdr, z_payload) = read_q4_payload_bounded(z_path, MAX_Q4_MERGE_PAYLOAD_LEN)?;
qkv_payload.extend_from_slice(&z_payload);
Ok(q4_sha256_hex(&qkv_payload))
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn merged_qkvz_file_fingerprint(
merged_path: &std::path::Path,
) -> Result<String, Box<dyn std::error::Error>> {
let (_hdr, payload) = read_q4_payload_bounded(merged_path, MAX_Q4_MERGE_PAYLOAD_LEN)?;
Ok(q4_sha256_hex(&payload))
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn merged_qkvz_cache_is_valid(
merged_path: &std::path::Path,
expected_size: u64,
qkv_path: &std::path::Path,
z_path: &std::path::Path,
) -> bool {
let Ok(metadata) = std::fs::metadata(merged_path) else {
return false;
};
if metadata.len() != expected_size {
return false;
}
let Ok(source_fp) = merged_qkvz_source_fingerprint(qkv_path, z_path) else {
return false;
};
let Ok(file_fp) = merged_qkvz_file_fingerprint(merged_path) else {
return false;
};
source_fp == file_fp
}
#[cfg(any(test, feature = "metal-gpu"))]
pub(crate) fn write_merged_qkvz(
qkv_path: &std::path::Path,
z_path: &std::path::Path,
out_path: &std::path::Path,
) -> Result<(), String> {
use std::io::Write;
let (qkv_hdr, qkv_payload) = read_q4_payload_bounded(qkv_path, MAX_Q4_MERGE_PAYLOAD_LEN)
.map_err(|e| format!("read {}: {e}", qkv_path.display()))?;
let (z_hdr, z_payload) = read_q4_payload_bounded(z_path, MAX_Q4_MERGE_PAYLOAD_LEN)
.map_err(|e| format!("read {}: {e}", z_path.display()))?;
if qkv_hdr.shape.len() != 2 {
return Err(format!(
"{}: qkv header shape {:?} is not 2-D (expected [rows, hidden])",
qkv_path.display(),
qkv_hdr.shape
));
}
if z_hdr.shape.len() != 2 {
return Err(format!(
"{}: z header shape {:?} is not 2-D (expected [rows, hidden])",
z_path.display(),
z_hdr.shape
));
}
if qkv_hdr.shape[1] != z_hdr.shape[1] {
return Err(format!(
"{}/{}: qkv hidden dimension {} does not match z hidden dimension {}",
qkv_path.display(),
z_path.display(),
qkv_hdr.shape[1],
z_hdr.shape[1]
));
}
let merged_rows = qkv_hdr.shape[0]
.checked_add(z_hdr.shape[0])
.ok_or("merged row count overflows usize")?;
let cols = qkv_hdr.shape[1];
let original_len = qkv_hdr
.original_len
.checked_add(z_hdr.original_len)
.ok_or("merged original_len overflows usize")?;
let tmp = out_path.with_extension("q4.tmp");
let write_result = (|| -> Result<(), String> {
let mut f = std::io::BufWriter::new(
std::fs::File::create(&tmp).map_err(|e| format!("create {}: {e}", tmp.display()))?,
);
f.write_all(b"KHQ4").map_err(|e| e.to_string())?;
f.write_all(&2u32.to_le_bytes())
.map_err(|e| e.to_string())?;
f.write_all(&2u32.to_le_bytes())
.map_err(|e| e.to_string())?;
f.write_all(&(merged_rows as u64).to_le_bytes())
.map_err(|e| e.to_string())?;
f.write_all(&(cols as u64).to_le_bytes())
.map_err(|e| e.to_string())?;
f.write_all(&(original_len as u64).to_le_bytes())
.map_err(|e| e.to_string())?;
f.write_all(&qkv_payload).map_err(|e| e.to_string())?;
f.write_all(&z_payload).map_err(|e| e.to_string())?;
Ok(())
})();
if write_result.is_err() {
let _ = std::fs::remove_file(&tmp);
return write_result;
}
std::fs::rename(&tmp, out_path).map_err(|e| {
let _ = std::fs::remove_file(&tmp);
format!("rename: {e}")
})
}
#[cfg(test)]
mod tests {
use super::*;
fn q4_file_bytes(shape: &[usize], original_len: usize, scale: u16, bias: u16) -> Vec<u8> {
let mut buf = Vec::new();
buf.extend_from_slice(b"KHQ4");
buf.extend_from_slice(&2u32.to_le_bytes());
buf.extend_from_slice(&(shape.len() as u32).to_le_bytes());
for &dim in shape {
buf.extend_from_slice(&(dim as u64).to_le_bytes());
}
buf.extend_from_slice(&(original_len as u64).to_le_bytes());
for _ in 0..original_len.div_ceil(32) {
buf.extend_from_slice(&scale.to_ne_bytes());
buf.extend_from_slice(&bias.to_ne_bytes());
buf.extend_from_slice(&[0u8; 16]);
}
buf
}
fn f16_file_bytes(shape: &[usize], values: &[u16]) -> Vec<u8> {
let mut buf = Vec::new();
buf.extend_from_slice(b"KHF1");
buf.extend_from_slice(&1u32.to_le_bytes());
buf.extend_from_slice(&(shape.len() as u32).to_le_bytes());
for &dim in shape {
buf.extend_from_slice(&(dim as u64).to_le_bytes());
}
buf.extend_from_slice(&(values.len() as u64).to_le_bytes());
for &value in values {
buf.extend_from_slice(&value.to_le_bytes());
}
buf
}
#[test]
fn test_q4_block_size() {
assert_eq!(std::mem::size_of::<Q4Block>(), 20);
let b = Q4Block {
scale: 0,
bias: 0,
packed: [0u8; 16],
};
let base = std::ptr::from_ref(&b) as usize;
let packed_off = std::ptr::from_ref(&b.packed) as usize - base;
assert_eq!(
packed_off, 4,
"packed field must start at byte offset 4 (after scale + bias, no padding)"
);
}
#[test]
fn test_quantize_dequantize_zeros() {
let data = quantize_row_q4_0(&vec![0.0f32; 64]).unwrap();
let out = dequantize_row_q4_0(&data, 64);
assert_eq!(out.len(), 64);
for v in &out {
assert!(v.abs() < 1e-6, "expected ~0, got {v}");
}
}
#[test]
fn test_quantize_dequantize_small_values() {
let src: Vec<f32> = (0..32).map(|i| i as f32 * 7.0 / 31.0).collect();
let data = quantize_row_q4_0(&src).unwrap();
let out = dequantize_row_q4_0(&data, 32);
let max_err = src
.iter()
.zip(&out)
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_err < 0.5,
"max abs error {max_err:.4} >= 0.5 for small values"
);
}
#[test]
fn test_quantize_dequantize_symmetric() {
let src: Vec<f32> = (0..32).map(|i| (i as f32 - 15.5) / 15.5 * 7.0).collect();
let data = quantize_row_q4_0(&src).unwrap();
let out = dequantize_row_q4_0(&data, 32);
let max_err = src
.iter()
.zip(&out)
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_err < 0.5,
"max abs error {max_err:.4} >= 0.5 for symmetric values"
);
}
#[test]
fn test_quantize_max_range() {
let mut src = vec![0.0f32; 32];
src[0] = 7.0;
src[1] = -7.0;
let data = quantize_row_q4_0(&src).unwrap();
let block_byte0 = data[4];
assert_eq!(
block_byte0 & 0x0f,
15,
"w[0]=7.0 (max) should produce low nibble 15"
);
assert_eq!(
block_byte0 >> 4,
0,
"w[1]=-7.0 (min) should produce high nibble 0"
);
}
#[test]
fn test_quantize_single_block() {
let src: Vec<f32> = (0..32).map(|i| (i as f32 / 31.0) * 14.0 - 7.0).collect();
let data = quantize_row_q4_0(&src).unwrap();
assert_eq!(data.len(), 20, "single block must be 20 bytes");
let out = dequantize_row_q4_0(&data, 32);
assert_eq!(out.len(), 32);
let max_err = src
.iter()
.zip(&out)
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_err <= 0.51,
"max abs error {max_err:.4} > 0.51 for single block"
);
}
#[test]
fn test_quantize_multiple_blocks() {
let src: Vec<f32> = (0..128).map(|i| (i as f32 - 64.0) / 10.0).collect();
let data = quantize_row_q4_0(&src).unwrap();
assert_eq!(data.len(), 4 * 20, "4 blocks must be 80 bytes");
let out = dequantize_row_q4_0(&data, 128);
assert_eq!(out.len(), 128);
let max_err = src
.iter()
.zip(&out)
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_err < 0.5,
"max abs error {max_err:.4} >= 0.5 for multiple blocks"
);
}
#[test]
fn test_f16_roundtrip() {
let values = [
0.0f32,
1.0,
-1.0,
0.5,
-0.5,
std::f32::consts::PI,
100.0,
-100.0,
0.001,
65504.0, ];
for &v in &values {
let bits = q4_f32_to_f16(v);
let back = q4_f16_to_f32(bits);
let rel_err = if v.abs() > 1e-4 {
(v - back).abs() / v.abs()
} else {
(v - back).abs()
};
assert!(
rel_err < 0.004,
"f16 roundtrip failed for {v}: got {back}, rel_err={rel_err:.6}"
);
}
}
#[test]
fn test_f16_special_values() {
assert_eq!(q4_f32_to_f16(0.0f32), 0x0000);
assert_eq!(q4_f32_to_f16(-0.0f32), 0x8000);
assert_eq!(q4_f16_to_f32(0x0000), 0.0f32);
let pos_inf = q4_f32_to_f16(f32::INFINITY);
assert_eq!(pos_inf, 0x7c00);
assert!(q4_f16_to_f32(pos_inf).is_infinite() && q4_f16_to_f32(pos_inf) > 0.0);
let neg_inf = q4_f32_to_f16(f32::NEG_INFINITY);
assert_eq!(neg_inf, 0xfc00);
assert!(q4_f16_to_f32(neg_inf).is_infinite() && q4_f16_to_f32(neg_inf) < 0.0);
let nan_bits = q4_f32_to_f16(f32::NAN);
assert!(
q4_f16_to_f32(nan_bits).is_nan(),
"NaN should round-trip to NaN"
);
let overflow = q4_f32_to_f16(1.0e10f32);
assert_eq!(overflow, 0x7c00, "overflow should produce +∞");
}
#[test]
fn test_nibble_packing_order() {
let mut src = vec![0.0f32; 32];
src[0] = 0.0;
src[1] = 7.0;
let data = quantize_row_q4_0(&src).unwrap();
let byte0 = data[4];
assert_eq!(
byte0, 0xF0,
"byte[0] should be 0xF0 for w[0]=0.0 (nibble=0), w[1]=7.0 (nibble=15)"
);
let out = dequantize_row_q4_0(&data, 32);
assert!(
(out[0] - 0.0).abs() < 1e-3,
"weight[0] should be ~0.0, got {}",
out[0]
);
assert!(
(out[1] - 7.0).abs() < 0.05,
"weight[1] should be ~7.0, got {}",
out[1]
);
}
#[test]
fn test_quantize_tensor_rows() {
let rows = 4usize;
let cols = 64usize;
let src: Vec<f32> = (0..rows * cols)
.map(|i| (i as f32 - 128.0) / 20.0)
.collect();
let data = quantize_tensor_q4_0(&src, rows, cols).unwrap();
let blocks_per_row = cols.div_ceil(32); assert_eq!(
data.len(),
rows * blocks_per_row * 20,
"tensor bytes mismatch"
);
for row_idx in 0..rows {
let row_bytes =
&data[row_idx * blocks_per_row * 20..(row_idx + 1) * blocks_per_row * 20];
let out = dequantize_row_q4_0(row_bytes, cols);
let row_src = &src[row_idx * cols..(row_idx + 1) * cols];
let max_err = row_src
.iter()
.zip(&out)
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_err < 0.5,
"row {row_idx}: max abs error {max_err:.4} >= 0.5"
);
}
}
fn to_bf16(vals: &[f32]) -> Vec<u16> {
vals.iter()
.map(|&v| {
let bits = v.to_bits();
(bits >> 16) as u16
})
.collect()
}
fn bf16_round_trip(v: f32) -> f32 {
bf16_to_f32((v.to_bits() >> 16) as u16)
}
#[test]
fn test_quantize_dequantize_round_trip_zeros_bf16() {
let data = vec![0u16; 64];
let tensor = quantize_bf16_to_q4(&data, &[64]).unwrap();
let out = dequantize_q4_to_f32(&tensor);
assert_eq!(out.len(), 64);
for v in &out {
assert!(v.abs() < 1e-6, "expected ~0, got {v}");
}
}
#[test]
fn test_quantize_dequantize_round_trip_positive_bf16() {
let f32_vals: Vec<f32> = (0..32).map(|i| i as f32 * 7.0 / 31.0).collect();
let bf16_vals = to_bf16(&f32_vals);
let tensor = quantize_bf16_to_q4(&bf16_vals, &[32]).unwrap();
let out = dequantize_q4_to_f32(&tensor);
let max_err = f32_vals
.iter()
.zip(&out)
.map(|(a, b)| (bf16_round_trip(*a) - b).abs())
.fold(0.0f32, f32::max);
assert!(max_err <= 0.51, "max abs error {max_err:.4} > 0.51");
}
#[test]
fn test_nibble_packing_byte_value_bf16() {
let mut f32_vals = [0.0f32; 32];
f32_vals[0] = 0.0;
f32_vals[1] = 7.0;
let bf16_vals = to_bf16(&f32_vals);
let tensor = quantize_bf16_to_q4(&bf16_vals, &[32]).unwrap();
assert_eq!(tensor.blocks.len(), 1);
assert_eq!(
tensor.blocks[0].packed[0], 0xF0,
"byte[0] should be 0xF0 for w[0]=0.0 (nibble 0), w[1]=7.0 (nibble 15)"
);
}
#[test]
fn test_max_value_clamps_to_nibble_15() {
let mut f32_vals = [0.0f32; 32];
f32_vals[0] = 100.0;
let bf16_vals = to_bf16(&f32_vals);
let tensor = quantize_bf16_to_q4(&bf16_vals, &[32]).unwrap();
let low_nibble = tensor.blocks[0].packed[0] & 0x0f;
assert_eq!(low_nibble, 15, "weight[0]=100 should clamp to nibble 15");
}
#[test]
fn test_block_boundary_continuity() {
let mut f32_vals = Vec::with_capacity(64);
for i in 0..32 {
f32_vals.push((i % 7) as f32 + 1.0);
} for i in 0..32 {
f32_vals.push(-((i % 7) as f32 + 1.0));
} let bf16_vals = to_bf16(&f32_vals);
let tensor = quantize_bf16_to_q4(&bf16_vals, &[64]).unwrap();
assert_eq!(tensor.blocks.len(), 2);
let out = dequantize_q4_to_f32(&tensor);
for v in &out[0..32] {
assert!(*v > 0.0, "block 0 weight should be positive, got {v}");
}
for v in &out[32..64] {
assert!(*v < 0.0, "block 1 weight should be negative, got {v}");
}
}
#[test]
fn test_save_load_round_trip() {
let f32_vals: Vec<f32> = (0..64).map(|i| (i as f32 - 32.0) / 4.0).collect();
let bf16_vals = to_bf16(&f32_vals);
let original = quantize_bf16_to_q4(&bf16_vals, &[8, 8]).unwrap();
let path = std::path::PathBuf::from("/tmp/test_q4_round_trip.q4");
save_q4_file(&path, &original).unwrap();
let loaded = load_q4_file(&path).unwrap();
assert_eq!(loaded.shape, original.shape);
assert_eq!(loaded.original_len, original.original_len);
assert_eq!(loaded.blocks.len(), original.blocks.len());
for (a, b) in original.blocks.iter().zip(&loaded.blocks) {
assert_eq!(a.scale, b.scale, "scale mismatch after load");
assert_eq!(a.packed, b.packed, "packed mismatch after load");
}
std::fs::remove_file(&path).ok();
}
#[test]
fn test_stream_quantize_shard_matches_batch() {
let f32_vals: Vec<f32> = (0..96).map(|i| i as f32 / 10.0).collect();
let bf16_vals = to_bf16(&f32_vals);
let batch_tensor = quantize_bf16_to_q4(&bf16_vals, &[96]).unwrap();
let bf16_bytes: Vec<u8> = bf16_vals.iter().flat_map(|v| v.to_ne_bytes()).collect();
let stream_blocks = stream_quantize_shard(&bf16_bytes).unwrap();
assert_eq!(stream_blocks.len(), batch_tensor.blocks.len());
for (a, b) in batch_tensor.blocks.iter().zip(&stream_blocks) {
assert_eq!(a.scale, b.scale, "stream vs batch scale mismatch");
assert_eq!(a.packed, b.packed, "stream vs batch packed mismatch");
}
}
#[test]
fn test_shape_preservation() {
let shape = vec![4usize, 8, 4]; let data = vec![0u16; 128];
let tensor = quantize_bf16_to_q4(&data, &shape).unwrap();
assert_eq!(tensor.shape, shape);
assert_eq!(tensor.original_len, 128);
assert_eq!(tensor.blocks.len(), 4);
let path = std::path::PathBuf::from("/tmp/test_q4_shape.q4");
save_q4_file(&path, &tensor).unwrap();
let loaded = load_q4_file(&path).unwrap();
assert_eq!(loaded.shape, shape);
assert_eq!(loaded.original_len, 128);
std::fs::remove_file(&path).ok();
}
#[test]
fn test_round_trip_accuracy_tolerance() {
let mut state = 12345u64;
let mut f32_vals = Vec::with_capacity(1024);
for _ in 0..1024 {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
let v = ((state >> 32) as f32 / u32::MAX as f32) * 14.0 - 7.0;
f32_vals.push(v);
}
let data = quantize_row_q4_0(&f32_vals).unwrap();
let out = dequantize_row_q4_0(&data, 1024);
assert_eq!(out.len(), 1024);
let mae = f32_vals
.iter()
.zip(&out)
.map(|(a, b)| (a - b).abs())
.sum::<f32>()
/ 1024.0;
assert!(
mae < 0.30,
"mean abs error {mae:.4} >= 0.30 (Q4 MAE for uniform [-7,7] expected ≈ 0.25)"
);
}
fn f32_to_bf16_bits(v: f32) -> u16 {
let bits = v.to_bits();
let lsb = (bits >> 16) & 1;
let rounding_bias = 0x7fff + lsb;
((bits.wrapping_add(rounding_bias)) >> 16) as u16
}
fn synthetic_f32_uniform(n: usize, seed: u64) -> Vec<f32> {
let mut state = seed;
(0..n)
.map(|_| {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
let u = (state >> 32) as f32 / u32::MAX as f32;
u * 2.0 - 1.0
})
.collect()
}
#[test]
fn quantize_f32_to_q4_shape_and_length() {
let src = synthetic_f32_uniform(96, 17);
let q = quantize_f32_to_q4(&src, &[3, 32]).unwrap();
assert_eq!(q.shape, vec![3, 32]);
assert_eq!(q.original_len, 96);
assert_eq!(q.blocks.len(), 3, "96 elems = 3 full Q4 blocks");
}
#[test]
fn quantize_f32_to_q4_pads_partial_block() {
let src = synthetic_f32_uniform(40, 19);
let q = quantize_f32_to_q4(&src, &[40]).unwrap();
assert_eq!(q.original_len, 40);
assert_eq!(q.blocks.len(), 2, "40 elems = 1 full + 1 partial Q4 block");
}
#[test]
fn quantize_f32_to_q4_partial_block_uses_real_tail_min_max() {
let src = [5.0f32, 6.0, 7.0];
let q = quantize_f32_to_q4(&src, &[3]).unwrap();
assert_eq!(q.original_len, 3);
assert_eq!(q.blocks.len(), 1);
let block = q.blocks[0];
assert_eq!(block.scale, q4_f32_to_f16(2.0f32 / 15.0));
assert_eq!(block.bias, q4_f32_to_f16(5.0));
assert_ne!(
block.scale,
q4_f32_to_f16(7.0f32 / 15.0),
"partial tail scale must not include padded zero in max-min range"
);
assert_ne!(
block.bias,
q4_f32_to_f16(0.0),
"partial tail bias must be the real tail min, not padded zero"
);
}
#[test]
fn quantize_f64_to_q4_symmetric_partial_block_is_bit_identical_to_padded_block() {
let src = [5.0f64, -6.0, 7.0];
let mut padded = [0.0f32; 32];
for (dst, src) in padded.iter_mut().zip(src.iter()) {
*dst = *src as f32;
}
let expected = quantize_block_with_mode_len(&padded, 32, true).unwrap();
let q = quantize_f64_to_q4_mode(&src, &[3], true).unwrap();
assert_eq!(
q.blocks[0], expected,
"symmetric partial blocks must stay byte-identical to the old padded path"
);
}
const TINY_SYMMETRIC_ABS_MAX: f32 = 3.363e-37;
const TINY_ASYMMETRIC_RANGE: f32 = 1.5e-37;
#[test]
fn symmetric_block_with_underflowing_scale_is_quantizable_and_reloadable() {
let mut vals = [0.0f32; 32];
vals[0] = TINY_SYMMETRIC_ABS_MAX;
vals[7] = -TINY_SYMMETRIC_ABS_MAX / 3.0;
assert!(
q4_f16_to_f32(q4_f32_to_f16(TINY_SYMMETRIC_ABS_MAX / 7.0)) == 0.0,
"fixture must actually underflow f16, else this test proves nothing"
);
let block = quantize_block_with_mode_len(&vals, 32, true)
.expect("a block with a tiny but nonzero range must quantize");
let scale = q4_f16_to_f32(block.scale);
assert!(
scale.is_finite() && scale > 0.0,
"serialized scale {scale} must be finite and strictly positive"
);
assert_eq!(scale, 1.0, "degenerate blocks take the 1.0 fallback");
let tensor = Q4Tensor {
blocks: vec![block],
shape: vec![32],
original_len: 32,
};
let out = dequantize_q4_to_f32(&tensor);
for (i, (&got, &want)) in out.iter().zip(vals.iter()).enumerate() {
assert!(
(got - want).abs() <= 2.0 * TINY_SYMMETRIC_ABS_MAX,
"element {i}: reconstruction error {} exceeds the block's own range",
(got - want).abs()
);
}
}
#[test]
fn asymmetric_block_with_underflowing_scale_is_quantizable_and_reloadable() {
let mut vals = [0.0f32; 32];
vals[3] = TINY_ASYMMETRIC_RANGE;
assert!(
q4_f16_to_f32(q4_f32_to_f16(TINY_ASYMMETRIC_RANGE / 15.0)) == 0.0,
"fixture must actually underflow f16, else this test proves nothing"
);
let block = quantize_block_with_mode_len(&vals, 32, false)
.expect("a block with a tiny but nonzero range must quantize");
let scale = q4_f16_to_f32(block.scale);
assert!(
scale.is_finite() && scale > 0.0,
"serialized scale {scale} must be finite and strictly positive"
);
assert_eq!(scale, 1.0, "degenerate blocks take the 1.0 fallback");
let tensor = Q4Tensor {
blocks: vec![block],
shape: vec![32],
original_len: 32,
};
let out = dequantize_q4_to_f32(&tensor);
for (i, (&got, &want)) in out.iter().zip(vals.iter()).enumerate() {
assert!(
(got - want).abs() <= 2.0 * TINY_ASYMMETRIC_RANGE,
"element {i}: reconstruction error {} exceeds the block's own range",
(got - want).abs()
);
}
}
#[test]
fn quarot_symmetric_write_pass_survives_a_tiny_range_row() {
let mut src = vec![0.0f64; 64];
src[0] = f64::from(TINY_SYMMETRIC_ABS_MAX);
src[40] = f64::from(-TINY_SYMMETRIC_ABS_MAX) / 2.0;
let q = quantize_f64_to_q4(&src, &[64]).expect("symmetric quantize must accept the row");
for (i, b) in q.blocks.iter().enumerate() {
let scale = q4_f16_to_f32(b.scale);
assert!(
scale.is_finite() && scale > 0.0,
"block {i} serialized scale {scale} is not a usable f16"
);
}
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("tiny_range.q4");
save_q4_file(&path, &q).unwrap();
let reloaded = load_q4_file(&path).expect("the loader must accept what the writer emits");
assert_eq!(reloaded.blocks, q.blocks);
}
#[test]
fn scale_too_large_for_f16_is_still_rejected() {
let mut vals = [0.0f32; 32];
vals[0] = 1.0e7;
let err = quantize_block_with_mode_len(&vals, 32, true)
.expect_err("a scale above f16's maximum must not be silently replaced");
assert!(
format!("{err}").contains("strictly positive f16"),
"unexpected error: {err}"
);
}
#[test]
fn bias_above_f16_range_is_still_rejected() {
let mut vals = [1.0e5f32; 32];
vals[0] = 1.0e5 + 1.0;
let err = quantize_block_with_mode_len(&vals, 32, false)
.expect_err("a bias outside f16 range must be rejected");
assert!(
format!("{err}").contains("finite f16"),
"unexpected error: {err}"
);
}
#[test]
fn quantize_f64_to_q4_matches_f32_path_after_downcast() {
let src_f64: Vec<f64> = synthetic_f32_uniform(256, 23)
.into_iter()
.map(f64::from)
.collect();
let src_f32: Vec<f32> = src_f64.iter().map(|&v| v as f32).collect();
let q_f64 = quantize_f64_to_q4_mode(&src_f64, &[256], false).unwrap();
let q_f32 = quantize_f32_to_q4(&src_f32, &[256]).unwrap();
assert_eq!(q_f64.shape, q_f32.shape);
assert_eq!(q_f64.original_len, q_f32.original_len);
assert_eq!(
q_f64.blocks.len(),
q_f32.blocks.len(),
"f64 path must produce same block count"
);
for (i, (a, b)) in q_f64.blocks.iter().zip(q_f32.blocks.iter()).enumerate() {
assert_eq!(a.scale, b.scale, "block {i} scale mismatch");
assert_eq!(a.bias, b.bias, "block {i} bias mismatch");
assert_eq!(a.packed, b.packed, "block {i} packed mismatch");
}
}
#[test]
fn quantize_f32_to_q4_matches_bf16_path_when_input_is_bf16_castable() {
let bf16_bits: Vec<u16> = synthetic_f32_uniform(256, 29)
.into_iter()
.map(f32_to_bf16_bits)
.collect();
let f32_from_bf16: Vec<f32> = bf16_bits.iter().map(|&b| bf16_to_f32(b)).collect();
let q_bf16 = quantize_bf16_to_q4(&bf16_bits, &[256]).unwrap();
let q_f32 = quantize_f32_to_q4(&f32_from_bf16, &[256]).unwrap();
assert_eq!(q_bf16.blocks.len(), q_f32.blocks.len());
for (i, (a, b)) in q_bf16.blocks.iter().zip(q_f32.blocks.iter()).enumerate() {
assert_eq!(a.scale, b.scale, "block {i} scale should match");
assert_eq!(a.packed, b.packed, "block {i} packed should match");
}
}
#[test]
fn quantize_f32_to_q4_lower_error_than_bf16_path_on_high_precision_input() {
let src = synthetic_f32_uniform(2048, 31);
let bf16_bits: Vec<u16> = src.iter().map(|&v| f32_to_bf16_bits(v)).collect();
let q_bf16 = quantize_bf16_to_q4(&bf16_bits, &[2048]).unwrap();
let q_f32 = quantize_f32_to_q4(&src, &[2048]).unwrap();
let deq_bf16 = dequantize_q4_to_f32(&q_bf16);
let deq_f32 = dequantize_q4_to_f32(&q_f32);
let err = |reconstructed: &[f32]| -> (f32, f32) {
let mut max_err = 0.0_f32;
let mut sum_err = 0.0_f32;
for (s, r) in src.iter().zip(reconstructed.iter()) {
let e = (s - r).abs();
max_err = max_err.max(e);
sum_err += e;
}
(max_err, sum_err / src.len() as f32)
};
let (max_bf16, mean_bf16) = err(&deq_bf16);
let (max_f32, mean_f32) = err(&deq_f32);
eprintln!(
"[3c-1 measurement] n=2048 source=f32 uniform [-1,1]: \
f32_path mean_abs_err={mean_f32:.6} max_abs_err={max_f32:.6}; \
bf16_path mean_abs_err={mean_bf16:.6} max_abs_err={max_bf16:.6}"
);
assert!(
mean_f32 < mean_bf16,
"f32 mean abs error ({mean_f32:.6}) should be < bf16 mean abs error ({mean_bf16:.6})"
);
assert!(
max_f32 <= max_bf16,
"f32 max abs error ({max_f32:.6}) should be <= bf16 max abs error ({max_bf16:.6})"
);
}
#[test]
fn quarot_rotated_q4_forward_matches_f64_reference() {
use std::collections::HashMap;
use crate::quant::quarot::hadamard::RandomizedHadamard;
use crate::quant::quarot::pipeline::{TensorEntry, absorb_rotations};
use crate::quant::quarot::plan::RotationPlan;
const HIDDEN: usize = 32;
const Q_ROWS: usize = 2; const O_ROWS: usize = 32;
let q_name = "model.language_model.layers.0.self_attn.q_proj.weight";
let o_name = "model.language_model.layers.0.self_attn.o_proj.weight";
fn lcg_f64(n: usize, seed: u64) -> Vec<f64> {
let mut state = seed;
(0..n)
.map(|_| {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
(state >> 32) as f64 / u32::MAX as f64 * 2.0 - 1.0
})
.collect()
}
fn matvec(w: &[f64], rows: usize, cols: usize, x: &[f64]) -> Vec<f64> {
(0..rows)
.map(|r| {
w[r * cols..(r + 1) * cols]
.iter()
.zip(x)
.map(|(a, b)| a * b)
.sum()
})
.collect()
}
fn max_diff(a: &[f64], b: &[f64]) -> f64 {
let mut max = 0.0_f64;
for (x, y) in a.iter().zip(b) {
let d = (x - y).abs();
if !d.is_finite() {
return d;
}
if d > max {
max = d;
}
}
max
}
fn ref_q4_dequant(data: &[f64]) -> Vec<f64> {
let mut out = Vec::with_capacity(data.len());
for chunk in data.chunks(32) {
let f32s: Vec<f32> = chunk.iter().map(|&v| v as f32).collect();
let abs_max = f32s.iter().map(|v| v.abs()).fold(0.0_f32, f32::max);
let scale_f32 = if abs_max == 0.0 {
1.0_f32
} else {
abs_max / 7.0
};
let bias_f32 = -8.0_f32 * scale_f32;
let scale_dq = q4_f16_to_f32(q4_f32_to_f16(scale_f32));
let bias_dq = q4_f16_to_f32(q4_f32_to_f16(bias_f32));
let inv_scale = 1.0 / scale_f32;
for &v in &f32s {
let nibble = ((v * inv_scale).round() + 8.0).clamp(0.0, 15.0) as u8;
out.push(f64::from(nibble as f32 * scale_dq + bias_dq));
}
}
out
}
let q_data_orig = lcg_f64(Q_ROWS * HIDDEN, 0x1111_1111_1111_1111);
let o_data_orig = lcg_f64(O_ROWS * HIDDEN, 0x2222_2222_2222_2222);
let x_q = lcg_f64(HIDDEN, 0x3333_3333_3333_3333);
let x_o = lcg_f64(HIDDEN, 0x4444_4444_4444_4444);
let rotation = RandomizedHadamard::new(0x3200_0001, HIDDEN).expect("rotation init");
let plan = RotationPlan::qwen35_residual_stream_linear_layers();
let mut tensors: HashMap<String, TensorEntry> = HashMap::new();
tensors.insert(
q_name.to_string(),
TensorEntry {
name: q_name.to_string(),
shape: vec![Q_ROWS, HIDDEN],
data: q_data_orig.clone(),
},
);
tensors.insert(
o_name.to_string(),
TensorEntry {
name: o_name.to_string(),
shape: vec![O_ROWS, HIDDEN],
data: o_data_orig.clone(),
},
);
absorb_rotations(&mut tensors, &plan, &rotation).expect("absorb_rotations");
let q_q4 =
quantize_f64_to_q4(&tensors[q_name].data, &[Q_ROWS, HIDDEN]).expect("q_proj quantize");
let o_q4 =
quantize_f64_to_q4(&tensors[o_name].data, &[O_ROWS, HIDDEN]).expect("o_proj quantize");
assert_eq!(q_q4.shape, vec![Q_ROWS, HIDDEN], "q_proj shape");
assert_eq!(q_q4.original_len, Q_ROWS * HIDDEN, "q_proj original_len");
assert_eq!(q_q4.blocks.len(), Q_ROWS, "[2,32] must produce 2 Q4 blocks");
assert_eq!(o_q4.shape, vec![O_ROWS, HIDDEN], "o_proj shape");
assert_eq!(o_q4.original_len, O_ROWS * HIDDEN, "o_proj original_len");
assert_eq!(
o_q4.blocks.len(),
O_ROWS,
"[32,32] must produce 32 Q4 blocks"
);
let q_deq: Vec<f64> = dequantize_q4_to_f32(&q_q4)
.into_iter()
.map(f64::from)
.collect();
let o_deq: Vec<f64> = dequantize_q4_to_f32(&o_q4)
.into_iter()
.map(f64::from)
.collect();
let prod_y_q = matvec(&q_deq, Q_ROWS, HIDDEN, &x_q);
let prod_y_o = matvec(&o_deq, O_ROWS, HIDDEN, &x_o);
let mut q_ref = q_data_orig.clone();
for r in 0..Q_ROWS {
rotation
.apply_f64(&mut q_ref[r * HIDDEN..(r + 1) * HIDDEN])
.expect("q_proj row rotation");
}
let mut o_ref = o_data_orig.clone();
let mut col_buf = vec![0.0_f64; O_ROWS];
for c in 0..HIDDEN {
for r in 0..O_ROWS {
col_buf[r] = o_ref[r * HIDDEN + c];
}
rotation
.apply_f64(&mut col_buf)
.expect("o_proj col rotation");
for r in 0..O_ROWS {
o_ref[r * HIDDEN + c] = col_buf[r];
}
}
let q_ref_deq = ref_q4_dequant(&q_ref);
let o_ref_deq = ref_q4_dequant(&o_ref);
let ref_y_q = matvec(&q_ref_deq, Q_ROWS, HIDDEN, &x_q);
let ref_y_o = matvec(&o_ref_deq, O_ROWS, HIDDEN, &x_o);
let max_q = max_diff(&prod_y_q, &ref_y_q);
let max_o = max_diff(&prod_y_o, &ref_y_o);
eprintln!("[quarot_q4_gate] max_abs_diff q_proj={max_q:.2e} o_proj={max_o:.2e}");
assert!(
max_q <= 1e-5,
"q_proj forward max_abs_diff {max_q:.2e} > 1e-5: composed rotated+Q4 path is broken"
);
assert!(
max_o <= 1e-5,
"o_proj forward max_abs_diff {max_o:.2e} > 1e-5: composed rotated+Q4 path is broken"
);
}
#[test]
fn quantize_f32_to_q4_rejects_shape_data_mismatch() {
let data = synthetic_f32_uniform(64, 41);
let err = quantize_f32_to_q4(&data, &[3, 32])
.expect_err("shape claiming 96 elements for 64 values must fail");
assert!(err.to_string().contains("shape product"));
}
#[test]
fn quantize_f64_to_q4_rejects_shape_data_mismatch() {
let data: Vec<f64> = synthetic_f32_uniform(64, 43)
.into_iter()
.map(f64::from)
.collect();
let err = quantize_f64_to_q4(&data, &[3, 32])
.expect_err("shape claiming 96 elements for 64 values must fail");
assert!(err.to_string().contains("shape product"));
}
#[test]
fn quantize_bf16_to_q4_rejects_shape_data_mismatch() {
let data: Vec<u16> = (0..64).map(|i| i as u16).collect();
let err = quantize_bf16_to_q4(&data, &[3, 32])
.expect_err("shape claiming 96 elements for 64 values must fail");
assert!(err.to_string().contains("shape product"));
}
#[test]
fn quantize_f32_to_q4_rejects_shape_product_overflow() {
let data = vec![0.0_f32; 32];
let err = quantize_f32_to_q4(&data, &[usize::MAX, 2])
.expect_err("overflowed shape product must fail");
assert!(err.to_string().contains("overflows usize"));
}
#[test]
fn quantize_f32_to_q4_block_layout_matches_quantize_row() {
let src = synthetic_f32_uniform(128, 37);
let q = quantize_f32_to_q4(&src, &[128]).unwrap();
let row_bytes = quantize_row_q4_0(&src).unwrap();
assert_eq!(row_bytes.len(), q.blocks.len() * 20);
let q_bytes: &[u8] = unsafe {
std::slice::from_raw_parts(q.blocks.as_ptr().cast::<u8>(), q.blocks.len() * 20)
};
assert_eq!(q_bytes, row_bytes.as_slice());
}
#[test]
fn dequantize_row_q4_0_misaligned_does_not_panic() {
let src: Vec<f32> = (0..32).map(|i| (i as f32 / 31.0) * 14.0 - 7.0).collect();
let mut buf = quantize_row_q4_0(&src).unwrap(); assert_eq!(buf.len(), 20);
buf.extend_from_slice(&[0xDE, 0xAD, 0xBE, 0xEF, 0xFF]);
assert_eq!(buf.len(), 25);
let out = dequantize_row_q4_0(&buf, 32);
assert_eq!(out.len(), 32);
let max_err = src
.iter()
.zip(&out)
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_err <= 0.51,
"max abs error {max_err:.4} > 0.51 for single-block misaligned input"
);
}
#[test]
fn dequantize_row_q4_0_truncated_below_one_block() {
let buf = vec![0xABu8; 10]; let out = dequantize_row_q4_0(&buf, 32);
assert!(
out.is_empty(),
"expected empty Vec for sub-block input, got {} values",
out.len()
);
}
#[test]
fn dequantize_row_q4_0_exact_blocks_unchanged() {
let src: Vec<f32> = (0..64).map(|i| (i as f32 - 32.0) / 10.0).collect();
let data = quantize_row_q4_0(&src).unwrap();
assert_eq!(data.len(), 40, "2-block input must be 40 bytes");
let out = dequantize_row_q4_0(&data, 64);
assert_eq!(out.len(), 64);
let max_err = src
.iter()
.zip(&out)
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_err < 0.5,
"max abs error {max_err:.4} >= 0.5 for exact 2-block input"
);
}
#[test]
fn test_q4_rejects_huge_ndim() {
let mut buf = Vec::new();
buf.extend_from_slice(b"KHQ4");
buf.extend_from_slice(&2u32.to_le_bytes());
buf.extend_from_slice(&u32::MAX.to_le_bytes());
let path = std::path::PathBuf::from("/tmp/test_q4_huge_ndim.q4");
std::fs::write(&path, &buf).unwrap();
let r = load_q4_file(&path);
std::fs::remove_file(&path).ok();
assert!(
r.is_err(),
"u32::MAX ndim must be rejected, not OOM-aborted"
);
}
#[test]
fn test_read_q4_header_rejects_huge_ndim() {
let mut buf = Vec::new();
buf.extend_from_slice(b"KHQ4");
buf.extend_from_slice(&2u32.to_le_bytes());
buf.extend_from_slice(&u32::MAX.to_le_bytes());
let path = std::path::PathBuf::from("/tmp/test_q4_header_huge_ndim.q4");
std::fs::write(&path, &buf).unwrap();
let mut file = std::fs::File::open(&path).unwrap();
let r = read_q4_header(&mut file);
std::fs::remove_file(&path).ok();
assert!(
r.is_err(),
"u32::MAX ndim in read_q4_header must be rejected"
);
}
#[test]
fn read_q4_header_positions_cursor_at_first_block() {
use std::io::Read;
let first_block = Q4Block {
scale: q4_f32_to_f16(0.5),
bias: q4_f32_to_f16(-1.0),
packed: [0xA5; 16],
};
let tensor = Q4Tensor {
blocks: vec![first_block],
shape: vec![32],
original_len: 32,
};
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("cursor.q4");
save_q4_file(&path, &tensor).unwrap();
let mut file = std::fs::File::open(&path).unwrap();
let header = read_q4_header(&mut file).unwrap();
assert_eq!(header.payload_offset, 28);
let mut actual = [0u8; Q4_BLOCK_BYTES];
file.read_exact(&mut actual).unwrap();
let mut expected = [0u8; Q4_BLOCK_BYTES];
expected[0..2].copy_from_slice(&first_block.scale.to_ne_bytes());
expected[2..4].copy_from_slice(&first_block.bias.to_ne_bytes());
expected[4..].copy_from_slice(&first_block.packed);
assert_eq!(actual, expected);
}
#[test]
fn test_q4_rejects_huge_original_len() {
let mut buf = Vec::new();
buf.extend_from_slice(b"KHQ4");
buf.extend_from_slice(&2u32.to_le_bytes());
buf.extend_from_slice(&1u32.to_le_bytes()); buf.extend_from_slice(&4u64.to_le_bytes()); buf.extend_from_slice(&(1u64 << 62).to_le_bytes()); let path = std::path::PathBuf::from("/tmp/test_q4_huge_len.q4");
std::fs::write(&path, &buf).unwrap();
let r = load_q4_file(&path);
std::fs::remove_file(&path).ok();
assert!(
r.is_err(),
"2^62 original_len must be rejected, not OOM-aborted"
);
}
#[test]
fn test_q4_rejects_shape_product_mismatch() {
let mut buf = Vec::new();
buf.extend_from_slice(b"KHQ4");
buf.extend_from_slice(&2u32.to_le_bytes()); buf.extend_from_slice(&2u32.to_le_bytes()); buf.extend_from_slice(&4u64.to_le_bytes()); buf.extend_from_slice(&16u64.to_le_bytes()); buf.extend_from_slice(&32u64.to_le_bytes()); buf.extend_from_slice(&q4_f32_to_f16(1.0).to_ne_bytes());
buf.extend_from_slice(&q4_f32_to_f16(0.0).to_ne_bytes());
buf.extend_from_slice(&[0u8; 16]); let path = std::path::PathBuf::from("/tmp/test_q4_shape_mismatch.q4");
std::fs::write(&path, &buf).unwrap();
let r = load_q4_file(&path);
std::fs::remove_file(&path).ok();
assert!(
r.is_err(),
"shape product 64 != original_len 32 must be rejected"
);
}
#[test]
fn q4_ingress_rejects_invalid_scale_and_bias_metadata() {
let cases = [
(q4_f32_to_f16(f32::NAN), q4_f32_to_f16(0.0), "NaN scale"),
(
q4_f32_to_f16(f32::INFINITY),
q4_f32_to_f16(0.0),
"infinite scale",
),
(q4_f32_to_f16(0.0), q4_f32_to_f16(0.0), "zero scale"),
(q4_f32_to_f16(-1.0), q4_f32_to_f16(0.0), "negative scale"),
(q4_f32_to_f16(1.0), q4_f32_to_f16(f32::NAN), "NaN bias"),
];
for (scale, bias, label) in cases {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("invalid_metadata.q4");
std::fs::write(&path, q4_file_bytes(&[32], 32, scale, bias)).unwrap();
let err = load_q4_file(&path).expect_err(label);
assert!(
err.to_string().contains("block 0"),
"{label} error must identify its block: {err}"
);
}
}
#[test]
fn validate_q4_file_does_not_scan_block_payload() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("garbage_blocks.q4");
std::fs::write(
&path,
q4_file_bytes(&[32], 32, q4_f32_to_f16(f32::NAN), q4_f32_to_f16(f32::NAN)),
)
.unwrap();
let mut file = std::fs::File::open(&path).unwrap();
let result = validate_q4_file(&mut file, &path, Some(&[32]));
assert!(
result.is_ok(),
"validate_q4_file must not scan block payload bytes, but got: {:?}",
result.err()
);
let load_result = load_q4_file(&path);
assert!(
load_result.is_err(),
"load_q4_file must still reject non-finite block metadata, folded into its own \
single payload read"
);
}
fn write_single_block_q4(tmp: &tempfile::TempDir, scale: u16, bias: u16) -> std::path::PathBuf {
let path = tmp.path().join("block_metadata.q4");
std::fs::write(&path, q4_file_bytes(&[32], 32, scale, bias)).unwrap();
path
}
#[test]
fn mmap_entry_point_rejects_nan_scale_when_caller_does_not_traverse() {
let tmp = tempfile::tempdir().unwrap();
let path = write_single_block_q4(&tmp, q4_f32_to_f16(f32::NAN), q4_f32_to_f16(0.0));
let Err(err) = open_and_mmap_q4_file(
&path,
Some(&[32]),
Q4BlockCheck::Now {
tensor_name: "nan scale weight",
},
) else {
panic!("NaN block scale must not reach a no-copy GPU buffer");
};
assert!(
err.contains("block 0"),
"rejection must name the offending block: {err}"
);
}
#[test]
fn mmap_entry_point_rejects_infinite_scale_when_caller_does_not_traverse() {
let tmp = tempfile::tempdir().unwrap();
let path = write_single_block_q4(&tmp, q4_f32_to_f16(f32::INFINITY), q4_f32_to_f16(0.0));
let Err(err) = open_and_mmap_q4_file(
&path,
Some(&[32]),
Q4BlockCheck::Now {
tensor_name: "infinite scale weight",
},
) else {
panic!("infinite block scale must not reach a no-copy GPU buffer");
};
assert!(
err.contains("block 0"),
"rejection must name the offending block: {err}"
);
}
#[test]
fn mmap_entry_point_rejects_nan_bias_when_caller_does_not_traverse() {
let tmp = tempfile::tempdir().unwrap();
let path = write_single_block_q4(&tmp, q4_f32_to_f16(1.0), q4_f32_to_f16(f32::NAN));
let Err(err) = open_and_mmap_q4_file(
&path,
Some(&[32]),
Q4BlockCheck::Now {
tensor_name: "nan bias weight",
},
) else {
panic!("NaN block bias must not reach a no-copy GPU buffer");
};
assert!(
err.contains("block 0"),
"rejection must name the offending block: {err}"
);
}
#[test]
fn mmap_entry_point_defers_block_check_to_a_traversing_caller() {
let tmp = tempfile::tempdir().unwrap();
let path = write_single_block_q4(&tmp, q4_f32_to_f16(f32::NAN), q4_f32_to_f16(f32::NAN));
let result = open_and_mmap_q4_file(
&path,
Some(&[32]),
Q4BlockCheck::InCallerTraversal {
traversal: "test stand-in for a caller decode loop",
},
);
let (_header, _mmap, checked) = result.unwrap_or_else(|e| {
panic!(
"a traversing caller validates during its own pass, so the mapping must be \
handed out here: {e}"
)
});
assert!(
checked.is_none(),
"a deferred check must not hand out the witness that lets bytes be published \
to a consumer which never decodes them"
);
}
#[test]
fn mmap_entry_point_accepts_well_formed_q4_file() {
let tmp = tempfile::tempdir().unwrap();
let path = write_single_block_q4(&tmp, q4_f32_to_f16(0.25), q4_f32_to_f16(-1.0));
let (header, mmap, checked) = open_and_mmap_q4_file(
&path,
Some(&[32]),
Q4BlockCheck::Now {
tensor_name: "well formed weight",
},
)
.expect("a well-formed Q4 file must still load through the eager-check path");
assert!(
checked.is_some(),
"an eagerly checked mapping must yield the witness a no-copy consumer needs"
);
assert_eq!(header.shape, vec![32]);
assert_eq!(header.original_len, 32);
assert_eq!(
mmap.len() as u64,
header.payload_offset + Q4_BLOCK_BYTES as u64
);
}
#[test]
fn q4_ingress_rejects_trailing_bytes() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("trailing.q4");
let mut bytes = q4_file_bytes(&[32], 32, q4_f32_to_f16(1.0), q4_f32_to_f16(0.0));
bytes.push(0xAA);
std::fs::write(&path, bytes).unwrap();
let err = load_q4_file(&path).expect_err("trailing byte must be rejected");
assert!(
err.to_string().contains("trailing"),
"unexpected error: {err}"
);
}
#[test]
fn checked_native_loaders_reject_same_numel_transposed_geometry() {
let tmp = tempfile::tempdir().unwrap();
let q4_path = tmp.path().join("transposed.q4");
std::fs::write(
&q4_path,
q4_file_bytes(&[2, 32], 64, q4_f32_to_f16(1.0), q4_f32_to_f16(0.0)),
)
.unwrap();
let f16_path = tmp.path().join("transposed.f16");
std::fs::write(&f16_path, f16_file_bytes(&[2, 32], &[0u16; 64])).unwrap();
let q4_err = load_q4_file_checked(&q4_path, &[32, 2])
.expect_err("same-numel transposed Q4 shape must be rejected");
assert!(q4_err.to_string().contains("expected [32, 2]"));
let f16_err = load_f16_tensor_file_checked(&f16_path, &[32, 2])
.expect_err("same-numel transposed F16 shape must be rejected");
assert!(f16_err.to_string().contains("expected [32, 2]"));
}
#[test]
fn test_validate_q4_header_payload_bounds_rejects_truncated_payload() {
let header = Q4FileHeader {
shape: vec![64],
original_len: 64,
payload_offset: 28,
};
let r = validate_q4_header_payload_bounds(&header, 28, &std::path::PathBuf::from("t.q4"));
assert!(
r.is_err(),
"file truncated to payload_offset must be rejected"
);
}
#[test]
fn test_validate_q4_header_payload_bounds_rejects_one_byte_short() {
let header = Q4FileHeader {
shape: vec![64],
original_len: 64,
payload_offset: 28,
};
let r = validate_q4_header_payload_bounds(&header, 67, &std::path::PathBuf::from("t.q4"));
assert!(r.is_err(), "payload one byte short of required must fail");
}
#[test]
fn test_validate_q4_header_payload_bounds_accepts_exact_length() {
let header = Q4FileHeader {
shape: vec![64],
original_len: 64,
payload_offset: 28,
};
let r = validate_q4_header_payload_bounds(&header, 68, &std::path::PathBuf::from("t.q4"));
assert!(
r.is_ok(),
"file with exactly the required payload bytes must be accepted: {r:?}"
);
}
#[test]
fn test_validate_q4_header_payload_bounds_rejects_trailing_byte() {
let header = Q4FileHeader {
shape: vec![32],
original_len: 32,
payload_offset: 28,
};
let r = validate_q4_header_payload_bounds(&header, 49, std::path::Path::new("t.q4"));
let err = r.expect_err("one trailing byte must be rejected");
assert!(err.to_string().contains("trailing"));
}
#[test]
fn test_validate_q4_header_payload_bounds_rejects_shape_mismatch() {
let header = Q4FileHeader {
shape: vec![4, 16], original_len: 32, payload_offset: 36,
};
let r = validate_q4_header_payload_bounds(&header, 56, std::path::Path::new("t.q4"));
assert!(
r.is_err(),
"shape product != original_len must be rejected before a payload-length check"
);
}
#[test]
fn test_validate_q4_header_payload_bounds_rejects_huge_original_len_overflow() {
let header = Q4FileHeader {
shape: vec![usize::MAX],
original_len: usize::MAX,
payload_offset: 28,
};
let r =
validate_q4_header_payload_bounds(&header, 1_000, &std::path::PathBuf::from("t.q4"));
assert!(
r.is_err(),
"huge original_len must be rejected, not panic on overflow"
);
}
#[test]
fn test_f16_rejects_huge_numel() {
let mut buf = Vec::new();
buf.extend_from_slice(b"KHF1");
buf.extend_from_slice(&1u32.to_le_bytes());
buf.extend_from_slice(&1u32.to_le_bytes()); buf.extend_from_slice(&(1u64 << 63).to_le_bytes()); buf.extend_from_slice(&(1u64 << 63).to_le_bytes()); let path = std::path::PathBuf::from("/tmp/test_f16_huge_numel.f16");
std::fs::write(&path, &buf).unwrap();
let r = load_f16_tensor_file(&path);
std::fs::remove_file(&path).ok();
assert!(
r.is_err(),
"2^63 numel must be rejected, not silently truncated to empty"
);
}
#[test]
fn test_f16_rejects_shape_numel_mismatch() {
let mut buf = Vec::new();
buf.extend_from_slice(b"KHF1");
buf.extend_from_slice(&1u32.to_le_bytes());
buf.extend_from_slice(&2u32.to_le_bytes()); buf.extend_from_slice(&2u64.to_le_bytes()); buf.extend_from_slice(&2u64.to_le_bytes()); buf.extend_from_slice(&1u64.to_le_bytes()); buf.extend_from_slice(&0u16.to_le_bytes()); let path = std::path::PathBuf::from("/tmp/test_f16_shape_numel_mismatch.f16");
std::fs::write(&path, &buf).unwrap();
let r = load_f16_tensor_file(&path);
std::fs::remove_file(&path).ok();
let err = r.expect_err("shape product != numel must be rejected");
assert!(
err.to_string().contains("shape product"),
"unexpected error: {err}"
);
}
#[test]
fn f16_ingress_rejects_non_finite_values_and_trailing_bytes() {
for (bits, label) in [
(q4_f32_to_f16(f32::NAN), "NaN"),
(q4_f32_to_f16(f32::INFINITY), "infinity"),
] {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("non_finite.f16");
std::fs::write(&path, f16_file_bytes(&[1], &[bits])).unwrap();
let err = load_f16_tensor_file(&path).expect_err(label);
assert!(err.to_string().contains("element index 0"));
}
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("trailing.f16");
let mut bytes = f16_file_bytes(&[1], &[q4_f32_to_f16(1.0)]);
bytes.push(0xAA);
std::fs::write(&path, bytes).unwrap();
let err = load_f16_tensor_file(&path).expect_err("trailing byte must be rejected");
assert!(
err.to_string().contains("trailing"),
"unexpected error: {err}"
);
}
#[test]
fn quantizers_return_errors_for_shape_mismatches() {
assert!(quantize_tensor_q4_0(&[0.0], 1, 2).is_err());
assert!(quantize_bf16_to_q4(&[0], &[2]).is_err());
assert!(quantize_f32_to_q4(&[0.0], &[2]).is_err());
assert!(quantize_f64_to_q4_mode(&[0.0], &[2], true).is_err());
}
#[test]
fn quantizer_rejects_f16_metadata_overflow() {
let err = quantize_row_q4_0(&[f32::MAX; 32])
.expect_err("finite source whose serialized bias overflows f16 must be rejected");
assert!(err.to_string().contains("f16"));
let mut symmetric = [0.0f64; 32];
symmetric[0] = 500_000.0;
let err = quantize_f64_to_q4_mode(&symmetric, &[32], true)
.expect_err("finite source whose serialized symmetric metadata overflows f16");
assert!(err.to_string().contains("f16"));
}
#[test]
fn test_f16_rejects_huge_ndim() {
let mut buf = Vec::new();
buf.extend_from_slice(b"KHF1");
buf.extend_from_slice(&1u32.to_le_bytes());
buf.extend_from_slice(&u32::MAX.to_le_bytes());
let path = std::path::PathBuf::from("/tmp/test_f16_huge_ndim.f16");
std::fs::write(&path, &buf).unwrap();
let r = load_f16_tensor_file(&path);
std::fs::remove_file(&path).ok();
assert!(
r.is_err(),
"u32::MAX ndim in .f16 must be rejected, not OOM-aborted"
);
}
#[test]
fn test_q4_rejects_original_len_near_usize_max() {
let huge = usize::MAX - 3;
let mut buf = Vec::new();
buf.extend_from_slice(b"KHQ4");
buf.extend_from_slice(&2u32.to_le_bytes()); buf.extend_from_slice(&1u32.to_le_bytes()); buf.extend_from_slice(&(huge as u64).to_le_bytes()); buf.extend_from_slice(&(huge as u64).to_le_bytes()); let path = std::path::PathBuf::from("/tmp/test_q4_original_len_near_usize_max.q4");
std::fs::write(&path, &buf).unwrap();
let r = load_q4_file(&path);
std::fs::remove_file(&path).ok();
let err = r.expect_err("original_len near usize::MAX must be rejected, not panic/OOM");
let msg = err.to_string();
assert!(
msg.contains("block payload") || msg.contains("header claims"),
"expected the block-payload allocation guard to fire, got: {msg}"
);
}
#[test]
fn test_f16_rejects_numel_whose_byte_count_exceeds_file_len() {
let huge = usize::MAX / 4;
let mut buf = Vec::new();
buf.extend_from_slice(b"KHF1");
buf.extend_from_slice(&1u32.to_le_bytes());
buf.extend_from_slice(&1u32.to_le_bytes()); buf.extend_from_slice(&(huge as u64).to_le_bytes()); buf.extend_from_slice(&(huge as u64).to_le_bytes()); let path = std::path::PathBuf::from("/tmp/test_f16_numel_exceeds_file_len.f16");
std::fs::write(&path, &buf).unwrap();
let r = load_f16_tensor_file(&path);
std::fs::remove_file(&path).ok();
let err = r.expect_err("oversized f16 numel must be rejected, not panic/OOM");
let msg = err.to_string();
assert!(
msg.contains("f16 data") && msg.contains("header claims"),
"expected the f16-data file_len-bound guard to fire, got: {msg}"
);
}
#[test]
fn test_f16_rejects_numel_that_wraps_to_small_value_on_overflow() {
let huge = usize::MAX / 2 + 5;
let mut buf = Vec::new();
buf.extend_from_slice(b"KHF1");
buf.extend_from_slice(&1u32.to_le_bytes());
buf.extend_from_slice(&1u32.to_le_bytes()); buf.extend_from_slice(&(huge as u64).to_le_bytes()); buf.extend_from_slice(&(huge as u64).to_le_bytes()); buf.extend_from_slice(&[0xABu8; 64]);
let path =
std::path::PathBuf::from("/tmp/test_f16_numel_wraps_to_small_value_on_overflow.f16");
std::fs::write(&path, &buf).unwrap();
let r = load_f16_tensor_file(&path);
std::fs::remove_file(&path).ok();
let err = r.expect_err(
"numel whose ×2 wraps to a small value must still be rejected via checked_mul, \
not silently accepted as a tiny (wrong) allocation",
);
let msg = err.to_string();
assert!(
msg.contains("overflows usize"),
"expected the checked_mul overflow branch specifically, got: {msg}"
);
}
#[test]
fn test_quantize_block_rejects_nan_input() {
let mut vals = vec![1.0f32; 32];
vals[7] = f32::NAN;
let result = quantize_row_q4_0(&vals);
assert!(
result.is_err(),
"NaN in weight block must be rejected with InvalidInput"
);
}
#[test]
fn test_quantize_block_rejects_inf_input() {
let mut vals = vec![1.0f32; 32];
vals[15] = f32::INFINITY;
let result = quantize_row_q4_0(&vals);
assert!(
result.is_err(),
"+inf in weight block must be rejected with InvalidInput"
);
}
fn write_test_q4_source(path: &std::path::Path, rows: usize, cols: usize, seed: f32) {
let n = rows * cols;
let f32_vals: Vec<f32> = (0..n).map(|i| (i as f32 + seed) % 7.0 - 3.0).collect();
let bf16_vals = to_bf16(&f32_vals);
let tensor = quantize_bf16_to_q4(&bf16_vals, &[rows, cols]).unwrap();
save_q4_file(path, &tensor).unwrap();
}
fn merge_test_paths(
name: &str,
) -> (std::path::PathBuf, std::path::PathBuf, std::path::PathBuf) {
let dir = std::env::temp_dir().join(format!("lattice_test_merged_qkvz_{name}"));
std::fs::create_dir_all(&dir).unwrap();
(dir.join("qkv.q4"), dir.join("z.q4"), dir.join("merged.q4"))
}
#[test]
fn test_merged_qkvz_expected_size_computes_correctly() {
let expected = merged_qkvz_expected_size(136, 76).unwrap();
assert_eq!(expected, 36 + 100 + 40);
}
#[test]
fn test_merged_qkvz_expected_size_rejects_truncated_source() {
let err = merged_qkvz_expected_size(20, 136).unwrap_err();
assert!(
err.contains("too small"),
"expected a too-small error, got: {err}"
);
}
#[test]
fn test_write_merged_qkvz_then_cache_is_valid() {
let (qkv_p, z_p, merged_p) = merge_test_paths("valid");
write_test_q4_source(&qkv_p, 4, 8, 1.0);
write_test_q4_source(&z_p, 4, 8, 5.0);
write_merged_qkvz(&qkv_p, &z_p, &merged_p).unwrap();
let qkv_len = std::fs::metadata(&qkv_p).unwrap().len();
let z_len = std::fs::metadata(&z_p).unwrap().len();
let expected_size = merged_qkvz_expected_size(qkv_len, z_len).unwrap();
assert!(
merged_qkvz_cache_is_valid(&merged_p, expected_size, &qkv_p, &z_p),
"freshly written merged cache must validate against its own sources"
);
std::fs::remove_dir_all(merged_p.parent().unwrap()).ok();
}
#[test]
fn test_write_merged_qkvz_rejects_oversized_source_payload() {
let (qkv_p, z_p, merged_p) = merge_test_paths("oversized_source");
write_test_q4_source(&qkv_p, 4, 8, 1.0);
write_test_q4_source(&z_p, 4, 8, 5.0);
let f = std::fs::File::options().write(true).open(&qkv_p).unwrap();
f.set_len(36 + MAX_Q4_MERGE_PAYLOAD_LEN + 1).unwrap();
drop(f);
let err = write_merged_qkvz(&qkv_p, &z_p, &merged_p)
.expect_err("oversized source payload must be rejected, not read to EOF");
assert!(
err.contains("payload too large"),
"expected a payload-cap error, got: {err}"
);
assert!(
!merged_p.exists(),
"no merged artifact may be produced from a rejected source"
);
std::fs::remove_dir_all(merged_p.parent().unwrap()).ok();
}
#[test]
fn test_write_merged_qkvz_rejects_rank0_qkv_header_instead_of_panicking() {
let (qkv_p, z_p, merged_p) = merge_test_paths("rank0_qkv");
std::fs::write(
&qkv_p,
q4_file_bytes(&[], 1, q4_f32_to_f16(1.0), q4_f32_to_f16(0.0)),
)
.unwrap();
write_test_q4_source(&z_p, 4, 8, 5.0);
let err = write_merged_qkvz(&qkv_p, &z_p, &merged_p)
.expect_err("rank-0 qkv header must be rejected, not panic");
assert!(
err.contains("not 2-D"),
"expected a 2-D shape error, got: {err}"
);
assert!(
!merged_p.exists(),
"no merged artifact may be produced from a rejected source"
);
std::fs::remove_dir_all(merged_p.parent().unwrap()).ok();
}
#[test]
fn test_write_merged_qkvz_rejects_rank0_z_header_instead_of_panicking() {
let (qkv_p, z_p, merged_p) = merge_test_paths("rank0_z");
write_test_q4_source(&qkv_p, 4, 8, 1.0);
std::fs::write(
&z_p,
q4_file_bytes(&[], 1, q4_f32_to_f16(1.0), q4_f32_to_f16(0.0)),
)
.unwrap();
let err = write_merged_qkvz(&qkv_p, &z_p, &merged_p)
.expect_err("rank-0 z header must be rejected, not panic");
assert!(
err.contains("not 2-D"),
"expected a 2-D shape error, got: {err}"
);
assert!(
!merged_p.exists(),
"no merged artifact may be produced from a rejected source"
);
std::fs::remove_dir_all(merged_p.parent().unwrap()).ok();
}
#[test]
fn test_merged_qkvz_cache_rejects_missing_file() {
let (qkv_p, z_p, merged_p) = merge_test_paths("missing");
write_test_q4_source(&qkv_p, 4, 8, 1.0);
write_test_q4_source(&z_p, 4, 8, 5.0);
let qkv_len = std::fs::metadata(&qkv_p).unwrap().len();
let z_len = std::fs::metadata(&z_p).unwrap().len();
let expected_size = merged_qkvz_expected_size(qkv_len, z_len).unwrap();
assert!(!merged_qkvz_cache_is_valid(
&merged_p,
expected_size,
&qkv_p,
&z_p
));
std::fs::remove_dir_all(merged_p.parent().unwrap()).ok();
}
#[test]
fn test_merged_qkvz_cache_rejects_wrong_size() {
let (qkv_p, z_p, merged_p) = merge_test_paths("wrongsize");
write_test_q4_source(&qkv_p, 4, 8, 1.0);
write_test_q4_source(&z_p, 4, 8, 5.0);
write_merged_qkvz(&qkv_p, &z_p, &merged_p).unwrap();
{
use std::io::Write;
let mut f = std::fs::OpenOptions::new()
.append(true)
.open(&merged_p)
.unwrap();
f.write_all(&[0xAA]).unwrap();
}
let qkv_len = std::fs::metadata(&qkv_p).unwrap().len();
let z_len = std::fs::metadata(&z_p).unwrap().len();
let expected_size = merged_qkvz_expected_size(qkv_len, z_len).unwrap();
assert!(!merged_qkvz_cache_is_valid(
&merged_p,
expected_size,
&qkv_p,
&z_p
));
std::fs::remove_dir_all(merged_p.parent().unwrap()).ok();
}
#[test]
fn test_merged_qkvz_cache_rejects_truncated_file() {
let (qkv_p, z_p, merged_p) = merge_test_paths("truncated");
write_test_q4_source(&qkv_p, 4, 8, 1.0);
write_test_q4_source(&z_p, 4, 8, 5.0);
write_merged_qkvz(&qkv_p, &z_p, &merged_p).unwrap();
let full_len = std::fs::metadata(&merged_p).unwrap().len();
let bytes = std::fs::read(&merged_p).unwrap();
std::fs::write(&merged_p, &bytes[..bytes.len() - 10]).unwrap();
let qkv_len = std::fs::metadata(&qkv_p).unwrap().len();
let z_len = std::fs::metadata(&z_p).unwrap().len();
let expected_size = merged_qkvz_expected_size(qkv_len, z_len).unwrap();
assert_eq!(expected_size, full_len, "sanity: source sizes unchanged");
assert!(!merged_qkvz_cache_is_valid(
&merged_p,
expected_size,
&qkv_p,
&z_p
));
std::fs::remove_dir_all(merged_p.parent().unwrap()).ok();
}
#[test]
fn test_merged_qkvz_cache_rejects_same_size_corrupted_payload() {
let (qkv_p, z_p, merged_p) = merge_test_paths("corrupted");
write_test_q4_source(&qkv_p, 4, 8, 1.0);
write_test_q4_source(&z_p, 4, 8, 5.0);
write_merged_qkvz(&qkv_p, &z_p, &merged_p).unwrap();
let qkv_len = std::fs::metadata(&qkv_p).unwrap().len();
let z_len = std::fs::metadata(&z_p).unwrap().len();
let expected_size = merged_qkvz_expected_size(qkv_len, z_len).unwrap();
let full_len = std::fs::metadata(&merged_p).unwrap().len();
assert_eq!(
full_len, expected_size,
"sanity: size unchanged by corruption"
);
let mut bytes = std::fs::read(&merged_p).unwrap();
let flip_at = bytes.len() - 5;
bytes[flip_at] ^= 0xFF;
std::fs::write(&merged_p, &bytes).unwrap();
assert_eq!(
std::fs::metadata(&merged_p).unwrap().len(),
expected_size,
"sanity: byte flip must not change file size"
);
assert!(
!merged_qkvz_cache_is_valid(&merged_p, expected_size, &qkv_p, &z_p),
"a same-size, bit-flipped merged payload must fail the content-integrity check \
even though the size-only check would have accepted it"
);
std::fs::remove_dir_all(merged_p.parent().unwrap()).ok();
}
#[test]
fn test_merged_qkvz_cache_rejects_same_size_stale_source() {
let (qkv_p, z_p, merged_p) = merge_test_paths("stale");
write_test_q4_source(&qkv_p, 4, 8, 1.0);
write_test_q4_source(&z_p, 4, 8, 5.0);
write_merged_qkvz(&qkv_p, &z_p, &merged_p).unwrap();
let qkv_len = std::fs::metadata(&qkv_p).unwrap().len();
let z_len_before = std::fs::metadata(&z_p).unwrap().len();
write_test_q4_source(&z_p, 4, 8, 99.0);
let z_len_after = std::fs::metadata(&z_p).unwrap().len();
assert_eq!(
z_len_before, z_len_after,
"sanity: same shape must produce the same file size"
);
let expected_size = merged_qkvz_expected_size(qkv_len, z_len_after).unwrap();
assert_eq!(
std::fs::metadata(&merged_p).unwrap().len(),
expected_size,
"sanity: merged file size still matches (source size unchanged)"
);
assert!(
!merged_qkvz_cache_is_valid(&merged_p, expected_size, &qkv_p, &z_p),
"a same-size stale source must invalidate the merged cache once content is checked"
);
std::fs::remove_dir_all(merged_p.parent().unwrap()).ok();
}
#[test]
fn test_merged_qkvz_source_fingerprint_matches_file_fingerprint_after_write() {
let (qkv_p, z_p, merged_p) = merge_test_paths("fingerprint");
write_test_q4_source(&qkv_p, 4, 8, 2.0);
write_test_q4_source(&z_p, 4, 8, 6.0);
write_merged_qkvz(&qkv_p, &z_p, &merged_p).unwrap();
let source_fp = merged_qkvz_source_fingerprint(&qkv_p, &z_p).unwrap();
let file_fp = merged_qkvz_file_fingerprint(&merged_p).unwrap();
assert_eq!(
source_fp, file_fp,
"a freshly written merged file's payload fingerprint must equal its sources' fingerprint"
);
std::fs::remove_dir_all(merged_p.parent().unwrap()).ok();
}
#[test]
fn f16_load_error_display_does_not_duplicate_its_source() {
use std::error::Error;
const CAUSE: &str = "sentinel-cause-text";
let err = F16LoadError::Other(Box::new(std::io::Error::other(CAUSE)));
let shown = format!("{err}");
assert!(
!shown.contains(CAUSE),
"Display must describe only this wrapper's own contribution, but it \
interpolated the wrapped cause: {shown:?}"
);
let source = err.source().expect("Other must keep its cause reachable");
assert!(
format!("{source}").contains(CAUSE),
"source() must yield the wrapped cause itself"
);
}
}