1use std::time::Instant;
20
21use maxsim_lut::{supported_kernels, Codes, ColbertPacking, DocView, Lut, Packing, PreparedQuery, Scorer};
22
23struct Rng(u64);
24impl Rng {
25 fn next(&mut self) -> u64 {
26 let mut x = self.0;
27 x ^= x << 13;
28 x ^= x >> 7;
29 x ^= x << 17;
30 self.0 = x;
31 x
32 }
33 fn f32(&mut self, lo: f32, hi: f32) -> f32 {
34 lo + (hi - lo) * ((self.next() >> 40) as f32 / (1u64 << 24) as f32)
35 }
36}
37
38struct Arm {
40 label: String,
41 lut: Lut,
42 ns: Vec<f64>,
43 checksum: f64,
44}
45
46fn median(sorted: &[f64]) -> f64 {
47 let n = sorted.len();
48 if n % 2 == 1 {
49 sorted[n / 2]
50 } else {
51 0.5 * (sorted[n / 2 - 1] + sorted[n / 2])
52 }
53}
54
55fn main() {
56 let args: Vec<usize> = std::env::args()
57 .skip(1)
58 .map(|a| a.parse().expect("integer arg"))
59 .collect();
60 let dim = args.first().copied().unwrap_or(128);
61 let nbits = args.get(1).copied().unwrap_or(4);
62 let nq = args.get(2).copied().unwrap_or(32);
63 let ntok = args.get(3).copied().unwrap_or(240);
64 let ndocs = args.get(4).copied().unwrap_or(1024);
65 let ncent = 16_384;
66 let reps = 9;
67 let mut rng = Rng(0x9E3779B97F4A7C15);
68
69 let p = ColbertPacking::new(nbits).unwrap();
70 let nb = 1usize << nbits;
71 let mut w: Vec<f32> = (0..nb).map(|_| rng.f32(-0.4, 0.4)).collect();
72 w.sort_by(|a, b| a.total_cmp(b));
73 let lut = Lut::new(&p, &w).unwrap();
74
75 let query: Vec<f32> = (0..nq * dim).map(|_| rng.f32(-1.0, 1.0)).collect();
76 let q = PreparedQuery::new(&lut, &query, nq, dim).unwrap();
77 let cdot: Vec<f32> = (0..ncent * nq).map(|_| rng.f32(-1.0, 1.0)).collect();
78
79 let pdim = dim / p.keys_per_byte();
80 let mut packed = vec![0u8; ndocs * ntok * pdim];
81 for b in packed.iter_mut() {
82 *b = (rng.next() >> 56) as u8;
83 }
84 let codes: Vec<u32> = (0..ndocs * ntok)
85 .map(|_| (rng.next() % ncent as u64) as u32)
86 .collect();
87 let inv: Vec<f32> = (0..ndocs * ntok).map(|_| rng.f32(0.8, 1.2)).collect();
88 let docs: Vec<DocView> = (0..ndocs)
89 .map(|d| {
90 DocView::new(&packed[d * ntok * pdim..(d + 1) * ntok * pdim], ntok, pdim)
91 .codes(Codes::U32(&codes[d * ntok..(d + 1) * ntok]))
92 .inv_norms(&inv[d * ntok..(d + 1) * ntok])
93 })
94 .collect();
95
96 let dispatched = lut.kernel(dim);
100 let mut arms: Vec<Arm> = Vec::new();
101 for &k in supported_kernels() {
102 arms.push(Arm {
103 label: if k == dispatched {
104 format!("{k} (dispatched)")
105 } else {
106 format!("{k}")
107 },
108 lut: lut.clone().pin_kernel(Some(k)),
109 ns: Vec::new(),
110 checksum: 0.0,
111 });
112 }
113 if !dispatched.is_simd() {
114 arms.push(Arm {
115 label: format!("{dispatched} (dispatched)"),
116 lut: lut.clone(),
117 ns: Vec::new(),
118 checksum: 0.0,
119 });
120 }
121 arms.push(Arm {
122 label: "scalar reference".to_string(),
123 lut: lut.clone().force_scalar(true),
124 ns: Vec::new(),
125 checksum: 0.0,
126 });
127
128 println!(
129 "dim {dim}, nbits {nbits}, {nq} query tokens, {ndocs} docs × {ntok} tokens, {ncent} centroids\narch {}, dispatched kernel: {dispatched}, {reps} interleaved rounds",
130 std::env::consts::ARCH,
131 );
132
133 let mut out = vec![0.0f32; ndocs];
134 for arm in arms.iter_mut() {
135 let s = Scorer::new(&arm.lut, &q)
136 .with_centroid_term(&cdot, ncent)
137 .unwrap();
138 s.score_many(docs.iter().copied(), &mut out); }
140 for _ in 0..reps {
141 for arm in arms.iter_mut() {
142 let s = Scorer::new(&arm.lut, &q)
143 .with_centroid_term(&cdot, ncent)
144 .unwrap();
145 let t = Instant::now();
146 s.score_many(docs.iter().copied(), &mut out);
147 arm.ns.push(t.elapsed().as_nanos() as f64 / (ndocs * ntok) as f64);
148 arm.checksum = out.iter().map(|&v| v as f64).sum();
149 }
150 }
151
152 let reference = arms.last().expect("at least the scalar arm");
153 let (slow, want) = {
154 let mut v = reference.ns.clone();
155 v.sort_by(f64::total_cmp);
156 (v[0], reference.checksum)
157 };
158 let mut worst_spread = 0.0f64;
159 for arm in &arms {
160 let mut v = arm.ns.clone();
161 v.sort_by(f64::total_cmp);
162 let (best, med) = (v[0], median(&v));
163 let spread = (med - best) / best;
164 worst_spread = worst_spread.max(spread);
165 assert_eq!(
166 arm.checksum.to_bits(),
167 want.to_bits(),
168 "{}: checksum {} differs from the scalar reference {want}",
169 arm.label,
170 arm.checksum
171 );
172 println!(
173 "{:>26}: {best:7.2} ns/token (median {med:7.2}, {:5.1} µs/doc, {:5.2}x scalar)",
174 arm.label,
175 best * ntok as f64 / 1e3,
176 slow / best,
177 );
178 }
179 println!(
180 "{:>26}: {:.1}% median-vs-best spread — {}",
181 "noise",
182 worst_spread * 100.0,
183 if worst_spread < 0.10 {
184 "quiet enough to compare kernels"
185 } else {
186 "TOO NOISY, differences under ~2x are not real; free the machine and rerun"
187 }
188 );
189}