-
-
Notifications
You must be signed in to change notification settings - Fork 18
feat(vector): fuse the BBQ rerank distance kernels #351
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,130 @@ | ||
| // SPDX-License-Identifier: Apache-2.0 | ||
|
|
||
| //! BBQ kernel benches: the zero-copy fused path against the | ||
| //! reconstruct-and-measure path it replaces. | ||
| //! | ||
| //! The unfused bench mirrors the removed path shape: decode the prepared | ||
| //! payload into a `Vec<f32>` per candidate, then reconstruct each dimension. | ||
| //! (Its residual scale is an approximation: the bench measures the allocation | ||
| //! and pass shape, not the codec's exact corrective factor.) | ||
| //! | ||
| //! Run with: cargo bench -p nodedb-vector --bench bbq_kernel | ||
|
|
||
| use fluxbench::bench; | ||
| use fluxbench::prelude::*; | ||
| use std::hint::black_box; | ||
|
|
||
| use nodedb_vector::rerank::codec::{PreparedQuery, RerankCodec}; | ||
| use nodedb_vector::rerank::codecs::BbqRerank; | ||
|
|
||
| /// Installs fluxbench's tracking allocator so the harness reports heap bytes | ||
| /// and allocation counts per benchmark — the fused path must show zero | ||
| /// allocations, the replaced path must show one `Vec<f32>` per candidate. | ||
| #[global_allocator] | ||
| static ALLOC: fluxbench::TrackingAllocator = fluxbench::TrackingAllocator; | ||
|
|
||
| const OVERSAMPLE: u8 = 4; | ||
| const CANDIDATES: usize = 256; | ||
|
|
||
| fn det_vec(i: usize, dim: usize) -> Vec<f32> { | ||
| (0..dim) | ||
| .map(|j| (((i * 31 + j) % 100) as f32 / 100.0) - 0.5) | ||
| .collect() | ||
| } | ||
|
|
||
| fn setup(dim: usize) -> (BbqRerank, PreparedQuery, Vec<Vec<u8>>) { | ||
| let vecs: Vec<Vec<f32>> = (0..CANDIDATES).map(|i| det_vec(i, dim)).collect(); | ||
| let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect(); | ||
| let mut codec = BbqRerank::new(dim, OVERSAMPLE); | ||
| codec.train(&refs).expect("train"); | ||
| let prepared = codec.prepare_query(&vecs[0]).expect("prepare_query"); | ||
| let encoded: Vec<Vec<u8>> = vecs | ||
| .iter() | ||
| .map(|v| codec.encode(v).expect("encode")) | ||
| .collect(); | ||
| (codec, prepared, encoded) | ||
| } | ||
|
|
||
| /// The replaced path: decode the prepared payload to `Vec<f32>` per candidate, | ||
| /// then reconstruct each dimension. `encoded` carries a 32-byte quant header | ||
| /// before the sign bits. | ||
| fn unfused_l2(payload: &[u8], encoded: &[u8], dim: usize) -> f32 { | ||
| let centered: Vec<f32> = payload[4..] | ||
| .as_chunks::<4>() | ||
| .0 | ||
| .iter() | ||
| .map(|b| f32::from_le_bytes(*b)) | ||
| .collect(); | ||
| let scale = 1.0f32 / (dim as f32).sqrt(); | ||
| let mut acc = 0.0f32; | ||
| for i in 0..dim { | ||
| let bit = (encoded[32 + i / 8] >> (7 - (i % 8))) & 1; | ||
| let recon = if bit != 0 { scale } else { -scale }; | ||
| let d = centered[i] - recon; | ||
| acc += d * d; | ||
| } | ||
| acc.sqrt() | ||
| } | ||
|
|
||
| #[bench(id = "bbq_fused_128", group = "bbq_kernel")] | ||
| fn bbq_fused_128(b: &mut Bencher) { | ||
| let (codec, prepared, encoded) = setup(128); | ||
| b.iter(|| { | ||
| let mut acc = 0.0f32; | ||
| for e in &encoded { | ||
| acc += codec.distance_prepared(&prepared, e).expect("distance"); | ||
| } | ||
| black_box(acc) | ||
| }); | ||
| } | ||
|
|
||
| #[bench(id = "bbq_unfused_128", group = "bbq_kernel")] | ||
| fn bbq_unfused_128(b: &mut Bencher) { | ||
| let (_codec, prepared, encoded) = setup(128); | ||
| let payload = match &prepared { | ||
| PreparedQuery::Bytes(b) => b.as_slice(), | ||
| _ => panic!("bbq prepared form is Bytes"), | ||
| }; | ||
| b.iter(|| { | ||
| let mut acc = 0.0f32; | ||
| for e in &encoded { | ||
| acc += unfused_l2(payload, e, 128); | ||
| } | ||
| black_box(acc) | ||
| }); | ||
| } | ||
|
|
||
| #[bench(id = "bbq_fused_768", group = "bbq_kernel")] | ||
| fn bbq_fused_768(b: &mut Bencher) { | ||
| let (codec, prepared, encoded) = setup(768); | ||
| b.iter(|| { | ||
| let mut acc = 0.0f32; | ||
| for e in &encoded { | ||
| acc += codec.distance_prepared(&prepared, e).expect("distance"); | ||
| } | ||
| black_box(acc) | ||
| }); | ||
| } | ||
|
|
||
| #[bench(id = "bbq_unfused_768", group = "bbq_kernel")] | ||
| fn bbq_unfused_768(b: &mut Bencher) { | ||
| let (_codec, prepared, encoded) = setup(768); | ||
| let payload = match &prepared { | ||
| PreparedQuery::Bytes(b) => b.as_slice(), | ||
| _ => panic!("bbq prepared form is Bytes"), | ||
| }; | ||
| b.iter(|| { | ||
| let mut acc = 0.0f32; | ||
| for e in &encoded { | ||
| acc += unfused_l2(payload, e, 768); | ||
| } | ||
| black_box(acc) | ||
| }); | ||
| } | ||
|
|
||
| fn main() { | ||
| if let Err(e) = fluxbench::run() { | ||
| eprintln!("Error: {e}"); | ||
| std::process::exit(1); | ||
| } | ||
| } |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -114,3 +114,62 @@ unsafe fn hsum256(v: std::arch::x86_64::__m256) -> f32 { | |
| let sums2 = _mm_add_ss(sums, shuf2); | ||
| _mm_cvtss_f32(sums2) | ||
| } | ||
|
|
||
| use super::bbq::{l2_scalar_from_bytes, recon_scale}; | ||
| /// Safe entry for `SimdRuntime`; the feature guard lives in `SimdRuntime::detect`. | ||
| pub fn l2_bbq(centered: &[u8], packed: &[u8], residual_norm: f32, dim: usize) -> f32 { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Blocker: this is a public safe function, but the SIMD loop reads through |
||
| // SAFETY: selected only when `detect()` observed this tier's features. | ||
| unsafe { l2_bbq_impl(centered, packed, residual_norm, dim) } | ||
| } | ||
|
|
||
| #[target_feature(enable = "avx2,fma")] | ||
| unsafe fn l2_bbq_impl(centered: &[u8], packed: &[u8], residual_norm: f32, dim: usize) -> f32 { | ||
| use std::arch::x86_64::*; | ||
|
|
||
| let scale = recon_scale(residual_norm, dim); | ||
| let scale_v = _mm256_set1_ps(scale); | ||
| let mut acc = _mm256_setzero_ps(); | ||
|
|
||
| let mut i = 0; | ||
| while i + 8 <= dim { | ||
| // SAFETY: `i + 8 <= dim` and the caller guarantees | ||
| // `centered.len() >= dim * 4`, so the 32-byte unaligned load stays in | ||
| // bounds; `i / 8` is in range because eight dims consume one byte. | ||
| let q = unsafe { _mm256_loadu_ps(centered.as_ptr().add(i * 4).cast::<f32>()) }; | ||
| let byte = unsafe { *packed.get_unchecked(i / 8) } as usize; | ||
| let signs = unsafe { _mm256_load_ps(SIGN_LANES[byte].0.as_ptr()) }; | ||
| let recon = _mm256_mul_ps(signs, scale_v); | ||
| let d = _mm256_sub_ps(q, recon); | ||
| acc = _mm256_fmadd_ps(d, d, acc); | ||
| i += 8; | ||
| } | ||
|
|
||
| let sum = unsafe { hsum256(acc) } + l2_scalar_from_bytes(centered, packed, scale, i, dim); | ||
| sum.sqrt() | ||
| } | ||
|
|
||
| /// `±1.0` lane patterns for every packed byte, MSB-first, 32-byte aligned for | ||
| /// an aligned load. Precomputed `reverse_bits` mapping (dim `k` → lane `k`). | ||
| #[cfg(all(target_arch = "x86_64", target_endian = "little"))] | ||
| #[derive(Clone, Copy)] | ||
| #[repr(align(32))] | ||
| struct Aligned8([f32; 8]); | ||
|
|
||
| #[cfg(all(target_arch = "x86_64", target_endian = "little"))] | ||
| static SIGN_LANES: [Aligned8; 256] = build_sign_lanes(); | ||
|
|
||
| #[cfg(all(target_arch = "x86_64", target_endian = "little"))] | ||
| const fn build_sign_lanes() -> [Aligned8; 256] { | ||
| let mut table = [Aligned8([0.0; 8]); 256]; | ||
| let mut byte = 0usize; | ||
| while byte < 256 { | ||
| let mut lane = 0usize; | ||
| while lane < 8 { | ||
| let bit = (byte >> (7 - lane)) & 1; | ||
| table[byte].0[lane] = if bit == 1 { 1.0 } else { -1.0 }; | ||
| lane += 1; | ||
| } | ||
| byte += 1; | ||
| } | ||
| table | ||
| } | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -99,3 +99,39 @@ unsafe fn ip_impl(a: &[f32], b: &[f32]) -> f32 { | |
| -dot | ||
| } | ||
| } | ||
|
|
||
| use super::bbq::{l2_scalar_from_bytes, recon_scale}; | ||
| /// Safe entry for `SimdRuntime`; the feature guard lives in `SimdRuntime::detect`. | ||
| pub fn l2_bbq(centered: &[u8], packed: &[u8], residual_norm: f32, dim: usize) -> f32 { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Blocker: this is a public safe function, but the SIMD loop reads through |
||
| // SAFETY: selected only when `detect()` observed this tier's features. | ||
| unsafe { l2_bbq_impl(centered, packed, residual_norm, dim) } | ||
| } | ||
|
|
||
| #[target_feature(enable = "avx512f")] | ||
| unsafe fn l2_bbq_impl(centered: &[u8], packed: &[u8], residual_norm: f32, dim: usize) -> f32 { | ||
| use std::arch::x86_64::*; | ||
|
|
||
| let scale = recon_scale(residual_norm, dim); | ||
| let pos = _mm512_set1_ps(scale); | ||
| let neg = _mm512_set1_ps(-scale); | ||
| let mut acc = _mm512_setzero_ps(); | ||
|
|
||
| let mut i = 0; | ||
| while i + 16 <= dim { | ||
| // SAFETY: `i + 16 <= dim` and the caller guarantees | ||
| // `centered.len() >= dim * 4`; two packed bytes are in range because | ||
| // 16 dims consume exactly two bytes. | ||
| // SAFETY: in-bounds per the comment above. | ||
| let q = unsafe { _mm512_loadu_ps(centered.as_ptr().add(i * 4).cast::<f32>()) }; | ||
| let b0 = unsafe { *packed.get_unchecked(i / 8) }; | ||
| let b1 = unsafe { *packed.get_unchecked(i / 8 + 1) }; | ||
| let mask: __mmask16 = (b0.reverse_bits() as u16) | ((b1.reverse_bits() as u16) << 8); | ||
| let recon = _mm512_mask_blend_ps(mask, neg, pos); | ||
| let d = _mm512_sub_ps(q, recon); | ||
| acc = _mm512_fmadd_ps(d, d, acc); | ||
| i += 16; | ||
| } | ||
|
|
||
| let sum = _mm512_reduce_add_ps(acc) + l2_scalar_from_bytes(centered, packed, scale, i, dim); | ||
| sum.sqrt() | ||
| } | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Move this
useto the top of the file with the other imports. The#[cfg(all(target_arch = "x86_64", ...))]on lines 153/158/161 repeats the module's own#![cfg(target_arch = "x86_64")]. Remove it. (The same applies to the mid-fileuseinavx512.rsandneon.rs.)