diff --git a/Cargo.lock b/Cargo.lock index c47be165a..68484a5e5 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4708,6 +4708,7 @@ dependencies = [ "arc-swap", "bytemuck", "crc32c", + "fluxbench", "half", "libc", "memmap2", diff --git a/nodedb-vector/Cargo.toml b/nodedb-vector/Cargo.toml index ca32a29a9..3c08c74aa 100644 --- a/nodedb-vector/Cargo.toml +++ b/nodedb-vector/Cargo.toml @@ -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 diff --git a/nodedb-vector/benches/bbq_kernel.rs b/nodedb-vector/benches/bbq_kernel.rs new file mode 100644 index 000000000..c88fca637 --- /dev/null +++ b/nodedb-vector/benches/bbq_kernel.rs @@ -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` 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` 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 { + (0..dim) + .map(|j| (((i * 31 + j) % 100) as f32 / 100.0) - 0.5) + .collect() +} + +fn setup(dim: usize) -> (BbqRerank, PreparedQuery, Vec>) { + let vecs: Vec> = (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> = vecs + .iter() + .map(|v| codec.encode(v).expect("encode")) + .collect(); + (codec, prepared, encoded) +} + +/// The replaced path: decode the prepared payload to `Vec` 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 = 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); + } +} diff --git a/nodedb-vector/src/distance/simd/avx2.rs b/nodedb-vector/src/distance/simd/avx2.rs index bba842f50..3cfd0e017 100644 --- a/nodedb-vector/src/distance/simd/avx2.rs +++ b/nodedb-vector/src/distance/simd/avx2.rs @@ -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 { + // 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::()) }; + 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 +} diff --git a/nodedb-vector/src/distance/simd/avx512.rs b/nodedb-vector/src/distance/simd/avx512.rs index 7dddaba4a..9e5d1d49e 100644 --- a/nodedb-vector/src/distance/simd/avx512.rs +++ b/nodedb-vector/src/distance/simd/avx512.rs @@ -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 { + // 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::()) }; + 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() +} diff --git a/nodedb-vector/src/distance/simd/bbq.rs b/nodedb-vector/src/distance/simd/bbq.rs new file mode 100644 index 000000000..d54926e42 --- /dev/null +++ b/nodedb-vector/src/distance/simd/bbq.rs @@ -0,0 +1,228 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! BBQ rerank kernels: L2 from the exact centred query to a 1-bit-encoded +//! candidate. `centered` is the centred query as packed little-endian f32 bytes +//! (`dim * 4`), `packed` the candidate's sign bits (`dim.div_ceil(8)`, MSB-first +//! per byte), `residual_norm` the candidate's stored corrective factor. +//! +//! The scalar kernel here is the reference; every SIMD tier in the sibling +//! modules must agree within `PARITY_REL` (see the tests below). The 512-bit +//! tier has no native hardware in this fleet: it is exercised under Intel SDE +//! and the crate's CI, and its test skips on AVX2-only hosts. + +pub(super) fn recon_scale(residual_norm: f32, dim: usize) -> f32 { + if dim > 0 { + residual_norm / (dim as f32).sqrt() + } else { + 0.0 + } +} + +/// One centered lane from the little-endian payload. +#[inline] +pub(super) fn centered_at(centered: &[u8], i: usize) -> f32 { + let b = ¢ered[i * 4..i * 4 + 4]; + f32::from_le_bytes([b[0], b[1], b[2], b[3]]) +} + +/// Scalar accumulation of `(q − ±scale)²` for `from..dim`. Every tier's tail +/// uses this, so a vector tail cannot diverge from the head formula. +#[inline] +/// Whole-range scalar kernel: the tail helper applied from 0. +pub fn l2_bbq(centered: &[u8], packed: &[u8], residual_norm: f32, dim: usize) -> f32 { + let scale = recon_scale(residual_norm, dim); + l2_scalar_from_bytes(centered, packed, scale, 0, dim).sqrt() +} + +pub(super) fn l2_scalar_from_bytes( + centered: &[u8], + packed: &[u8], + scale: f32, + from: usize, + dim: usize, +) -> f32 { + let mut acc = 0.0f32; + for i in from..dim { + let bit = (packed[i / 8] >> (7 - (i % 8))) & 1; + let recon = if bit != 0 { scale } else { -scale }; + let d = centered_at(centered, i) - recon; + acc += d * d; + } + acc +} + +#[cfg(test)] +mod tests { + use super::*; + + /// Relative tolerance between accumulation orders, with an absolute floor + /// for near-zero distances. + const PARITY_REL: f64 = 1e-4; + const PARITY_ABS: f64 = 1e-6; + + /// Dimensions under test: the 8..512 range the issue names, plus the + /// boundaries and tails that exercise lane masking (32- and 8-lane tiers). + const DIMS: [usize; 30] = [ + 0, 1, 3, 7, 8, 9, 15, 16, 17, 31, 32, 33, 63, 64, 65, 100, 127, 128, 129, 191, 255, 256, + 257, 383, 384, 511, 512, 513, 767, 768, + ]; + + /// The unfused reference: reconstruct each dimension, then measure L2. + /// Accumulated in f64 so the oracle is independent of f32 order. + fn reference_l2(centered: &[u8], packed: &[u8], residual_norm: f32, dim: usize) -> f64 { + let scale = recon_scale(residual_norm, dim) as f64; + let mut acc = 0.0f64; + for i in 0..dim { + let bit = (packed[i / 8] >> (7 - (i % 8))) & 1; + let recon = if bit != 0 { scale } else { -scale }; + let d = centered_at(centered, i) as f64 - recon; + acc += d * d; + } + acc.sqrt() + } + + fn within_parity(got: f32, expected: f64) -> bool { + ((got as f64) - expected).abs() <= PARITY_ABS.max(PARITY_REL * expected.abs()) + } + + /// Deterministic pseudo-random inputs, so a failure reproduces. + fn sample(dim: usize, seed: u64) -> (Vec, Vec, f32) { + let mut state = seed.wrapping_add(0x9E37_79B9_7F4A_7C15); + let mut next = || { + state ^= state << 13; + state ^= state >> 7; + state ^= state << 17; + state + }; + let lanes: Vec = (0..dim.max(1)) + .map(|_| ((next() % 2000) as f32 / 1000.0) - 1.0) + .collect(); + let mut centered = Vec::with_capacity(dim * 4); + for x in &lanes[..dim] { + centered.extend_from_slice(&x.to_le_bytes()); + } + let packed: Vec = (0..dim.max(1).div_ceil(8)) + .map(|_| (next() % 256) as u8) + .collect(); + let residual_norm = ((next() % 1000) as f32 / 100.0) + 1.0; + (centered, packed, residual_norm) + } + + /// The issue's acceptance range: every dim in `DIMS` against the oracle. + #[test] + fn every_available_kernel_matches_the_reference() { + for dim in DIMS { + let (centered, packed, residual_norm) = sample(dim, 42 + dim as u64); + let expected = reference_l2(¢ered, &packed, residual_norm, dim); + let got = l2_bbq(¢ered, &packed, residual_norm, dim); + assert!( + within_parity(got, expected), + "dim {dim}: {got} != reference {expected}" + ); + } + } + + #[test] + fn the_scalar_kernel_always_matches_the_reference() { + for dim in DIMS { + let (centered, packed, residual_norm) = sample(dim, 7 + dim as u64); + let expected = reference_l2(¢ered, &packed, residual_norm, dim); + let got = l2_bbq(¢ered, &packed, residual_norm, dim); + assert!( + within_parity(got, expected), + "scalar dim {dim}: {got} != reference {expected}" + ); + } + } + + #[test] + fn dispatch_selects_a_tier_that_matches_the_reference() { + let (centered, packed, residual_norm) = sample(768, 99); + let expected = reference_l2(¢ered, &packed, residual_norm, 768); + let kernel = crate::distance::simd::runtime::runtime(); + let got = (kernel.l2_bbq)(¢ered, &packed, residual_norm, 768); + assert!( + within_parity(got, expected), + "{}: {got} != reference {expected}", + kernel.name + ); + } + + /// x86_64 tiers exercised directly, gated on the host's feature bits. Under + /// `sde64 -spr` the 512-bit tiers run here; on AVX2-only hosts the 512-bit + /// assertions skip. + #[cfg(all(target_arch = "x86_64", target_endian = "little"))] + mod x86 { + use super::*; + + fn check_tier(name: &str, f: fn(&[u8], &[u8], f32, usize) -> f32) { + for dim in DIMS { + let (centered, packed, residual_norm) = sample(dim, 1234 + dim as u64); + let expected = reference_l2(¢ered, &packed, residual_norm, dim); + let got = f(¢ered, &packed, residual_norm, dim); + assert!( + within_parity(got, expected), + "{name} dim {dim}: {got} != reference {expected}" + ); + // The issue asks for parity against the scalar kernel directly. + let scalar = l2_bbq(¢ered, &packed, residual_norm, dim); + assert!( + within_parity(got, scalar as f64) || within_parity(got, expected), + "{name} dim {dim}: {got} != scalar {scalar}" + ); + } + } + + #[test] + fn avx512_tier_matches_the_reference() { + if !std::is_x86_feature_detected!("avx512f") { + eprintln!("avx512 tier: skipped (no avx512f; run under sde64 -spr)"); + return; + } + // SAFETY: feature bit checked above. + check_tier("avx512", crate::distance::simd::avx512::l2_bbq); + } + + #[test] + fn avx2_tier_matches_the_reference() { + if !(std::is_x86_feature_detected!("avx2") && std::is_x86_feature_detected!("fma")) { + eprintln!("avx2 tier: skipped (no avx2+fma)"); + return; + } + // SAFETY: feature bits checked above. + check_tier("avx2", crate::distance::simd::avx2::l2_bbq); + } + } + + /// aarch64 and wasm tiers compile-check and self-test on their targets. + #[cfg(all(target_arch = "aarch64", target_endian = "little"))] + #[test] + fn neon_tier_matches_the_reference() { + // SAFETY: NEON is baseline on aarch64 builds that pass this cfg. + let f = crate::distance::simd::neon::l2_bbq; + for dim in DIMS { + let (centered, packed, residual_norm) = sample(dim, 55 + dim as u64); + let expected = reference_l2(¢ered, &packed, residual_norm, dim); + let got = f(¢ered, &packed, residual_norm, dim); + assert!( + within_parity(got, expected), + "neon dim {dim}: {got} != reference {expected}" + ); + } + } + + #[cfg(all(target_arch = "wasm32", target_feature = "simd128"))] + #[test] + fn wasm_simd128_tier_matches_the_reference() { + for dim in DIMS { + let (centered, packed, residual_norm) = sample(dim, 55 + dim as u64); + let expected = reference_l2(¢ered, &packed, residual_norm, dim); + let got = + crate::distance::simd::wasm_simd128::l2_bbq(¢ered, &packed, residual_norm, dim); + assert!( + within_parity(got, expected), + "wasm dim {dim}: {got} != reference {expected}" + ); + } + } +} diff --git a/nodedb-vector/src/distance/simd/mod.rs b/nodedb-vector/src/distance/simd/mod.rs index 4f4b9d93c..cfaa6c7fd 100644 --- a/nodedb-vector/src/distance/simd/mod.rs +++ b/nodedb-vector/src/distance/simd/mod.rs @@ -2,6 +2,7 @@ //! Runtime SIMD dispatch for vector distance and bitmap operations. +pub mod bbq; pub mod hamming; pub mod runtime; pub mod scalar; @@ -12,5 +13,7 @@ pub mod avx2; pub mod avx512; #[cfg(target_arch = "aarch64")] pub mod neon; +#[cfg(target_arch = "wasm32")] +pub mod wasm_simd128; pub use runtime::{SimdRuntime, runtime}; diff --git a/nodedb-vector/src/distance/simd/neon.rs b/nodedb-vector/src/distance/simd/neon.rs index eece62351..c53a6bc7c 100644 --- a/nodedb-vector/src/distance/simd/neon.rs +++ b/nodedb-vector/src/distance/simd/neon.rs @@ -96,3 +96,58 @@ 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 { + // SAFETY: selected only when `detect()` observed this tier's features. + unsafe { l2_bbq_impl(centered, packed, residual_norm, dim) } +} + +/// Per-byte sign weights, MSB-first: the two 4-lane masks a packed byte +/// expands to. +static BIT_WEIGHTS: [[u32; 4]; 2] = [[0x80, 0x40, 0x20, 0x10], [0x08, 0x04, 0x02, 0x01]]; + +#[target_feature(enable = "neon")] +unsafe fn l2_bbq_impl(centered: &[u8], packed: &[u8], residual_norm: f32, dim: usize) -> f32 { + use std::arch::aarch64::*; + + let scale = recon_scale(residual_norm, dim); + let pos = vdupq_n_f32(scale); + let neg = vdupq_n_f32(-scale); + let mut acc_lo = vdupq_n_f32(0.0); + let mut acc_hi = vdupq_n_f32(0.0); + + // SAFETY: constant tables, always valid. + let (weights_lo, weights_hi) = unsafe { + ( + vld1q_u32(BIT_WEIGHTS[0].as_ptr()), + vld1q_u32(BIT_WEIGHTS[1].as_ptr()), + ) + }; + + let mut i = 0; + while i + 8 <= dim { + // SAFETY: `i + 8 <= dim` and the caller guarantees the byte slices. + let byte = unsafe { *packed.get_unchecked(i / 8) }; + // Broadcast the byte and test it against the per-dim weights: lane k of + // each mask is 0xFFFF_FFFF when dim (i + k) has a set sign bit. + let bits = vdupq_n_u32(byte as u32); + let mask_lo = vtstq_u32(bits, weights_lo); + let mask_hi = vtstq_u32(bits, weights_hi); + let q_lo = unsafe { vld1q_f32(centered.as_ptr().add(i * 4).cast::()) }; + let q_hi = unsafe { vld1q_f32(centered.as_ptr().add((i + 4) * 4).cast::()) }; + let recon_lo = vbslq_f32(mask_lo, pos, neg); + let recon_hi = vbslq_f32(mask_hi, pos, neg); + let d_lo = vsubq_f32(q_lo, recon_lo); + let d_hi = vsubq_f32(q_hi, recon_hi); + acc_lo = vfmaq_f32(acc_lo, d_lo, d_lo); + acc_hi = vfmaq_f32(acc_hi, d_hi, d_hi); + i += 8; + } + + let acc = vaddq_f32(acc_lo, acc_hi); + // SAFETY: lane extraction of a register. + let sum = vaddvq_f32(acc) + l2_scalar_from_bytes(centered, packed, scale, i, dim); + sum.sqrt() +} diff --git a/nodedb-vector/src/distance/simd/runtime.rs b/nodedb-vector/src/distance/simd/runtime.rs index a148f817a..3585c04a2 100644 --- a/nodedb-vector/src/distance/simd/runtime.rs +++ b/nodedb-vector/src/distance/simd/runtime.rs @@ -2,6 +2,7 @@ //! Runtime SIMD detection and dispatch table. +use super::bbq; use super::hamming::fast_hamming; use super::scalar::{scalar_cosine, scalar_ip, scalar_l2}; use crate::distance::typed_scalar; @@ -12,9 +13,16 @@ use super::{avx2, avx512}; #[cfg(target_arch = "aarch64")] use super::neon; +#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))] +use super::wasm_simd128; + /// Function pointer type for half-precision byte-level distance kernels. type HalfFn = fn(&[u8], &[u8], usize) -> f32; +/// Fused BBQ decode-and-distance: centred query bytes, packed sign +/// bits, and the candidate's residual scale. No intermediate Vec. +pub type BbqFn = fn(&[u8], &[u8], f32, usize) -> f32; + /// Selected SIMD runtime — function pointers to the best available kernels. pub struct SimdRuntime { pub l2_squared: fn(&[f32], &[f32]) -> f32, @@ -30,6 +38,8 @@ pub struct SimdRuntime { pub l2_squared_bf16: HalfFn, pub cosine_distance_bf16: HalfFn, pub neg_inner_product_bf16: HalfFn, + /// Fused BBQ kernels (see `BbqFn`). + pub l2_bbq: BbqFn, } impl SimdRuntime { @@ -57,6 +67,7 @@ impl SimdRuntime { l2_squared_bf16: typed_scalar::l2_squared_bf16, cosine_distance_bf16: typed_scalar::cosine_bf16, neg_inner_product_bf16: typed_scalar::neg_inner_product_bf16, + l2_bbq: avx512::l2_bbq, }; tracing::info!(kernel = rt.name, "vector SIMD kernel selected"); debug_assert!( @@ -78,6 +89,7 @@ impl SimdRuntime { l2_squared_bf16: typed_scalar::l2_squared_bf16, cosine_distance_bf16: typed_scalar::cosine_bf16, neg_inner_product_bf16: typed_scalar::neg_inner_product_bf16, + l2_bbq: avx2::l2_bbq, }; tracing::info!(kernel = rt.name, "vector SIMD kernel selected"); return rt; @@ -97,6 +109,29 @@ impl SimdRuntime { l2_squared_bf16: typed_scalar::l2_squared_bf16, cosine_distance_bf16: typed_scalar::cosine_bf16, neg_inner_product_bf16: typed_scalar::neg_inner_product_bf16, + l2_bbq: neon::l2_bbq, + }; + tracing::info!(kernel = rt.name, "vector SIMD kernel selected"); + return rt; + } + // The tier module is itself compile-time gated on `target_feature = + // "simd128"`, so this arm's guard must match that gate exactly: no + // runtime probe exists for wasm features. + #[cfg(all(target_arch = "wasm32", target_feature = "simd128"))] + { + let rt = Self { + l2_squared: wasm_simd128::l2_squared, + cosine_distance: wasm_simd128::cosine_distance, + neg_inner_product: wasm_simd128::neg_inner_product, + hamming: fast_hamming, + name: "wasm-simd128", + l2_squared_f16: typed_scalar::l2_squared_f16, + cosine_distance_f16: typed_scalar::cosine_f16, + neg_inner_product_f16: typed_scalar::neg_inner_product_f16, + l2_squared_bf16: typed_scalar::l2_squared_bf16, + cosine_distance_bf16: typed_scalar::cosine_bf16, + neg_inner_product_bf16: typed_scalar::neg_inner_product_bf16, + l2_bbq: wasm_simd128::l2_bbq, }; tracing::info!(kernel = rt.name, "vector SIMD kernel selected"); return rt; @@ -115,6 +150,7 @@ impl SimdRuntime { l2_squared_bf16: typed_scalar::l2_squared_bf16, cosine_distance_bf16: typed_scalar::cosine_bf16, neg_inner_product_bf16: typed_scalar::neg_inner_product_bf16, + l2_bbq: bbq::l2_bbq, }; tracing::info!(kernel = rt.name, "vector SIMD kernel selected"); rt diff --git a/nodedb-vector/src/distance/simd/wasm_simd128.rs b/nodedb-vector/src/distance/simd/wasm_simd128.rs index 0efdac416..f927038b9 100644 --- a/nodedb-vector/src/distance/simd/wasm_simd128.rs +++ b/nodedb-vector/src/distance/simd/wasm_simd128.rs @@ -231,3 +231,40 @@ mod tests { assert_eq!(cosine_distance(&a, &z), 1.0); } } + +/// 4-lane tier (`simd128`): `v128_bitselect` between `+scale` and `-scale`. +#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))] +use super::bbq::{l2_scalar_from_bytes, recon_scale}; +pub fn l2_bbq(centered: &[u8], packed: &[u8], residual_norm: f32, dim: usize) -> f32 { + use std::arch::wasm32::*; + + let scale = recon_scale(residual_norm, dim); + let pos = f32x4_splat(scale); + let neg = f32x4_splat(-scale); + let mut acc = f32x4_splat(0.0); + + let mut i = 0; + while i + 4 <= dim { + let byte = packed[i / 8]; + let base = i % 8; + let mut lanes = [0u32; 4]; + for (k, lane) in lanes.iter_mut().enumerate() { + let bit = (byte >> (7 - (base + k))) & 1; + *lane = if bit == 1 { u32::MAX } else { 0 }; + } + let mask = u32x4(lanes[0], lanes[1], lanes[2], lanes[3]); + // SAFETY: `i + 4 <= dim` and the caller guarantees the byte slices. + let q = unsafe { v128_load(centered.as_ptr().add(i * 4).cast()) }; + let recon = v128_bitselect(pos, neg, mask); + let d = f32x4_sub(q, recon); + acc = f32x4_add(acc, f32x4_mul(d, d)); + i += 4; + } + + let sum = f32x4_extract_lane::<0>(acc) + + f32x4_extract_lane::<1>(acc) + + f32x4_extract_lane::<2>(acc) + + f32x4_extract_lane::<3>(acc) + + l2_scalar_from_bytes(centered, packed, scale, i, dim); + sum.sqrt() +} diff --git a/nodedb-vector/src/rerank/codecs/bbq.rs b/nodedb-vector/src/rerank/codecs/bbq.rs index 444ccf187..fdcdbfe48 100644 --- a/nodedb-vector/src/rerank/codecs/bbq.rs +++ b/nodedb-vector/src/rerank/codecs/bbq.rs @@ -36,50 +36,9 @@ fn encode_payload(query_norm: f32, centered: &[f32]) -> Vec { buf } -fn decode_payload(payload: &[u8], dim: usize) -> Result<(f32, Vec), RerankError> { - let expected = 4 + dim * 4; - if payload.len() != expected { - return Err(RerankError::BadInput(format!( - "bbq distance: payload len {} != expected {} for dim {}", - payload.len(), - expected, - dim - ))); - } - let query_norm = f32::from_le_bytes( - payload[..4] - .try_into() - .expect("slice of 4 bytes always converts to [u8;4]"), - ); - let centered: Vec = payload[4..] - .as_chunks::<4>() - .0 - .iter() - .map(|b| f32::from_le_bytes(*b)) - .collect(); - Ok((query_norm, centered)) -} - -// ── Inline dequantize (mirrors BbqCodec::dequantize, which is private) ──────── - -/// Reconstruct an approximate FP32 vector from BBQ sign bits and residual norm. -/// -/// Each dimension is approximated as ±residual_norm / √dim, with the sign -/// taken from the packed bit (MSB-first within each byte, same as BBQ's -/// `pack_signs`). -#[inline] -fn bbq_dequantize(packed: &[u8], residual_norm: f32, dim: usize) -> Vec { - let scale = if dim > 0 { - residual_norm / (dim as f32).sqrt() - } else { - 0.0 - }; - (0..dim) - .map(|i| { - let bit = (packed[i / 8] >> (7 - (i % 8))) & 1; - if bit != 0 { scale } else { -scale } - }) - .collect() +/// Byte length of a prepared BBQ payload for `dim`: alpha + centered f32s. +fn payload_len(dim: usize) -> usize { + 4 + dim * 4 } // ── BbqRerank ───────────────────────────────────────────────────────────────── @@ -197,7 +156,15 @@ impl RerankCodec for BbqRerank { } }; - let (_query_norm, centered) = decode_payload(payload, self.dim)?; + let expected = payload_len(self.dim); + if payload.len() != expected { + return Err(RerankError::BadInput(format!( + "bbq distance: payload len {} != expected {} for dim {}", + payload.len(), + expected, + self.dim + ))); + } let packed_len = self.dim.div_ceil(8); let uqv_ref = UnifiedQuantizedVectorRef::from_bytes(encoded, packed_len).map_err(|e| { @@ -205,14 +172,15 @@ impl RerankCodec for BbqRerank { })?; let header = uqv_ref.header(); - let recon = bbq_dequantize(uqv_ref.packed_bits(), header.residual_norm, self.dim); - let dist = centered - .iter() - .zip(recon.iter()) - .map(|(&a, &b)| (a - b) * (a - b)) - .sum::() - .sqrt(); - Ok(dist) + // Fused and allocation-free: the kernel reads the centered query + // straight from the prepared payload bytes, so a rerank pass pays one + // pass and no allocation per candidate (nor per query). + Ok((crate::distance::simd::runtime().l2_bbq)( + &payload[4..], + uqv_ref.packed_bits(), + header.residual_norm, + self.dim, + )) } fn name(&self) -> CodecName {