pub trait Packing {
// Required methods
fn nbits(&self) -> usize;
fn bucket_index(&self, byte: u8, key: usize) -> usize;
// Provided methods
fn keys_per_byte(&self) -> usize { ... }
fn pack_row(&self, buckets: &[usize], out: &mut [u8]) { ... }
}Expand description
Describes a bit-packing layout of nbits-wide bucket indices.
Key k of a byte is the k-th embedding dimension that byte carries, in
dimension order: a token’s packed row bytes[0..dim·nbits/8] holds dims
i·keys_per_byte + k at (bytes[i], key k).
Required Methods§
Sourcefn bucket_index(&self, byte: u8, key: usize) -> usize
fn bucket_index(&self, byte: u8, key: usize) -> usize
Bucket index (0 .. 2^nbits) stored at key position key
(0 .. 8/nbits) of byte.
Provided Methods§
Sourcefn keys_per_byte(&self) -> usize
fn keys_per_byte(&self) -> usize
8 / nbits.
Examples found in repository?
examples/bench.rs (line 79)
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}Sourcefn pack_row(&self, buckets: &[usize], out: &mut [u8])
fn pack_row(&self, buckets: &[usize], out: &mut [u8])
Reference packer, the inverse of Packing::bucket_index, for hosts
that want to produce rows the same way the tests do. Not optimised.
Dyn Compatibility§
This trait is dyn compatible.
In older versions of Rust, dyn compatibility was called "object safety".