Skip to main content

urna_format/encoding/
int4.rs

1//! int4 block-64 quantized embeddings (`encoding=7`).
2//!
3//! On disk:
4//! ```text
5//!   u32 LE  payload_version = 1
6//!   u32 LE  scale_kind      = 1  (per-group, block-64)
7//!   f16 LE * (n * dim/64)   per-group block absmax scales, row-major
8//!   u8     * (n * dim/2)    packed 4-bit signed codes, two nibbles/byte
9//! ```
10//!
11//! Each row is split into `dim/64` contiguous 64-dim blocks; `dim` must be
12//! divisible by 64. Every block carries one f16 absmax scale, so an outlier
13//! in one block cannot crush another (per-group idea, the int8 per-vector
14//! step taken finer). A code `c` in `[-7, 7]` reconstructs as
15//! `f32_value ~= c * scale`; the range is symmetric and reserves `-8`.
16//!
17//! STORED-PRECISION codec: like int8, the returned cosine is real AT THIS
18//! PRECISION (disclosed via the dtype). It never zstds/shuffles, so the
19//! runtime scores it off mmap with the fused dequant+dot kernel.
20
21use crate::bytes::le_u32;
22use crate::error::UrnaError;
23
24pub const INT4_PAYLOAD_VERSION: u32 = 1;
25pub const INT4_SCALE_KIND_PER_GROUP: u32 = 1;
26pub const INT4_PREFIX_SIZE: usize = 8;
27/// Block size for the per-group absmax scale. `dim` must be a multiple.
28pub const INT4_BLOCK: usize = 64;
29
30/// Number of 64-dim blocks per row. `dim` must be divisible by `INT4_BLOCK`.
31#[inline]
32pub fn int4_blocks_per_row(dim: usize) -> usize {
33    dim / INT4_BLOCK
34}
35
36/// Quantize one L2-normalized f32 row to int4 with per-64-block absmax
37/// scales. Returns `(scales, codes)` where `scales[g]` is the f16 absmax
38/// scale of block `g` and `codes[j]` in `[-7, 7]` reconstructs as
39/// `codes[j] as f32 * scales[j / 64]`. Mirrors `quantize_f32_to_i8` but per
40/// 64-dim group; a zero block maps to all-zero codes with scale 1.
41pub fn quantize_f32_to_i4(values: &[f32], dim: usize) -> (Vec<half::f16>, Vec<i8>) {
42    let blocks = int4_blocks_per_row(dim);
43    let mut scales: Vec<half::f16> = Vec::with_capacity(blocks);
44    let mut codes: Vec<i8> = Vec::with_capacity(dim);
45    for g in 0..blocks {
46        let blk = &values[g * INT4_BLOCK..(g + 1) * INT4_BLOCK];
47        let max_abs = blk.iter().fold(0.0f32, |acc, &v| acc.max(v.abs()));
48        // Quantize against the f16-rounded scale so the codes match the
49        // stored (f16) scale the reader sees, keeping round-trip tight.
50        let scale_f16 = half::f16::from_f32(if max_abs == 0.0 { 1.0 } else { max_abs / 7.0 });
51        let scale = scale_f16.to_f32();
52        let inv = if scale > 0.0 { 1.0 / scale } else { 0.0 };
53        for &v in blk {
54            let q = (v * inv).round().clamp(-7.0, 7.0);
55            codes.push(q as i8);
56        }
57        scales.push(scale_f16);
58    }
59    (scales, codes)
60}
61
62/// Pack signed 4-bit codes (`[-7, 7]`) into bytes, two nibbles per byte,
63/// low nibble first. The nibble is the two's-complement low 4 bits, so
64/// `-7..=7` maps to `0x9..=0x7` and unpacks back exactly via sign extension.
65#[inline]
66pub fn pack_nibbles(codes: &[i8]) -> Vec<u8> {
67    let mut out = Vec::with_capacity(codes.len().div_ceil(2));
68    for pair in codes.chunks(2) {
69        let lo = (pair[0] as u8) & 0x0F;
70        let hi = pair.get(1).map(|&c| (c as u8) & 0x0F).unwrap_or(0);
71        out.push(lo | (hi << 4));
72    }
73    out
74}
75
76/// Sign-extend a 4-bit nibble (low 4 bits of `b`) to an `i8` in `[-8, 7]`.
77#[inline]
78pub fn nibble_to_i4(b: u8) -> i8 {
79    let n = b & 0x0F;
80    if n & 0x08 != 0 {
81        (n | 0xF0) as i8
82    } else {
83        n as i8
84    }
85}
86
87/// Encode the int4 embeddings section payload. `embeddings` is `n * dim`
88/// row-major f32; `dim` must be divisible by `INT4_BLOCK`.
89pub fn encode_int4_embeddings(embeddings: &[f32], n: usize, dim: usize) -> crate::Result<Vec<u8>> {
90    if dim == 0 || dim % INT4_BLOCK != 0 {
91        return Err(UrnaError::InvalidInput(format!(
92            "encode_int4_embeddings: dim={dim} must be a nonzero multiple of {INT4_BLOCK}"
93        )));
94    }
95    if embeddings.len() != n * dim {
96        return Err(UrnaError::InvalidInput(format!(
97            "encode_int4_embeddings: got {} f32 values for n={n} dim={dim}",
98            embeddings.len()
99        )));
100    }
101    let blocks = int4_blocks_per_row(dim);
102    let mut out = Vec::with_capacity(INT4_PREFIX_SIZE + n * blocks * 2 + n * dim / 2);
103    out.extend_from_slice(&INT4_PAYLOAD_VERSION.to_le_bytes());
104    out.extend_from_slice(&INT4_SCALE_KIND_PER_GROUP.to_le_bytes());
105    let mut scale_bytes: Vec<u8> = Vec::with_capacity(n * blocks * 2);
106    let mut body: Vec<u8> = Vec::with_capacity(n * dim / 2);
107    for i in 0..n {
108        let row = &embeddings[i * dim..(i + 1) * dim];
109        let (scales, codes) = quantize_f32_to_i4(row, dim);
110        for s in &scales {
111            scale_bytes.extend_from_slice(&s.to_le_bytes());
112        }
113        body.extend_from_slice(&pack_nibbles(&codes));
114    }
115    out.extend_from_slice(&scale_bytes);
116    out.extend_from_slice(&body);
117    Ok(out)
118}
119
120/// Decoded view over an int4 embeddings payload. Slices borrow the input
121/// bytes (no copy); accessors decode scales/codes on demand.
122pub struct Int4EmbeddingsView<'a> {
123    /// f16 LE group scales, `n * blocks` of them, row-major.
124    pub scales: &'a [u8],
125    /// packed nibble codes, `n * dim/2` bytes.
126    pub codes: &'a [u8],
127    pub n: usize,
128    pub dim: usize,
129    pub blocks: usize,
130}
131
132impl<'a> Int4EmbeddingsView<'a> {
133    pub fn parse(bytes: &'a [u8], n: usize, dim: usize) -> crate::Result<Self> {
134        if dim == 0 || dim % INT4_BLOCK != 0 {
135            return Err(UrnaError::MalformedSectionPayload {
136                section_id: crate::layout::SECTION_EMBEDDINGS,
137                reason: format!("int4 dim={dim} must be a nonzero multiple of {INT4_BLOCK}"),
138            });
139        }
140        let blocks = int4_blocks_per_row(dim);
141        // checked: `n` / `dim` are header-controlled; an overflowed product
142        // must be a typed mismatch, never a wrapped "match".
143        let want = super::expected_embeddings_size("int4", n, dim).unwrap_or(usize::MAX);
144        if bytes.len() != want {
145            return Err(UrnaError::EmbeddingSizeMismatch {
146                expected: want,
147                got: bytes.len(),
148            });
149        }
150        let version = le_u32(&bytes[0..4])?;
151        if version != INT4_PAYLOAD_VERSION {
152            return Err(UrnaError::UnsupportedSectionVersion {
153                section_id: crate::layout::SECTION_EMBEDDINGS,
154                version,
155            });
156        }
157        let kind = le_u32(&bytes[4..8])?;
158        if kind != INT4_SCALE_KIND_PER_GROUP {
159            return Err(UrnaError::MalformedSectionPayload {
160                section_id: crate::layout::SECTION_EMBEDDINGS,
161                reason: format!("int4 scale_kind {kind} not supported"),
162            });
163        }
164        let scales_end = INT4_PREFIX_SIZE + n * blocks * 2;
165        Ok(Self {
166            scales: &bytes[INT4_PREFIX_SIZE..scales_end],
167            codes: &bytes[scales_end..],
168            n,
169            dim,
170            blocks,
171        })
172    }
173
174    /// Read the f16 group scale `g` of row `i` as f32.
175    #[inline]
176    pub fn group_scale(&self, i: usize, g: usize) -> f32 {
177        let off = (i * self.blocks + g) * 2;
178        half::f16::from_le_bytes([self.scales[off], self.scales[off + 1]]).to_f32()
179    }
180
181    /// Borrow row `i`'s packed nibble bytes (`dim/2` of them).
182    #[inline]
183    pub fn row_codes(&self, i: usize) -> &'a [u8] {
184        let rs = self.dim / 2;
185        let start = i * rs;
186        &self.codes[start..start + rs]
187    }
188
189    /// Decode row `i`'s f16 group scales into a fresh `Vec<f32>` (one per
190    /// 64-dim block). One-shot use only (the packed ANN store materializes
191    /// every row once); the per-candidate rerank uses `row_scales_into`.
192    #[inline]
193    pub fn row_scales_f32(&self, i: usize) -> Vec<f32> {
194        (0..self.blocks).map(|g| self.group_scale(i, g)).collect()
195    }
196
197    /// Decode row `i`'s f16 group scales into a caller-owned buffer of
198    /// exactly `blocks` f32s, so a rerank over thousands of candidates
199    /// reuses one allocation instead of one `Vec` per row.
200    ///
201    /// # Panics
202    ///
203    /// When `out.len() != self.blocks`.
204    #[inline]
205    pub fn row_scales_into(&self, i: usize, out: &mut [f32]) {
206        assert_eq!(out.len(), self.blocks, "int4: one scale slot per block");
207        for (g, slot) in out.iter_mut().enumerate() {
208            *slot = self.group_scale(i, g);
209        }
210    }
211}
212
213// Unit + round-trip coverage lives in `tests/int4_roundtrip.rs` (pack /
214// unpack, quantize clamping, the section view, and the typed malformed-
215// payload rejections), so this file holds the codec alone. Negative file-level paths live in `tests/negative_int4.rs`.