1use crate::kernel::{self, Kernel};
4use crate::packing::Packing;
5use crate::Error;
6
7pub const MAX_DIM: usize = 256;
11
12#[derive(Debug, Clone, PartialEq, Eq)]
22pub struct NibbleTables {
23 pub tables: [[i8; 16]; 8],
25 pub from_hi: [bool; 8],
27}
28
29#[derive(Debug, Clone)]
36pub struct Lut {
37 fused: Vec<i8>,
39 keys_per_byte: usize,
40 nbits: usize,
41 scale: f32,
43 nibble: Option<NibbleTables>,
44 force_scalar: bool,
45 pin: Option<Kernel>,
46}
47
48impl Lut {
49 pub fn new<P: Packing>(packing: &P, bucket_weights: &[f32]) -> Result<Self, Error> {
54 let nbits = packing.nbits();
55 if !matches!(nbits, 1 | 2 | 4 | 8) {
56 return Err(Error::NbitsUnsupported(nbits));
57 }
58 let expected = 1usize << nbits;
59 if bucket_weights.len() != expected {
60 return Err(Error::BucketCount {
61 expected,
62 got: bucket_weights.len(),
63 });
64 }
65 let max_abs = bucket_weights.iter().fold(0.0f32, |m, &x| m.max(x.abs()));
66 let scale = (max_abs / 127.0).max(1e-12);
67 let vals: Vec<i8> = bucket_weights
68 .iter()
69 .map(|&w| (w / scale).round().clamp(-127.0, 127.0) as i8)
70 .collect();
71 let keys_per_byte = 8 / nbits;
72 let mut fused = vec![0i8; 256 * keys_per_byte];
73 for byte in 0..256usize {
74 for k in 0..keys_per_byte {
75 let bi = packing.bucket_index(byte as u8, k);
76 assert!(
77 bi < expected,
78 "Packing::bucket_index returned {bi} >= 2^nbits for byte {byte} key {k}"
79 );
80 fused[byte * keys_per_byte + k] = vals[bi];
81 }
82 }
83 let nibble = derive_nibble_tables(&fused, keys_per_byte);
84 Ok(Self {
85 fused,
86 keys_per_byte,
87 nbits,
88 scale,
89 nibble,
90 force_scalar: false,
91 pin: None,
92 })
93 }
94
95 pub fn colbert(nbits: usize, bucket_weights: &[f32]) -> Result<Self, Error> {
97 Self::new(&crate::ColbertPacking::new(nbits)?, bucket_weights)
98 }
99
100 pub fn force_scalar(mut self, yes: bool) -> Self {
105 self.force_scalar = yes;
106 self
107 }
108
109 pub fn pin_kernel(mut self, kernel: Option<Kernel>) -> Self {
120 self.pin = kernel;
121 self
122 }
123
124 pub fn nbits(&self) -> usize {
126 self.nbits
127 }
128
129 pub fn keys_per_byte(&self) -> usize {
131 self.keys_per_byte
132 }
133
134 pub fn scale(&self) -> f32 {
136 self.scale
137 }
138
139 #[inline]
141 pub fn expand(&self, byte: u8) -> &[i8] {
142 let base = byte as usize * self.keys_per_byte;
143 &self.fused[base..base + self.keys_per_byte]
144 }
145
146 pub fn fused_table(&self) -> &[i8] {
148 &self.fused
149 }
150
151 pub fn nibble_tables(&self) -> Option<&NibbleTables> {
154 self.nibble.as_ref()
155 }
156
157 pub fn kernel(&self, dim: usize) -> Kernel {
162 kernel::select(self, dim)
163 }
164
165 pub(crate) fn force_scalar_set(&self) -> bool {
166 self.force_scalar
167 }
168
169 pub(crate) fn pinned_kernel(&self) -> Option<Kernel> {
170 self.pin
171 }
172
173 pub(crate) fn fingerprint(&self) -> (usize, u32) {
177 (self.nbits, self.scale.to_bits())
178 }
179}
180
181fn derive_nibble_tables(fused: &[i8], keys_per_byte: usize) -> Option<NibbleTables> {
184 if keys_per_byte > 8 || keys_per_byte == 1 {
185 return None;
187 }
188 let mut tables = [[0i8; 16]; 8];
189 let mut from_hi = [false; 8];
190 for k in 0..keys_per_byte {
191 let hi: [i8; 16] = std::array::from_fn(|x| fused[(x << 4) * keys_per_byte + k]);
192 if (0..256).all(|b| fused[b * keys_per_byte + k] == hi[b >> 4]) {
193 tables[k] = hi;
194 from_hi[k] = true;
195 continue;
196 }
197 let lo: [i8; 16] = std::array::from_fn(|x| fused[x * keys_per_byte + k]);
198 if (0..256).all(|b| fused[b * keys_per_byte + k] == lo[b & 15]) {
199 tables[k] = lo;
200 from_hi[k] = false;
201 continue;
202 }
203 return None;
204 }
205 Some(NibbleTables { tables, from_hi })
206}
207
208#[cfg(test)]
209mod tests {
210 use super::*;
211 use crate::ColbertPacking;
212
213 fn weights(nbits: usize) -> Vec<f32> {
214 let n = 1usize << nbits;
215 (0..n)
216 .map(|i| -0.35 + 0.7 * (i as f32 + 0.5) / n as f32)
217 .collect()
218 }
219
220 #[test]
221 fn fused_table_expands_to_quantised_bucket_weights() {
222 for nbits in [1usize, 2, 4, 8] {
223 let p = ColbertPacking::new(nbits).unwrap();
224 let w = weights(nbits);
225 let lut = Lut::new(&p, &w).unwrap();
226 for byte in 0..=255u8 {
227 for k in 0..lut.keys_per_byte() {
228 let bi = p.bucket_index(byte, k);
229 let want = (w[bi] / lut.scale()).round().clamp(-127.0, 127.0) as i8;
230 assert_eq!(lut.expand(byte)[k], want, "nbits={nbits} byte={byte} k={k}");
231 }
232 }
233 }
234 }
235
236 #[test]
237 fn nibble_factorisation_holds_for_sub_byte_codes() {
238 for nbits in [1usize, 2, 4] {
239 let lut = Lut::colbert(nbits, &weights(nbits)).unwrap();
240 let nib = lut
241 .nibble_tables()
242 .unwrap_or_else(|| panic!("nbits={nbits}: not nibble-separable"));
243 for b in 0..256usize {
244 for k in 0..lut.keys_per_byte() {
245 let nibble = if nib.from_hi[k] { b >> 4 } else { b & 15 };
246 assert_eq!(
247 lut.fused_table()[b * lut.keys_per_byte() + k],
248 nib.tables[k][nibble]
249 );
250 }
251 }
252 }
253 assert!(Lut::colbert(8, &weights(8)).unwrap().nibble_tables().is_none());
254 }
255
256 #[test]
257 fn rejects_bad_inputs() {
258 assert_eq!(ColbertPacking::new(3).unwrap_err(), Error::NbitsUnsupported(3));
259 assert_eq!(
260 Lut::colbert(4, &[0.0; 15]).unwrap_err(),
261 Error::BucketCount {
262 expected: 16,
263 got: 15
264 }
265 );
266 }
267}