Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

5 changes: 5 additions & 0 deletions nodedb-vector/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -44,3 +44,8 @@ rand = { workspace = true }
tempfile = { workspace = true }
libc = { workspace = true }
nodedb-wal = { workspace = true }
fluxbench = { workspace = true }

[[bench]]
name = "bbq_kernel"
harness = false
130 changes: 130 additions & 0 deletions nodedb-vector/benches/bbq_kernel.rs
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);
}
}
59 changes: 59 additions & 0 deletions nodedb-vector/src/distance/simd/avx2.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Move this use to 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-file use in avx512.rs and neon.rs.)

/// 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 {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The 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 loadu/get_unchecked on the assumption that centered.len() >= dim * 4 and packed.len() >= dim.div_ceil(8). Nothing checks either length. nodedb_vector::distance::simd is public and so is SimdRuntime::l2_bbq, so safe code such as l2_bbq(&[], &[], 1.0, 16) reads out of bounds, which is undefined behaviour. The existing kernels in this file guard their entry point (assert_eq!(a.len(), b.len(), "avx2 l2: length mismatch")). Check both slice lengths against dim here the same way, before the unsafe block, and in every tier.

// 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
}
36 changes: 36 additions & 0 deletions nodedb-vector/src/distance/simd/avx512.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The 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 loadu/get_unchecked on the assumption that centered.len() >= dim * 4 and packed.len() >= dim.div_ceil(8). Nothing checks either length. nodedb_vector::distance::simd is public and so is SimdRuntime::l2_bbq, so safe code such as l2_bbq(&[], &[], 1.0, 16) reads out of bounds, which is undefined behaviour. The existing kernels in avx2.rs guard their entry point (assert_eq!(a.len(), b.len(), "avx2 l2: length mismatch")). Check both slice lengths against dim here the same way, before the unsafe block, and in every tier.

// 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()
}
Loading
Loading