nodedb_vector/
quantize.rs1use serde::{Deserialize, Serialize};
13
14#[derive(Clone, Serialize, Deserialize)]
16pub struct Sq8Codec {
17 dim: usize,
18 mins: Vec<f32>,
20 maxs: Vec<f32>,
22 scales: Vec<f32>,
24 inv_scales: Vec<f32>,
26}
27
28impl Sq8Codec {
29 pub fn calibrate(vectors: &[&[f32]], dim: usize) -> Self {
34 assert!(!vectors.is_empty(), "cannot calibrate on empty set");
35 assert!(dim > 0);
36
37 let mut mins = vec![f32::MAX; dim];
38 let mut maxs = vec![f32::MIN; dim];
39
40 for v in vectors {
41 debug_assert_eq!(v.len(), dim);
42 for d in 0..dim {
43 if v[d] < mins[d] {
44 mins[d] = v[d];
45 }
46 if v[d] > maxs[d] {
47 maxs[d] = v[d];
48 }
49 }
50 }
51
52 let mut scales = vec![0.0f32; dim];
53 let mut inv_scales = vec![0.0f32; dim];
54 for d in 0..dim {
55 let range = maxs[d] - mins[d];
56 if range > f32::EPSILON {
57 scales[d] = range / 255.0;
58 inv_scales[d] = 255.0 / range;
59 }
60 }
61
62 Self {
63 dim,
64 mins,
65 maxs,
66 scales,
67 inv_scales,
68 }
69 }
70
71 pub fn quantize(&self, vector: &[f32]) -> Vec<u8> {
73 debug_assert_eq!(vector.len(), self.dim);
74 let mut out = Vec::with_capacity(self.dim);
75 for ((&v, &min), (&max, &inv_scale)) in vector
76 .iter()
77 .zip(self.mins.iter())
78 .zip(self.maxs.iter().zip(self.inv_scales.iter()))
79 {
80 let clamped = v.clamp(min, max);
81 let q = ((clamped - min) * inv_scale).round() as u8;
82 out.push(q);
83 }
84 out
85 }
86
87 pub fn quantize_batch(&self, vectors: &[&[f32]]) -> Vec<u8> {
91 let mut out = Vec::with_capacity(self.dim * vectors.len());
92 for v in vectors {
93 out.extend(self.quantize(v));
94 }
95 out
96 }
97
98 pub fn dequantize(&self, quantized: &[u8]) -> Vec<f32> {
100 debug_assert_eq!(quantized.len(), self.dim);
101 let mut out = Vec::with_capacity(self.dim);
102 for ((&q, &min), &scale) in quantized
103 .iter()
104 .zip(self.mins.iter())
105 .zip(self.scales.iter())
106 {
107 out.push(min + q as f32 * scale);
108 }
109 out
110 }
111
112 #[inline]
116 pub fn asymmetric_l2(&self, query: &[f32], candidate: &[u8]) -> f32 {
117 debug_assert_eq!(query.len(), self.dim);
118 debug_assert_eq!(candidate.len(), self.dim);
119 let mut sum = 0.0f32;
120 for d in 0..self.dim {
121 let dequant = self.mins[d] + candidate[d] as f32 * self.scales[d];
122 let diff = query[d] - dequant;
123 sum += diff * diff;
124 }
125 sum
126 }
127
128 #[inline]
130 pub fn asymmetric_cosine(&self, query: &[f32], candidate: &[u8]) -> f32 {
131 debug_assert_eq!(query.len(), self.dim);
132 debug_assert_eq!(candidate.len(), self.dim);
133 let mut dot = 0.0f32;
134 let mut norm_q = 0.0f32;
135 let mut norm_c = 0.0f32;
136 for d in 0..self.dim {
137 let dequant = self.mins[d] + candidate[d] as f32 * self.scales[d];
138 dot += query[d] * dequant;
139 norm_q += query[d] * query[d];
140 norm_c += dequant * dequant;
141 }
142 let denom = (norm_q * norm_c).sqrt();
143 if denom < f32::EPSILON {
144 return 1.0;
145 }
146 (1.0 - dot / denom).max(0.0)
147 }
148
149 #[inline]
151 pub fn asymmetric_ip(&self, query: &[f32], candidate: &[u8]) -> f32 {
152 debug_assert_eq!(query.len(), self.dim);
153 debug_assert_eq!(candidate.len(), self.dim);
154 let mut dot = 0.0f32;
155 for d in 0..self.dim {
156 let dequant = self.mins[d] + candidate[d] as f32 * self.scales[d];
157 dot += query[d] * dequant;
158 }
159 -dot
160 }
161
162 pub fn dimensions(&self) -> usize {
164 self.dim
165 }
166}
167
168#[cfg(test)]
169mod tests {
170 use super::*;
171
172 fn make_vectors() -> Vec<Vec<f32>> {
173 (0..100)
174 .map(|i| vec![i as f32 * 0.1, (i as f32).sin(), (i as f32).cos()])
175 .collect()
176 }
177
178 #[test]
179 fn quantize_dequantize_roundtrip() {
180 let vecs = make_vectors();
181 let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
182 let codec = Sq8Codec::calibrate(&refs, 3);
183
184 for v in &vecs {
185 let q = codec.quantize(v);
186 let dq = codec.dequantize(&q);
187 for d in 0..3 {
188 let error = (v[d] - dq[d]).abs();
189 let range = codec.maxs[d] - codec.mins[d];
190 assert!(
191 error <= range / 255.0 + 1e-6,
192 "d={d}: error={error}, max_step={}",
193 range / 255.0
194 );
195 }
196 }
197 }
198
199 #[test]
200 fn asymmetric_l2_close_to_exact() {
201 let vecs = make_vectors();
202 let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
203 let codec = Sq8Codec::calibrate(&refs, 3);
204
205 let query = &[5.0, 0.5, -0.5];
206 for v in &vecs {
207 let q = codec.quantize(v);
208 let exact = crate::distance::l2_squared(query, v);
209 let approx = codec.asymmetric_l2(query, &q);
210 let rel_error = if exact > 0.01 {
211 (exact - approx).abs() / exact
212 } else {
213 (exact - approx).abs()
214 };
215 assert!(
216 rel_error < 0.05 || (exact - approx).abs() < 0.1,
217 "exact={exact}, approx={approx}, rel_error={rel_error}"
218 );
219 }
220 }
221
222 #[test]
223 fn batch_quantize() {
224 let vecs = make_vectors();
225 let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
226 let codec = Sq8Codec::calibrate(&refs, 3);
227
228 let batch = codec.quantize_batch(&refs);
229 assert_eq!(batch.len(), 3 * 100);
230
231 let single = codec.quantize(&vecs[0]);
232 assert_eq!(&batch[0..3], &single[..]);
233 }
234
235 #[test]
236 fn constant_dimension_handled() {
237 let vecs: Vec<Vec<f32>> = (0..10).map(|i| vec![5.0, i as f32]).collect();
238 let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
239 let codec = Sq8Codec::calibrate(&refs, 2);
240
241 let q = codec.quantize(&[5.0, 3.0]);
242 assert_eq!(q[0], 0); }
244}