#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub enum Quant {
#[default]
None,
Int8,
Bin,
}
impl Quant {
#[must_use]
pub fn token(self) -> &'static str {
match self {
Quant::None => "f32",
Quant::Int8 => "int8",
Quant::Bin => "bin",
}
}
#[must_use]
pub fn code_bytes(self, dim: usize) -> usize {
match self {
Quant::None => dim * 4,
Quant::Int8 => dim,
Quant::Bin => dim.div_ceil(64) * 8,
}
}
}
#[derive(Debug, Clone)]
pub struct Squeezed {
pub dir: Vec<f32>,
pub norm: f32,
pub range: f32,
}
#[must_use]
pub fn norm(v: &[f32]) -> f32 {
let mut sum = 0.0f32;
let (four, rest) = v.as_chunks::<4>();
for c in four {
let block = c[0].mul_add(c[0], c[1] * c[1]);
let block = c[2].mul_add(c[2], block);
let block = c[3].mul_add(c[3], block);
sum += block;
}
for &x in rest {
sum = x.mul_add(x, sum);
}
sum.sqrt()
}
#[must_use]
pub fn squeeze(quant: Quant, v: &[f32]) -> Squeezed {
let norm = norm(v);
if norm <= 0.0 || !norm.is_finite() {
return Squeezed {
dir: vec![0.0; v.len()],
norm: 0.0,
range: 0.0,
};
}
let mut dir: Vec<f32> = v.iter().map(|x| x / norm).collect();
let range = dir.iter().fold(0.0f32, |wide, x| wide.max(x.abs()));
match quant {
Quant::None => {}
Quant::Int8 => {
let step = 127.0 / range;
for x in &mut dir {
*x = f32::from(eighth(*x, step)) * range / 127.0;
}
}
Quant::Bin => {
#[allow(clippy::cast_precision_loss)]
let w = (v.len() as f32).sqrt().recip();
for x in &mut dir {
*x = if *x > 0.0 { w } else { -w };
}
}
}
Squeezed { dir, norm, range }
}
#[must_use]
pub fn restore(quant: Quant, dir: &[f32], norm: f32) -> Vec<f32> {
match quant {
Quant::Bin => dir
.iter()
.map(|x| if *x > 0.0 { 1.0 } else { -1.0 })
.collect(),
_ => dir.iter().map(|x| x * norm).collect(),
}
}
#[must_use]
pub fn raw(quant: Quant, dir: &[f32], range: f32) -> Vec<u8> {
let mut bytes = Vec::with_capacity(quant.code_bytes(dir.len()));
match quant {
Quant::None => {
for x in dir {
bytes.extend_from_slice(&x.to_le_bytes());
}
}
Quant::Int8 => {
for x in dir {
bytes.push(code(*x, range) as u8);
}
}
Quant::Bin => {
bytes.resize(quant.code_bytes(dir.len()), 0);
for (at, x) in dir.iter().enumerate() {
if *x > 0.0 {
bytes[at / 8] |= 1 << (at % 8);
}
}
}
}
bytes
}
fn code(x: f32, range: f32) -> i8 {
if range <= 0.0 || !range.is_finite() {
return 0;
}
eighth(x, 127.0 / range)
}
#[allow(clippy::cast_possible_truncation)]
fn eighth(x: f32, step: f32) -> i8 {
(x * step).round().clamp(-127.0, 127.0) as i8
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_q8_vector_is_squeezed_the_way_a_real_server_squeezes_it() {
let v = [
-1.057f32, -2.095, 0.906, -2.565, 0.215, -0.806, -2.652, 0.045,
];
let s = squeeze(Quant::Int8, &v);
assert_eq!(s.norm, 4.542_832_4);
assert_eq!(s.range, 0.583_776_8);
let bytes = raw(Quant::Int8, &s.dir, s.range);
assert_eq!(bytes, [0xcd, 0x9c, 0x2b, 0x85, 0x0a, 0xd9, 0x81, 0x02]);
let back = restore(Quant::Int8, &s.dir, s.norm);
assert_eq!(back[0], -1.064_976_5);
assert_eq!(back[7], 0.041_763_78);
}
#[test]
fn the_length_is_the_length_a_real_server_measures() {
let six = [
-0.937_271_f32,
-0.990_583_06,
-3.973_563_2,
-3.560_724_7,
-4.341_83,
-2.311_353_2,
];
assert_eq!(norm(&six), 7.383_869_6);
let eight = [
3.911_460_9_f32,
-4.397_685_5,
1.571_724_4,
1.238_847_4,
4.993_485_5,
4.326_879,
-0.686_563_8,
2.032_343_4,
];
assert_eq!(norm(&eight), 9.322_167);
assert_eq!(norm(&[3.0, 4.0]), 5.0);
}
#[test]
fn nothing_is_lost_when_nothing_is_squeezed() {
let v = [3.0f32, 4.0];
let s = squeeze(Quant::None, &v);
assert_eq!(s.norm, 5.0);
assert_eq!(restore(Quant::None, &s.dir, s.norm), [3.0, 4.0]);
assert_eq!(raw(Quant::None, &s.dir, s.range).len(), 8);
}
#[test]
fn a_binary_vector_is_its_signs_and_keeps_no_length() {
let v = [1.0f32, -2.0, 0.0, 4.0];
let s = squeeze(Quant::Bin, &v);
assert_eq!(restore(Quant::Bin, &s.dir, s.norm), [1.0, -1.0, -1.0, 1.0]);
assert_eq!(
raw(Quant::Bin, &s.dir, s.range),
[0b1001, 0, 0, 0, 0, 0, 0, 0]
);
}
#[test]
fn a_squeezed_direction_is_still_a_direction() {
let v = [0.3f32, -1.7, 2.2, 0.9, -0.4];
for quant in [Quant::None, Quant::Int8, Quant::Bin] {
let s = squeeze(quant, &v);
let len = norm(&s.dir);
assert!((len - 1.0).abs() < 1e-3, "{} is {len} long", quant.token());
}
}
#[test]
fn a_vector_of_no_length_is_stored_as_the_origin() {
let s = squeeze(Quant::Int8, &[0.0, 0.0, 0.0]);
assert_eq!(s.norm, 0.0);
assert_eq!(s.dir, [0.0, 0.0, 0.0]);
assert_eq!(raw(Quant::Int8, &s.dir, s.range), [0, 0, 0]);
}
#[test]
fn how_wide_the_bytes_are_is_how_wide_they_turn_out_to_be() {
for dim in [1usize, 7, 8, 63, 64, 65, 300] {
let v = vec![0.5f32; dim];
for quant in [Quant::None, Quant::Int8, Quant::Bin] {
let s = squeeze(quant, &v);
let bytes = raw(quant, &s.dir, s.range);
assert_eq!(
bytes.len(),
quant.code_bytes(dim),
"{} at {dim}",
quant.token()
);
}
}
}
}