Skip to main content

kernel/
vecquant.rs

1//! 2g: data-oblivious vector fingerprints (TurboQuant-style, arXiv 2504.19874).
2//!
3//! Beginner's map of what happens and why it needs NO training:
4//!
5//! 1. ROTATE the vector by a fixed random rotation. A rotation preserves all
6//!    distances and dot products, but after it every coordinate of a unit
7//!    vector looks like a small Gaussian: sigma = 1/sqrt(dim), regardless of
8//!    what the data means. That statistical guarantee is the whole trick --
9//!    it replaces the per-dataset codebook training PQ needs.
10//! 2. QUANTIZE each rotated coordinate independently against a FIXED ladder
11//!    of 2^BITS levels sized for that Gaussian. The recipe depends only on
12//!    (dimension, seed); both live in the catalog, so any process can encode
13//!    or decode any vector at any time, forever.
14//! 3. Store the original LENGTH (one f32): the code approximates the vector's
15//!    direction; the norm restores its scale.
16//!
17//! The rotation is a fast Walsh-Hadamard transform (FWHT) with seeded random
18//! sign flips, ROUNDS times: O(d log d) instead of a d*d matrix multiply,
19//! and zero bytes of stored matrix. FWHT needs a power-of-two width, so
20//! vectors are zero-padded up to one (1536 -> 2048).
21//!
22//! Search math: dot(x, q) = |x| * dot(unit_code(x), rotate(q)) because
23//! rotations preserve dots. One estimated dot yields L2, cosine and dot
24//! scores alike. The estimate is APPROXIMATE -- callers oversample and
25//! rescore against the exact f32 rows (D23), so approximation can only ever
26//! cost a MISS, never a wrong distance in the final ranking.
27
28/// Default code width. 4-bit = 16 levels, dim/2 bytes; 2-bit = 4 levels,
29/// dim/4 bytes -- half the scan I/O and half the unpack work, the lever
30/// that matters when the code keyspace outgrows the pool. Recorded per
31/// store in the catalog; the recall gate decides the default.
32pub const DEFAULT_BITS: usize = 2;
33pub const MAX_LEVELS: usize = 16;
34const ROUNDS: usize = 1;
35/// Quantizer span in sigmas: rotated unit-vector coordinates scaled by
36/// sqrt(d) are ~N(0,1); +-2.5 covers 98.8% of mass, the tails clamp.
37const SPAN: f32 = 2.5;
38
39/// Bytes of code per vector for a padded width (norm f32 NOT included).
40pub fn code_len(padded: usize, bits: usize) -> usize { padded * bits / 8 }
41
42pub fn pad_dim(dim: usize) -> usize {
43    // Floor at 64: the estimate kernels consume whole 16-byte lanes, and
44    // padded*bits/8 TRUNCATED TO ZERO for dim <= 2 at 2 bits -- every
45    // set_vec of a tiny vector then indexed an empty code buffer and
46    // panicked (found by e1's dim-2 integration tests). A 64-floor makes
47    // every code an exact lane multiple at 1/2/4 bits.
48    dim.next_power_of_two().max(64)
49}
50
51/// Split a u64 seed into one sign per (round, coordinate) via a tiny PRNG.
52fn signs(seed: u64, round: usize, padded: usize) -> Vec<f32> {
53    let mut s = seed ^ (round as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15);
54    (0..padded).map(|_| {
55        s ^= s << 13; s ^= s >> 7; s ^= s << 17;
56        if s & 1 == 1 { 1.0 } else { -1.0 }
57    }).collect()
58}
59
60/// In-place FWHT. Classic butterfly; unnormalised (we fold the 1/sqrt(n)
61/// into a single scale at the end of the rotation).
62fn fwht(v: &mut [f32]) {
63    let n = v.len();
64    let mut h = 1;
65    while h < n {
66        for block in v.chunks_exact_mut(2 * h) {
67            let (a, b) = block.split_at_mut(h);
68            // contiguous halves: this loop auto-vectorizes; the strided
69            // v[j]/v[j+h] form did not (2k measured 49.6us/encode before).
70            for j in 0..h {
71                let (x, y) = (a[j], b[j]);
72                a[j] = x + y;
73                b[j] = x - y;
74            }
75        }
76        h *= 2;
77    }
78}
79
80/// The prepared recipe: sign vectors and scales are FIXED per (seed, width),
81/// so they are computed once here, never per vector. Before this hoist the
82/// encode path allocated fresh sign vectors and re-ran the PRNG for EVERY
83/// vector -- measured 3.3x on the 250K ingest ladder; the hoist plus
84/// ROUNDS 3 -> 1 (recall unchanged at 1.000 on both the gaussian and the
85/// hostile sparse dataset; the identity-rotation mutation fails the sparse
86/// case at 0.555) brought bulk 15.9s -> 8.6s.
87pub struct Encoder {
88    padded: usize,
89    bits: usize,
90    levels: usize,
91    signs: Vec<Vec<f32>>,
92    /// (1/sqrt(n))^ROUNDS folded into one multiply at the end.
93    scale: f32,
94    table: [f32; MAX_LEVELS],
95}
96
97impl Encoder {
98    pub fn new(dim: usize, seed: u64, bits: usize) -> Encoder {
99        let padded = pad_dim(dim);
100        Encoder {
101            padded,
102            bits,
103            levels: 1 << bits,
104            signs: (0..ROUNDS).map(|r| signs(seed, r, padded)).collect(),
105            scale: (1.0 / (padded as f32).sqrt()).powi(ROUNDS as i32),
106            table: level_table(padded, bits),
107        }
108    }
109
110    pub fn bits(&self) -> usize { self.bits }
111
112    pub fn padded(&self) -> usize { self.padded }
113
114    /// Rotate `x` (already padded) in place.
115    pub fn rotate(&self, x: &mut [f32]) {
116        for sg in &self.signs {
117            for (xi, s) in x.iter_mut().zip(sg) { *xi *= s; }
118            fwht(x);
119        }
120        for xi in x.iter_mut() { *xi *= self.scale; }
121    }
122
123    /// Encode: returns (norm, packed code) for a raw vector (dim <= padded).
124    pub fn encode(&self, v: &[f32]) -> (f32, Vec<u8>) {
125        let norm = v.iter().map(|a| a * a).sum::<f32>().sqrt();
126        let mut x = vec![0f32; self.padded];
127        if norm > 0.0 {
128            for (xi, vi) in x.iter_mut().zip(v) { *xi = vi / norm; }
129        }
130        self.rotate(&mut x);
131        let sd = (self.padded as f32).sqrt();
132        let per = 8 / self.bits;
133        let mut code = vec![0u8; self.padded * self.bits / 8];
134        for (j, &xi) in x.iter().enumerate() {
135            let z = (xi * sd).clamp(-SPAN, SPAN);
136            let q = (((z + SPAN) / (2.0 * SPAN)) * (self.levels as f32 - 1.0)).round() as usize;
137            let q = q.min(self.levels - 1) as u8;
138            code[j / per] |= q << ((j % per) * self.bits);
139        }
140        (norm, code)
141    }
142
143    /// Pad + rotate a query once per search.
144    pub fn rotate_query(&self, q: &[f32]) -> Vec<f32> {
145        let mut x = vec![0f32; self.padded];
146        x[..q.len()].copy_from_slice(q);
147        self.rotate(&mut x);
148        x
149    }
150
151}
152
153/// The affine estimate (2g.2): because the quantizer ladder is UNIFORM,
154/// level value v(l) = l * step - SPAN (in coordinate units), so
155///   dot(x, q) ~ norm * ( step * SUM(l_j * q_j)  -  SPAN * SUM(q_j) ) / sd
156/// The second sum is one constant per query; the first needs no table at
157/// all -- unpack each nibble to an integer, convert, multiply-add. Every
158/// step of that loop is straight-line arithmetic the compiler can
159/// vectorize, unlike a table lookup which never vectorizes. Four
160/// accumulators break the dependency chain as before.
161pub struct AffineQuery {
162    pub qrot: Vec<f32>,
163    /// step / sqrt(padded), premultiplied.
164    pub a: f32,
165    /// -SPAN/sqrt(padded) * SUM(qrot), premultiplied.
166    pub b: f32,
167}
168
169impl Encoder {
170    pub fn affine_query(&self, q: &[f32]) -> AffineQuery {
171        let qrot = self.rotate_query(q);
172        let sd = (self.padded as f32).sqrt();
173        let step = 2.0 * SPAN / (self.levels as f32 - 1.0);
174        let sum_q: f32 = qrot.iter().sum();
175        AffineQuery { qrot, a: step / sd, b: -SPAN / sd * sum_q }
176    }
177}
178
179/// dot estimate via the affine form, 2-bit codes (4 coords per byte).
180/// 256 x 4 lane-value table: row b = the four 2-bit lane values of byte b
181/// as f32. 4KB, L1-resident; turns per-lane shift+convert (which never
182/// vectorized -- 1592ns/code) into one contiguous 4-float load per byte.
183static LANE4: [[f32; 4]; 256] = {
184    let mut t = [[0f32; 4]; 256];
185    let mut b = 0usize;
186    while b < 256 {
187        t[b] = [(b & 3) as f32, ((b >> 2) & 3) as f32,
188                ((b >> 4) & 3) as f32, ((b >> 6) & 3) as f32];
189        b += 1;
190    }
191    t
192};
193
194pub fn dot_est_affine2(norm: f32, code: &[u8], aq: &AffineQuery) -> f32 {
195    let q = &aq.qrot;
196    let (mut a0, mut a1, mut a2, mut a3) = (0f32, 0f32, 0f32, 0f32);
197    // 16 code bytes = 64 lanes per iteration; both sides chunked exact so
198    // the compiler sees fixed trip counts and no bounds checks (2k: the
199    // per-byte form with open indexing measured 1993ns/code).
200    let cb = code.chunks_exact(16);
201    let rest = cb.remainder();
202    let qb = q.chunks_exact(64);
203    for (ch, qk) in cb.zip(qb) {
204        for i in 0..16 {
205            let l = &LANE4[ch[i] as usize];
206            let base = i * 4;
207            a0 += l[0] * qk[base];
208            a1 += l[1] * qk[base + 1];
209            a2 += l[2] * qk[base + 2];
210            a3 += l[3] * qk[base + 3];
211        }
212    }
213    let mut j = (code.len() - rest.len()) * 4;
214    for &b in rest {
215        a0 += (b & 3) as f32 * q[j];
216        a1 += ((b >> 2) & 3) as f32 * q[j + 1];
217        a2 += ((b >> 4) & 3) as f32 * q[j + 2];
218        a3 += ((b >> 6) & 3) as f32 * q[j + 3];
219        j += 4;
220    }
221    let s = (a0 + a1) + (a2 + a3);
222    norm * (aq.a * s + aq.b)
223}
224
225/// dot estimate via the affine form, 4-bit codes (nibble-packed).
226pub fn dot_est_affine(norm: f32, code: &[u8], aq: &AffineQuery) -> f32 {
227    let q = &aq.qrot;
228    let (mut a0, mut a1, mut a2, mut a3) = (0f32, 0f32, 0f32, 0f32);
229    let chunks = code.chunks_exact(4);
230    let rest = chunks.remainder();
231    let mut j = 0usize;
232    for ch in chunks {
233        a0 += (ch[0] & 0x0F) as f32 * q[j]     + (ch[0] >> 4) as f32 * q[j + 1];
234        a1 += (ch[1] & 0x0F) as f32 * q[j + 2] + (ch[1] >> 4) as f32 * q[j + 3];
235        a2 += (ch[2] & 0x0F) as f32 * q[j + 4] + (ch[2] >> 4) as f32 * q[j + 5];
236        a3 += (ch[3] & 0x0F) as f32 * q[j + 6] + (ch[3] >> 4) as f32 * q[j + 7];
237        j += 8;
238    }
239    for &b in rest {
240        a0 += (b & 0x0F) as f32 * q[j] + (b >> 4) as f32 * q[j + 1];
241        j += 2;
242    }
243    let s = (a0 + a1) + (a2 + a3);
244    norm * (aq.a * s + aq.b)
245}
246
247
248/// The 16 reconstruction values, in coordinate units (already divided by
249/// sqrt(padded)): dequant(level) * sd = the ladder midpoint.
250pub fn level_table(padded: usize, bits: usize) -> [f32; MAX_LEVELS] {
251    let sd = (padded as f32).sqrt();
252    let levels = 1 << bits;
253    let mut t = [0f32; MAX_LEVELS];
254    for l in 0..levels {
255        let z = (l as f32) / (levels as f32 - 1.0) * (2.0 * SPAN) - SPAN;
256        t[l] = z / sd;
257    }
258    t
259}