use anyhow::{anyhow, Result};
use mlx_native::{DType, MlxBuffer, MlxDevice};
const QK8_0: usize = 32;
const BLOCK_Q8_0_BYTES: usize = 2 + QK8_0;
pub fn quantize_f32_to_q8_0_buffer(
f32_data: &[f32],
shape: Vec<usize>,
device: &MlxDevice,
) -> Result<MlxBuffer> {
let numel = f32_data.len();
if numel == 0 {
return Err(anyhow!("quantize_f32_to_q8_0_buffer: empty input"));
}
if numel % QK8_0 != 0 {
return Err(anyhow!(
"quantize_f32_to_q8_0_buffer: numel={} not divisible by Q8_0 block size {}",
numel,
QK8_0
));
}
let shape_numel: usize = shape.iter().product();
if shape_numel != numel {
return Err(anyhow!(
"quantize_f32_to_q8_0_buffer: shape product {} != numel {}",
shape_numel,
numel
));
}
let num_blocks = numel / QK8_0;
let total_bytes = num_blocks * BLOCK_Q8_0_BYTES;
let mut buf = device
.alloc_buffer(total_bytes, DType::U8, shape)
.map_err(|e| anyhow!("quantize_f32_to_q8_0_buffer: alloc_buffer: {e}"))?;
super::weight_pool::register_weight_buffer(device, &buf)
.map_err(|e| anyhow!("quantize_f32_to_q8_0_buffer: register_weight_buffer: {e}"))?;
{
let dst: &mut [u8] = buf
.as_mut_slice()
.map_err(|e| anyhow!("quantize_f32_to_q8_0_buffer: as_mut_slice: {e}"))?;
for block_idx in 0..num_blocks {
let block = &f32_data[block_idx * QK8_0..(block_idx + 1) * QK8_0];
let absmax = block.iter().fold(0.0f32, |acc, &x| acc.max(x.abs()));
let d = absmax / 127.0;
let inv_d = if d == 0.0 { 0.0f32 } else { 1.0 / d };
let block_off = block_idx * BLOCK_Q8_0_BYTES;
let scale_f16 = half::f16::from_f32(d);
let scale_bytes = scale_f16.to_le_bytes();
dst[block_off] = scale_bytes[0];
dst[block_off + 1] = scale_bytes[1];
for (i, &val) in block.iter().enumerate() {
let q = (val * inv_d).round().clamp(-128.0, 127.0) as i32;
dst[block_off + 2 + i] = (q as i8) as u8;
}
}
}
Ok(buf)
}
pub fn bf16_bytes_to_f32(src: &[u8], dst: &mut Vec<f32>) {
debug_assert_eq!(src.len() % 2, 0, "BF16 byte length must be even");
dst.clear();
dst.reserve(src.len() / 2);
for chunk in src.chunks_exact(2) {
let bits = u16::from_le_bytes([chunk[0], chunk[1]]);
let f32_bits = (bits as u32) << 16;
dst.push(f32::from_bits(f32_bits));
}
}
pub fn f16_bytes_to_f32(src: &[u8], dst: &mut Vec<f32>) {
debug_assert_eq!(src.len() % 2, 0, "F16 byte length must be even");
dst.clear();
dst.reserve(src.len() / 2);
for chunk in src.chunks_exact(2) {
let bits = u16::from_le_bytes([chunk[0], chunk[1]]);
dst.push(half::f16::from_bits(bits).to_f32());
}
}
pub fn f32_bytes_to_f32(src: &[u8], dst: &mut Vec<f32>) {
debug_assert_eq!(src.len() % 4, 0, "F32 byte length must be divisible by 4");
dst.clear();
dst.reserve(src.len() / 4);
for chunk in src.chunks_exact(4) {
dst.push(f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]));
}
}
#[cfg(test)]
mod tests {
use super::*;
use mlx_native::MlxDevice;
#[test]
fn q8_0_round_trip_zero_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let zeros = vec![0.0f32; 32];
let buf = quantize_f32_to_q8_0_buffer(&zeros, vec![32], &device).expect("quantize");
let bytes = buf.as_slice::<u8>().expect("slice");
assert_eq!(bytes.len(), 34);
assert_eq!(bytes[0], 0);
assert_eq!(bytes[1], 0);
for i in 0..32 {
assert_eq!(bytes[2 + i], 0);
}
}
#[test]
fn q8_0_block_round_trip_within_tolerance() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let mut input = Vec::with_capacity(32);
for i in 0..32 {
input.push((i as f32) - 15.5);
}
let buf = quantize_f32_to_q8_0_buffer(&input, vec![32], &device).expect("quantize");
let bytes = buf.as_slice::<u8>().expect("slice");
let scale_bits = u16::from_le_bytes([bytes[0], bytes[1]]);
let scale = half::f16::from_bits(scale_bits).to_f32();
let mut max_err = 0.0f32;
for i in 0..32 {
let q = bytes[2 + i] as i8;
let recon = (q as f32) * scale;
let err = (recon - input[i]).abs();
if err > max_err {
max_err = err;
}
}
assert!(max_err < 0.07, "max_err {} >= 0.07", max_err);
}
#[test]
fn q8_0_rejects_unaligned_numel() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let unaligned = vec![1.0f32; 31];
let res = quantize_f32_to_q8_0_buffer(&unaligned, vec![31], &device);
assert!(res.is_err());
}
#[test]
fn q8_0_rejects_shape_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = MlxDevice::new().expect("device");
let v = vec![1.0f32; 32];
let res = quantize_f32_to_q8_0_buffer(&v, vec![16, 4], &device);
assert!(res.is_err());
}
#[test]
fn bf16_to_f32_round_trip() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let bytes: Vec<u8> = vec![0x80, 0x3F, 0x00, 0x40]; let mut out = Vec::new();
bf16_bytes_to_f32(&bytes, &mut out);
assert_eq!(out.len(), 2);
assert_eq!(out[0], 1.0);
assert_eq!(out[1], 2.0);
}
#[test]
fn f32_bytes_round_trip() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let f32s = [1.5f32, -2.25, 100.0, 0.0];
let mut bytes = Vec::new();
for v in &f32s {
bytes.extend_from_slice(&v.to_le_bytes());
}
let mut out = Vec::new();
f32_bytes_to_f32(&bytes, &mut out);
assert_eq!(out, f32s);
}
}