const PER_BYTE: usize = 2;
const BIAS: i32 = 8;
#[derive(Clone, Debug, PartialEq)]
pub struct QuantizedMatrixQ4 {
pub data: Vec<u8>,
pub scales: Vec<f32>,
pub n: usize,
pub k: usize,
}
impl QuantizedMatrixQ4 {
#[must_use]
pub fn quantize(weight: &[f32], n: usize, k: usize) -> Self {
assert_eq!(weight.len(), n * k, "weight must be [n, k]");
let packed_row = k.div_ceil(PER_BYTE);
let mut data = vec![0_u8; n * packed_row];
let mut scales = Vec::with_capacity(n);
for row in 0..n {
let source = &weight[row * k..row * k + k];
let mut maximum = 0.0_f32;
for (index, &value) in source.iter().enumerate() {
assert!(
value.is_finite(),
"non-finite value {value} at index {index} reached the Q4 quantizer"
);
maximum = maximum.max(value.abs());
}
let scale = if maximum == 0.0 { 0.0 } else { maximum / 7.0 };
scales.push(if scale == 0.0 { 1.0 } else { scale });
let target = &mut data[row * packed_row..(row + 1) * packed_row];
if scale == 0.0 {
target.fill(((BIAS as u8) << 4) | BIAS as u8);
continue;
}
for (index, &value) in source.iter().enumerate() {
let level = (value / scale).clamp(-7.0, 7.0).round_ties_even() as i32;
let biased = (level + BIAS) as u8;
let byte = &mut target[index / PER_BYTE];
if index % PER_BYTE == 0 {
*byte = (*byte & 0xF0) | biased;
} else {
*byte = (*byte & 0x0F) | (biased << 4);
}
}
}
Self { data, scales, n, k }
}
#[must_use]
pub fn dequantize_row(&self, row: usize) -> Vec<f32> {
let packed_row = self.k.div_ceil(PER_BYTE);
let bytes = &self.data[row * packed_row..(row + 1) * packed_row];
let scale = self.scales[row];
(0..self.k)
.map(|index| {
let byte = bytes[index / PER_BYTE];
let nibble = if index % PER_BYTE == 0 {
i32::from(byte & 0x0F)
} else {
i32::from(byte >> 4)
};
#[allow(clippy::cast_precision_loss)]
{
(nibble - BIAS) as f32 * scale
}
})
.collect()
}
#[must_use]
pub fn packed_bytes(&self) -> usize {
self.data.len()
}
}
#[must_use]
pub fn dot_i32_q4(x: &[i8], packed: &[u8], k: usize) -> i32 {
assert!(
packed.len() >= k.div_ceil(PER_BYTE),
"packed row shorter than k nibbles"
);
assert!(x.len() >= k, "activation row shorter than k");
let mut unsigned_accumulator = 0_i32;
let mut activation_sum = 0_i32;
let pairs = k / PER_BYTE;
for pair in 0..pairs {
let byte = packed[pair];
let low = i32::from(byte & 0x0F);
let high = i32::from(byte >> 4);
let first = i32::from(x[pair * PER_BYTE]);
let second = i32::from(x[pair * PER_BYTE + 1]);
unsigned_accumulator += first * low + second * high;
activation_sum += first + second;
}
if k % PER_BYTE == 1 {
let byte = packed[pairs];
let value = i32::from(x[k - 1]);
unsigned_accumulator += value * i32::from(byte & 0x0F);
activation_sum += value;
}
unsigned_accumulator - BIAS * activation_sum
}
pub fn linear_q4(
x_q: &[i8],
x_scales: &[f32],
weight: &QuantizedMatrixQ4,
bias: Option<&[f32]>,
m: usize,
out: &mut [f32],
) {
let (n, k) = (weight.n, weight.k);
assert_eq!(x_q.len(), m * k, "activations must be [m, k]");
assert_eq!(x_scales.len(), m, "one activation scale per row");
assert_eq!(out.len(), m * n, "out must be [m, n]");
let packed_row = k.div_ceil(PER_BYTE);
for row in 0..m {
let x_row = &x_q[row * k..row * k + k];
for column in 0..n {
let w_row = &weight.data[column * packed_row..(column + 1) * packed_row];
let accumulated = dot_i32_q4(x_row, w_row, k);
#[allow(clippy::cast_precision_loss)]
let value = accumulated as f32 * (x_scales[row] * weight.scales[column]);
out[row * n + column] = bias.map_or(value, |values| value + values[column]);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::int8::quantize_row_q8;
fn deterministic(count: usize, seed: u64) -> Vec<f32> {
let mut state = seed | 1;
(0..count)
.map(|_| {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
((state >> 40) as f32 / 8192.0) - 0.5
})
.collect()
}
#[test]
fn bias_cancellation_is_exact_against_a_signed_reference() {
for (index, &k) in [1_usize, 2, 3, 7, 8, 15, 64, 127, 1024].iter().enumerate() {
let weight = deterministic(k, 0x4B1A_0000 + index as u64);
let matrix = QuantizedMatrixQ4::quantize(&weight, 1, k);
let activation = deterministic(k, 0xA0C7_0000 + index as u64);
let mut x_q = vec![0_i8; k];
quantize_row_q8(&activation, &mut x_q);
let packed_row = k.div_ceil(PER_BYTE);
let mut expected = 0_i32;
for (position, &activation) in x_q.iter().enumerate().take(k) {
let byte = matrix.data[position / PER_BYTE];
let nibble = if position % PER_BYTE == 0 {
i32::from(byte & 0x0F)
} else {
i32::from(byte >> 4)
};
expected += i32::from(activation) * (nibble - BIAS);
}
assert_eq!(
dot_i32_q4(&x_q, &matrix.data[..packed_row], k),
expected,
"k={k}: biased accumulation with a single correction diverged from the signed dot"
);
}
}
#[test]
fn levels_are_symmetric_and_never_emit_negative_eight() {
let k = 4096;
let mut weight = deterministic(k, 0xD00D);
weight[0] = 1.0;
weight[1] = -1.0;
weight[2] = 0.0;
let matrix = QuantizedMatrixQ4::quantize(&weight, 1, k);
let scale = matrix.scales[0];
for position in 0..k {
let byte = matrix.data[position / PER_BYTE];
let nibble = if position % PER_BYTE == 0 {
i32::from(byte & 0x0F)
} else {
i32::from(byte >> 4)
};
let level = nibble - BIAS;
assert!(
(-7..=7).contains(&level),
"level {level} outside the symmetric range at {position}"
);
}
let negated: Vec<f32> = weight.iter().map(|value| -value).collect();
let mirror = QuantizedMatrixQ4::quantize(&negated, 1, k);
assert!((mirror.scales[0] - scale).abs() <= f32::EPSILON * scale.max(1.0));
let forward = matrix.dequantize_row(0);
let backward = mirror.dequantize_row(0);
for position in 0..k {
assert!(
forward[position] == -backward[position],
"negation was not exact at {position}: {} vs {}",
forward[position],
backward[position]
);
}
}
#[test]
fn packed_storage_is_half_of_q8() {
let (n, k) = (2048, 1024);
let weight = deterministic(n * k, 0xFEED);
let q4 = QuantizedMatrixQ4::quantize(&weight, n, k);
assert_eq!(q4.packed_bytes(), n * k / 2);
let q8 = crate::int8::QuantizedMatrix::quantize(&weight, n, k);
assert_eq!(q4.packed_bytes() * 2, q8.data.len());
}
#[test]
fn quantization_error_is_bounded_by_the_level_step() {
let k = 8192;
let weight = deterministic(k, 0xBEEF);
let matrix = QuantizedMatrixQ4::quantize(&weight, 1, k);
let restored = matrix.dequantize_row(0);
let scale = matrix.scales[0];
for (position, (&original, &back)) in weight.iter().zip(restored.iter()).enumerate() {
assert!(
(original - back).abs() <= scale * 0.5 + f32::EPSILON * 8.0,
"position {position}: |{original} - {back}| exceeds half a level ({scale})"
);
}
}
}