Skip to main content

urna_format/encoding/
int8.rs

1//! int8 quantized embeddings (`encoding=3`).
2//!
3//! On disk:
4//! ```text
5//!   u32 LE  payload_version = 1
6//!   u32 LE  scale_kind      = 0  (per-vector f32)
7//!   f32 LE * n              scales (one per vector)
8//!   i8     * (n * dim)      quantized embeddings, row-major
9//! ```
10//!
11//! The scale is the multiplier such that `f32_value ≈ i8_value * scale`.
12//! We pick per-vector scales so a single outlier vector cannot crush
13//! resolution for the whole corpus.
14
15use crate::bytes::le_u32;
16use crate::error::UrnaError;
17
18pub const INT8_PAYLOAD_VERSION: u32 = 1;
19pub const INT8_SCALE_KIND_PER_VECTOR: u32 = 0;
20pub const INT8_PREFIX_SIZE: usize = 8;
21
22/// Quantize an L2-normalized float32 embedding to int8 with a per-vector
23/// scale. Returns `(scale, i8_bytes)` where `f32_value ≈ i8 * scale`.
24///
25/// L2-normalized vectors live in `[-1, 1]`; in practice the largest
26/// component is well below 1 so we map `max(|v|)` to 127 to use the
27/// full int8 range. Re-normalization at query time accumulates in f32.
28pub fn quantize_f32_to_i8(values: &[f32]) -> (f32, Vec<i8>) {
29    let max_abs = values.iter().fold(0.0f32, |acc, &v| acc.max(v.abs()));
30    if max_abs == 0.0 {
31        // Pathological zero vector - quantize to all zeros with scale 1.
32        // The reader's zero-norm guard will reject queries against this.
33        return (1.0, vec![0i8; values.len()]);
34    }
35    let scale = max_abs / 127.0;
36    let inv_scale = 1.0 / scale;
37    let q: Vec<i8> = values
38        .iter()
39        .map(|&v| {
40            let scaled = (v * inv_scale).round();
41            scaled.clamp(-127.0, 127.0) as i8
42        })
43        .collect();
44    (scale, q)
45}
46
47/// Encode the int8 embeddings section payload. Layout matches
48/// `INT8_PAYLOAD_VERSION` / `INT8_SCALE_KIND_PER_VECTOR`.
49///
50/// `embeddings` is `n * dim` row-major f32 values. Returns a buffer
51/// ready to write as the embeddings section (encoding = INT8).
52pub fn encode_int8_embeddings(embeddings: &[f32], n: usize, dim: usize) -> crate::Result<Vec<u8>> {
53    if embeddings.len() != n * dim {
54        return Err(UrnaError::InvalidInput(format!(
55            "encode_int8_embeddings: got {} f32 values for n={} dim={}",
56            embeddings.len(),
57            n,
58            dim
59        )));
60    }
61    let mut out = Vec::with_capacity(INT8_PREFIX_SIZE + n * 4 + n * dim);
62    out.extend_from_slice(&INT8_PAYLOAD_VERSION.to_le_bytes());
63    out.extend_from_slice(&INT8_SCALE_KIND_PER_VECTOR.to_le_bytes());
64    let mut scales: Vec<u8> = Vec::with_capacity(n * 4);
65    let mut bodies: Vec<u8> = Vec::with_capacity(n * dim);
66    for i in 0..n {
67        let row = &embeddings[i * dim..(i + 1) * dim];
68        let (scale, q) = quantize_f32_to_i8(row);
69        scales.extend_from_slice(&scale.to_le_bytes());
70        // i8 -> u8 is a bitcast (two's complement preserved).
71        bodies.extend(q.iter().map(|&v| v as u8));
72    }
73    out.extend_from_slice(&scales);
74    out.extend_from_slice(&bodies);
75    Ok(out)
76}
77
78/// Decoded view over an int8 embeddings payload. The slices borrow
79/// from the input bytes (no copy).
80pub struct Int8EmbeddingsView<'a> {
81    pub scales: &'a [u8], // n * 4 bytes (f32 LE)
82    pub bodies: &'a [u8], // n * dim bytes (i8, bitcast to u8)
83    pub n: usize,
84    pub dim: usize,
85}
86
87impl<'a> Int8EmbeddingsView<'a> {
88    pub fn parse(bytes: &'a [u8], n: usize, dim: usize) -> crate::Result<Self> {
89        // checked: `n` / `dim` are header-controlled; an overflowed product
90        // must be a typed mismatch, never a wrapped "match".
91        let want = super::expected_embeddings_size("int8", n, dim).unwrap_or(usize::MAX);
92        if bytes.len() != want {
93            return Err(UrnaError::EmbeddingSizeMismatch {
94                expected: want,
95                got: bytes.len(),
96            });
97        }
98        let version = le_u32(&bytes[0..4])?;
99        if version != INT8_PAYLOAD_VERSION {
100            return Err(UrnaError::UnsupportedSectionVersion {
101                section_id: crate::layout::SECTION_EMBEDDINGS,
102                version,
103            });
104        }
105        let kind = le_u32(&bytes[4..8])?;
106        if kind != INT8_SCALE_KIND_PER_VECTOR {
107            return Err(UrnaError::MalformedSectionPayload {
108                section_id: crate::layout::SECTION_EMBEDDINGS,
109                reason: format!("int8 scale_kind {} not supported", kind),
110            });
111        }
112        let scales_end = INT8_PREFIX_SIZE + n * 4;
113        Ok(Self {
114            scales: &bytes[INT8_PREFIX_SIZE..scales_end],
115            bodies: &bytes[scales_end..],
116            n,
117            dim,
118        })
119    }
120
121    /// Read scale[i] as f32.
122    #[inline]
123    pub fn scale(&self, i: usize) -> f32 {
124        let off = i * 4;
125        f32::from_le_bytes([
126            self.scales[off],
127            self.scales[off + 1],
128            self.scales[off + 2],
129            self.scales[off + 3],
130        ])
131    }
132
133    /// Borrow row[i] as `&[i8]`.
134    #[inline]
135    pub fn row(&self, i: usize) -> &'a [i8] {
136        let start = i * self.dim;
137        let end = start + self.dim;
138        // i8 has the same size and alignment as u8; bytemuck checks both at
139        // compile time, so this reinterpretation carries no `unsafe`.
140        bytemuck::cast_slice(&self.bodies[start..end])
141    }
142}