Skip to content
Merged
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.

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

# `fluxbench` pulls `tokio`, which does not build for wasm32. Keeping it out of
# the wasm dev-dependency set is what lets the crate's own wasm tests build and
# run under a wasm runner, instead of only being compile-checked.
[target.'cfg(not(target_arch = "wasm32"))'.dev-dependencies]
fluxbench = { workspace = true }

[[bench]]
name = "bbq_kernel"
harness = false
147 changes: 147 additions & 0 deletions nodedb-vector/benches/bbq_kernel.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,147 @@
// SPDX-License-Identifier: Apache-2.0

//! BBQ kernel benches: the fused path against a reconstruct-and-measure
//! baseline.
//!
//! The baseline models the shape of the pre-fusion rerank pass — decode the
//! prepared payload into a `Vec<f32>` per candidate, then reconstruct each
//! dimension — so the allocation and pass counts are comparable. It is an
//! approximation of that path, not a copy of it: it uses a fixed `1/√dim` scale
//! in place of the candidate's stored corrective factor, and reads the sign bits
//! at a fixed header offset. It measures pass and allocation shape, not the
//! codec's arithmetic.
//!
//! Fixtures are built once per thread, outside the timed region, so only the
//! kernel loop is measured.
//!
//! 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 shows zero allocations per
/// candidate, the baseline shows one `Vec<f32>`.
#[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()
}

/// Trained codec, prepared query, and 256 encoded candidates for `dim`.
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)
}

thread_local! {
static STATE_128: (BbqRerank, PreparedQuery, Vec<Vec<u8>>) = setup(128);
static STATE_768: (BbqRerank, PreparedQuery, Vec<Vec<u8>>) = setup(768);
}

/// The baseline: decode the prepared payload to `Vec<f32>` per candidate, then
/// reconstruct each dimension. `encoded` carries a 32-byte quant header before
/// the sign bits, and the reconstruction uses a fixed `1/√dim` scale rather than
/// the candidate's stored `residual_norm` — see the module docs.
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) {
STATE_128.with(|(codec, prepared, encoded)| {
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) {
STATE_128.with(|(_codec, prepared, encoded)| {
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) {
STATE_768.with(|(codec, prepared, encoded)| {
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) {
STATE_768.with(|(_codec, prepared, encoded)| {
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 @@ -4,6 +4,8 @@

#![cfg(target_arch = "x86_64")]

use super::bbq::{assert_payload_shapes, l2_scalar_from_bytes, recon_scale};

pub fn l2_squared(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "avx2 l2: length mismatch");
// SAFETY: caller verified avx2+fma via is_x86_feature_detected.
Expand Down Expand Up @@ -114,3 +116,60 @@ unsafe fn hsum256(v: std::arch::x86_64::__m256) -> f32 {
let sums2 = _mm_add_ss(sums, shuf2);
_mm_cvtss_f32(sums2)
}

/// 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 {
assert_payload_shapes(centered, packed, dim);
// SAFETY: selected only when `detect()` observed this tier's features, and
// both slices were just checked against `dim`.
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 `l2_bbq` asserted
// `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`).
#[derive(Clone, Copy)]
#[repr(align(32))]
struct Aligned8([f32; 8]);

static SIGN_LANES: [Aligned8; 256] = build_sign_lanes();

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

#![cfg(target_arch = "x86_64")]

use super::bbq::{assert_payload_shapes, l2_scalar_from_bytes, recon_scale};

pub fn l2_squared(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "avx512 l2: length mismatch");
unsafe { l2_impl(a, b) }
Expand Down Expand Up @@ -99,3 +101,39 @@ unsafe fn ip_impl(a: &[f32], b: &[f32]) -> f32 {
-dot
}
}

/// 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 {
assert_payload_shapes(centered, packed, dim);
// SAFETY: selected only when `detect()` observed this tier's features, and
// both slices were just checked against `dim`.
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 `l2_bbq` asserted
// `centered.len() >= dim * 4`; two packed bytes are in range because
// 16 dims consume exactly two bytes.
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