Skip to main content

bench/
bench.rs

1//! Nanoseconds per scored document token for **every kernel this CPU can
2//! run**, plus the scalar reference, on synthetic ColBERT-shaped data.
3//!
4//!   cargo run --release --example bench -- [dim] [nbits] [nq] [doc_tokens] [n_docs]
5//!
6//! The arms are interleaved round by round inside one process, and each is
7//! reduced by its minimum, so a scheduling hiccup has to hit every round of
8//! one arm to change the ranking. Comparing separate `cargo run` invocations
9//! instead is how a busy machine invents a regression: on a hybrid CPU one
10//! run can land on an efficiency core and read 2× slow.
11//!
12//! The spread between an arm's minimum and its median is printed as a noise
13//! verdict. Above about 10% the machine is too busy for the small
14//! differences between kernels to mean anything.
15//!
16//! On an Apple Silicon machine whose rustup default is x86_64, pass
17//! `--target aarch64-apple-darwin` or you will benchmark Rosetta.
18
19use 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
38/// One measured arm: a label, the table configured for it, and its timings.
39struct 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    // One arm per executable kernel, plus the scalar reference. `dispatch`
97    // is what an unpinned host would get, named so the calibrated choice is
98    // visible next to the kernels it chose between.
99    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); // warm caches and branch predictors
139    }
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}