From fff1264799e4590e7b7de20607b0876e76a75d1b Mon Sep 17 00:00:00 2001 From: EnRaiha <15997552+EnRaiha@users.noreply.github.com> Date: Mon, 14 Sep 2026 12:25:30 +0800 Subject: [PATCH] feat(vector): fuse the BBQ rerank distance into SimdRuntime's kernel table MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Reranking against 1-bit BBQ candidates materialized a `Vec` per candidate and per query; at an oversample x ef candidate count that dominated the pass and defeated SIMD. The kernel now reads the centred query straight from the prepared payload bytes: zero allocation per candidate and per query, one pass. It is a `SimdRuntime` field (`l2_bbq`), selected in `detect()` beside the f16/bf16 fused kernels, with the per-tier kernels beside their siblings in `distance/simd/{avx512,avx2,neon,wasm_simd128}.rs`. Every safe entry point validates both byte slices against `dim` before any pointer arithmetic (`bbq::assert_payload_shapes`): dispatch guarantees the host's features, never the shape of the data. Each tier carries a `should_panic` case on that guard, and the crate's `simd_length_safety` suite gains the same contract for the dispatched kernel. The x86 tier modules are `pub(crate)`. Their entry points are safe `pub fn`s that call `#[target_feature]` implementations, so exporting them would let safe code outside the crate execute AVX2 or AVX-512 on a host that does not have it. `SimdRuntime::detect()` is the only supported way to reach a tier, and it selects one under a runtime feature probe. `fluxbench` moves to the non-wasm dev-dependency set. It pulls `tokio`, which does not build for wasm32, so the crate's own wasm unit tests could not be built at all — and this change moves the wasm f32 kernels from scalar to SIMD, which would have left them unverified. With the dev-dependency gated, the wasm lib suite builds and runs: 374 pass on wasm32-wasip1 +simd128 under Node's WASI runner, including the simd128 tier parity test. `distance_prepared` rejects a candidate whose header names another dimension or another quantizer instead of scoring it, and the prepared-payload length is computed with checked arithmetic. A candidate buffer long enough to parse is not proof that it belongs to this codec. The wasm simd128 arm is wired into `detect()` and not merely compiled, gated on the same compile-time condition as the tier module, since wasm has no runtime feature probe. The NEON arm is gated to little-endian aarch64 because it reinterprets the little-endian payload; its loads go through `vld1q_u8` plus `vreinterpretq_f32_u8` so they do not depend on 4-byte alignment. Evidence on the host (avx2+fma): `cargo nextest run -p nodedb-vector --all-features --cargo-profile ci --profile ci` passes 468 tests with no failures; the wasm lib suite passes 374 under Node's WASI runner; `cargo check --workspace --all-features` is clean; `cargo fmt --all --check` and `cargo clippy -p nodedb-vector --profile ci --all-targets -- -D warnings` are clean; the crate checks for `aarch64-unknown-linux-gnu` and for `wasm32-wasip1` with and without `+simd128`. Both jobs in CI run on `ubuntu-24.04-arm`, so CI exercises the NEON tier and no x86 tier; the x86 tiers are checked on an x86_64 host, and the 512-bit tier under Intel SDE. No big-endian aarch64 target ships a prebuilt std, so that path is not compiled here. Benches against the reconstruct-and-measure baseline, re-run on this revision: 6.4x at dim 128 and 10.2x at dim 768, the fused path reporting zero allocations and the baseline one `Vec` per candidate. The baseline approximates the pre-fusion pass (fixed 1/sqrt(dim) scale, fixed header offset), so the ratio is indicative rather than exact. The AVX10/256 arm the first revision carried is gone: `AVX512VL` implies `AVX512F`, so no host could reach it while AVX10 detection is absent from `std::arch`. --- Cargo.lock | 1 + nodedb-vector/Cargo.toml | 10 + nodedb-vector/benches/bbq_kernel.rs | 147 +++++++++ nodedb-vector/src/distance/simd/avx2.rs | 59 ++++ nodedb-vector/src/distance/simd/avx512.rs | 38 +++ nodedb-vector/src/distance/simd/bbq.rs | 281 ++++++++++++++++++ nodedb-vector/src/distance/simd/mod.rs | 14 +- nodedb-vector/src/distance/simd/neon.rs | 61 ++++ nodedb-vector/src/distance/simd/runtime.rs | 42 ++- .../src/distance/simd/wasm_simd128.rs | 58 ++++ nodedb-vector/src/hnsw/graph/index/state.rs | 8 + nodedb-vector/src/rerank/codecs/bbq.rs | 199 +++++++++---- .../vector_suite/cases/simd_length_safety.rs | 26 ++ 13 files changed, 886 insertions(+), 58 deletions(-) create mode 100644 nodedb-vector/benches/bbq_kernel.rs create mode 100644 nodedb-vector/src/distance/simd/bbq.rs diff --git a/Cargo.lock b/Cargo.lock index 466f5ce30..0cdae75ef 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4710,6 +4710,7 @@ dependencies = [ "arc-swap", "bytemuck", "crc32c", + "fluxbench", "half", "libc", "memmap2", diff --git a/nodedb-vector/Cargo.toml b/nodedb-vector/Cargo.toml index ca32a29a9..0291ef1ef 100644 --- a/nodedb-vector/Cargo.toml +++ b/nodedb-vector/Cargo.toml @@ -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 diff --git a/nodedb-vector/benches/bbq_kernel.rs b/nodedb-vector/benches/bbq_kernel.rs new file mode 100644 index 000000000..1352c41af --- /dev/null +++ b/nodedb-vector/benches/bbq_kernel.rs @@ -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` 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`. +#[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() +} + +/// Trained codec, prepared query, and 256 encoded candidates for `dim`. +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) +} + +thread_local! { + static STATE_128: (BbqRerank, PreparedQuery, Vec>) = setup(128); + static STATE_768: (BbqRerank, PreparedQuery, Vec>) = setup(768); +} + +/// The baseline: decode the prepared payload to `Vec` 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 = 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); + } +} diff --git a/nodedb-vector/src/distance/simd/avx2.rs b/nodedb-vector/src/distance/simd/avx2.rs index bba842f50..b7e0e0d92 100644 --- a/nodedb-vector/src/distance/simd/avx2.rs +++ b/nodedb-vector/src/distance/simd/avx2.rs @@ -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. @@ -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::()) }; + 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 +} diff --git a/nodedb-vector/src/distance/simd/avx512.rs b/nodedb-vector/src/distance/simd/avx512.rs index 7dddaba4a..8b3ad25b3 100644 --- a/nodedb-vector/src/distance/simd/avx512.rs +++ b/nodedb-vector/src/distance/simd/avx512.rs @@ -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) } @@ -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::()) }; + 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..d57e05fff --- /dev/null +++ b/nodedb-vector/src/distance/simd/bbq.rs @@ -0,0 +1,281 @@ +// 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). +//! +//! Tier coverage is host-dependent, and CI covers only one of them. Both jobs in +//! `.github/workflows/test.yml` run on `ubuntu-24.04-arm`, so CI exercises the +//! NEON tier and never an x86 one. The x86 tiers are checked on an x86_64 host; +//! the 512-bit tier additionally needs hardware or Intel SDE (`sde64 -spr`). +//! Each x86 test returns early, with a printed note, on a host that lacks the +//! feature its tier needs. + +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]]) +} + +/// Panics unless both byte slices hold what `dim` dimensions need: `dim * 4` +/// bytes of little-endian f32 query lanes, and `dim.div_ceil(8)` bytes of sign +/// bits. Every safe entry point runs this before any pointer arithmetic, so a +/// caller handing over short slices gets a panic instead of a read past the end. +pub(super) fn assert_payload_shapes(centered: &[u8], packed: &[u8], dim: usize) { + // Saturating: `dim * 4` must not wrap to a small value and slip past the + // check on an absurd `dim`. + let need_centered = dim.saturating_mul(4); + let need_packed = dim.div_ceil(8); + assert!( + centered.len() >= need_centered, + "bbq l2: centered payload too short: {} bytes, need {need_centered} for dim {dim}", + centered.len() + ); + assert!( + packed.len() >= need_packed, + "bbq l2: packed bits too short: {} bytes, need {need_packed} for dim {dim}", + packed.len() + ); +} + +/// Whole-range scalar kernel: the tail helper applied from 0. +pub fn l2_bbq(centered: &[u8], packed: &[u8], residual_norm: f32, dim: usize) -> f32 { + assert_payload_shapes(centered, packed, dim); + let scale = recon_scale(residual_norm, dim); + l2_scalar_from_bytes(centered, packed, scale, 0, dim).sqrt() +} + +/// 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] +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 lane widths (16 f32 lanes for AVX-512, 8 for + /// AVX2, 4 for NEON and wasm), their multiples, and the partial tails + /// between them. + 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) + } + + #[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}" + ); + } + } + + /// The length guard runs before any pointer arithmetic, so a short slice is + /// rejected rather than read past its end. + #[test] + #[should_panic(expected = "centered payload too short")] + fn short_centered_payload_is_rejected() { + let _ = l2_bbq(&[], &[0u8; 2], 1.0, 16); + } + + #[test] + #[should_panic(expected = "packed bits too short")] + fn short_packed_payload_is_rejected() { + let centered = vec![0u8; 16 * 4]; + let _ = l2_bbq(¢ered, &[], 1.0, 16); + } + + #[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 tier runs here; on an AVX2-only host its test + /// returns early. + #[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}" + ); + // Parity against the scalar kernel directly, not only the oracle. + let scalar = l2_bbq(¢ered, &packed, residual_norm, dim); + assert!( + within_parity(got, scalar as f64), + "{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; + } + // The entry point is safe but does not probe the CPU itself; the + // check above is what keeps this off a host without avx512f. + 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; + } + // The entry point is safe but does not probe the CPU itself; the + // checks above are what keep this off a host without avx2+fma. + check_tier("avx2", crate::distance::simd::avx2::l2_bbq); + } + + /// Each tier's safe entry rejects short slices. The guard runs before + /// the first vector load, so these panic on any host and need no + /// AVX-512 or AVX2 hardware to be valid. + #[test] + #[should_panic(expected = "centered payload too short")] + fn avx512_rejects_empty_slices() { + let _ = crate::distance::simd::avx512::l2_bbq(&[], &[], 1.0, 16); + } + + #[test] + #[should_panic(expected = "centered payload too short")] + fn avx2_rejects_empty_slices() { + let _ = crate::distance::simd::avx2::l2_bbq(&[], &[], 1.0, 16); + } + } + + /// 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() { + // 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 = "aarch64", target_endian = "little"))] + #[test] + #[should_panic(expected = "centered payload too short")] + fn neon_rejects_empty_slices() { + let _ = crate::distance::simd::neon::l2_bbq(&[], &[], 1.0, 16); + } + + #[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..eacdbe56d 100644 --- a/nodedb-vector/src/distance/simd/mod.rs +++ b/nodedb-vector/src/distance/simd/mod.rs @@ -2,15 +2,25 @@ //! Runtime SIMD dispatch for vector distance and bitmap operations. +pub mod bbq; pub mod hamming; pub mod runtime; pub mod scalar; +// The x86 tier modules are crate-private on purpose. Their entry points are +// safe `pub fn`s that call `#[target_feature]` implementations, so exposing the +// module would let safe code outside the crate execute AVX2 or AVX-512 +// instructions on a host that does not have them. `SimdRuntime::detect()` is the +// only supported way to reach a tier, and it selects one under a runtime feature +// probe. NEON is baseline on aarch64 and wasm simd128 is a compile-time gate, so +// those two stay public. #[cfg(target_arch = "x86_64")] -pub mod avx2; +pub(crate) mod avx2; #[cfg(target_arch = "x86_64")] -pub mod avx512; +pub(crate) 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..100569515 100644 --- a/nodedb-vector/src/distance/simd/neon.rs +++ b/nodedb-vector/src/distance/simd/neon.rs @@ -4,6 +4,8 @@ #![cfg(target_arch = "aarch64")] +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(), "neon l2: length mismatch"); unsafe { l2_impl(a, b) } @@ -96,3 +98,62 @@ 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) } +} + +/// 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 `l2_bbq` asserted 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); + // Loaded as bytes and reinterpreted: `vld1q_f32` would require 4-byte + // alignment, which a `&[u8]` payload offset does not promise. `vld1q_u8` + // is byte-aligned by definition. + let q_lo = unsafe { vreinterpretq_f32_u8(vld1q_u8(centered.as_ptr().add(i * 4))) }; + let q_hi = unsafe { vreinterpretq_f32_u8(vld1q_u8(centered.as_ptr().add((i + 4) * 4))) }; + 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..3a392db61 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; @@ -9,12 +10,21 @@ use crate::distance::typed_scalar; #[cfg(target_arch = "x86_64")] use super::{avx2, avx512}; -#[cfg(target_arch = "aarch64")] +// The NEON tier reinterprets the little-endian payload, so a big-endian +// aarch64 build falls through to the scalar kernel instead. +#[cfg(all(target_arch = "aarch64", target_endian = "little"))] 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 +40,8 @@ pub struct SimdRuntime { pub l2_squared_bf16: HalfFn, pub cosine_distance_bf16: HalfFn, pub neg_inner_product_bf16: HalfFn, + /// Fused BBQ decode-and-distance kernel (see `BbqFn`). + pub l2_bbq: BbqFn, } impl SimdRuntime { @@ -57,6 +69,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,12 +91,13 @@ 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; } } - #[cfg(target_arch = "aarch64")] + #[cfg(all(target_arch = "aarch64", target_endian = "little"))] { let rt = Self { l2_squared: neon::l2_squared, @@ -97,6 +111,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 +152,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..0656d0c21 100644 --- a/nodedb-vector/src/distance/simd/wasm_simd128.rs +++ b/nodedb-vector/src/distance/simd/wasm_simd128.rs @@ -17,6 +17,7 @@ //! explicit `unsafe {}` blocks even within `unsafe fn` bodies, as required by //! the `unsafe_op_in_unsafe_fn` lint default in edition 2024. +use super::bbq::{assert_payload_shapes, l2_scalar_from_bytes, recon_scale}; use std::arch::wasm32::*; /// L2-squared distance between two F32 slices using WASM SIMD128. @@ -129,6 +130,63 @@ unsafe fn ip_impl(a: &[f32], b: &[f32]) -> f32 { -dot } +/// Lane masks for every packed byte, MSB-first: `MASKS[byte][0]` covers dims +/// 0..4 of the byte, `MASKS[byte][1]` covers dims 4..8. `u32::MAX` marks a set +/// sign bit, so `v128_bitselect` takes `+scale` for that lane. +static MASKS: [[[u32; 4]; 2]; 256] = build_masks(); + +const fn build_masks() -> [[[u32; 4]; 2]; 256] { + let mut table = [[[0u32; 4]; 2]; 256]; + let mut byte = 0usize; + while byte < 256 { + let mut half = 0usize; + while half < 2 { + let mut lane = 0usize; + while lane < 4 { + let bit = (byte >> (7 - (half * 4 + lane))) & 1; + table[byte][half][lane] = if bit == 1 { u32::MAX } else { 0 }; + lane += 1; + } + half += 1; + } + byte += 1; + } + table +} + +/// 4-lane tier (`simd128`): `v128_bitselect` between `+scale` and `-scale`. +/// Safe entry for `SimdRuntime`; the feature gate is this module's `#![cfg]`. +pub fn l2_bbq(centered: &[u8], packed: &[u8], residual_norm: f32, dim: usize) -> f32 { + assert_payload_shapes(centered, packed, dim); + + 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 { + // `i` advances by 4, so `i % 8` is 0 or 4 and one byte holds all four + // dims of this step. + let lanes = MASKS[packed[i / 8] as usize][(i % 8) / 4]; + let mask = u32x4(lanes[0], lanes[1], lanes[2], lanes[3]); + // SAFETY: `i + 4 <= dim`, and `assert_payload_shapes` checked + // `centered.len() >= dim * 4`. + 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() +} + #[cfg(target_arch = "wasm32")] #[cfg(test)] mod tests { diff --git a/nodedb-vector/src/hnsw/graph/index/state.rs b/nodedb-vector/src/hnsw/graph/index/state.rs index ae2c62440..7d13547b2 100644 --- a/nodedb-vector/src/hnsw/graph/index/state.rs +++ b/nodedb-vector/src/hnsw/graph/index/state.rs @@ -128,6 +128,8 @@ impl HnswIndex { #[cfg(test)] mod tests { + // Only the backing test uses this, and that test is gated off on wasm32. + #[cfg(not(target_arch = "wasm32"))] use std::sync::Arc; use super::{HnswIndex, HnswParams}; @@ -255,6 +257,9 @@ mod tests { /// Rerank must see the vector for a narrow dtype too. Before the `Cow` return /// this yielded `None` for F16/BF16, which made the FP32 rerank path fail with /// "fetch_vector returned None" for every candidate in the collection. + /// + /// `get_vector_or_backing` is gated off on wasm32, so the test is too. + #[cfg(not(target_arch = "wasm32"))] #[test] fn get_vector_or_backing_serves_narrow_dtypes() { for dtype in [VectorStorageDtype::F16, VectorStorageDtype::BF16] { @@ -307,6 +312,9 @@ mod tests { /// backing is the worst case: the graph looks healthy, so search proceeds and /// then scores a node that has no vector — which is what made one poisoned /// segment panic the daemon on every query. + /// + /// `segment_backing` is gated off on wasm32, so the test is too. + #[cfg(not(target_arch = "wasm32"))] #[test] fn with_backing_refuses_a_backing_that_cannot_serve_the_index() { use crate::segment_backing::VectorSegmentBacking; diff --git a/nodedb-vector/src/rerank/codecs/bbq.rs b/nodedb-vector/src/rerank/codecs/bbq.rs index 444ccf187..92e6a5bfc 100644 --- a/nodedb-vector/src/rerank/codecs/bbq.rs +++ b/nodedb-vector/src/rerank/codecs/bbq.rs @@ -17,7 +17,7 @@ use nodedb_codec::vector_quant::bbq::BbqCodec; use nodedb_codec::vector_quant::codec::VectorCodec as _; -use nodedb_codec::vector_quant::layout::UnifiedQuantizedVectorRef; +use nodedb_codec::vector_quant::layout::{QuantMode, UnifiedQuantizedVectorRef}; use crate::{ rerank::codec::{CodecName, PreparedQuery, RerankCodec}, @@ -36,50 +36,10 @@ 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. +/// `None` when `dim` is large enough that the length itself overflows. +fn payload_len(dim: usize) -> Option { + dim.checked_mul(4)?.checked_add(4) } // ── BbqRerank ───────────────────────────────────────────────────────────────── @@ -197,22 +157,54 @@ impl RerankCodec for BbqRerank { } }; - let (_query_norm, centered) = decode_payload(payload, self.dim)?; + let Some(expected) = payload_len(self.dim) else { + return Err(RerankError::BadInput(format!( + "bbq distance: dim {} overflows the prepared payload length", + 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| { RerankError::BadInput(format!("bbq distance: failed to parse encoded bytes: {e}")) })?; + // The header decides how the candidate is reconstructed, so a candidate + // encoded for another dimension or another quantizer must be rejected + // rather than scored: `from_bytes` only proves the buffer is long enough + // to parse, not that it belongs to this codec. 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) + if usize::from(header.dim) != self.dim { + return Err(RerankError::BadInput(format!( + "bbq distance: candidate dim {} != codec dim {}", + header.dim, self.dim + ))); + } + if header.quant_mode != QuantMode::Bbq as u16 { + return Err(RerankError::BadInput(format!( + "bbq distance: candidate quant mode {} is not BBQ ({})", + header.quant_mode, + QuantMode::Bbq as u16 + ))); + } + + // 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 { @@ -291,6 +283,62 @@ mod tests { assert!(dist >= 0.0, "distance must be non-negative, got {dist}"); } + /// Pins the value `distance_prepared` returns: the asymmetric L2 between the + /// exact centred query and the `±residual_norm/√dim` reconstruction, computed + /// here in f64 by plain indexing. The fused kernel sits behind this seam, so + /// the reference is built from the wire layout rather than from the kernel. + #[test] + fn distance_prepared_matches_the_unfused_l2() { + let codec = trained(); + let v = det_vec(7, DIM); + let enc = codec.encode(&v).expect("encode"); + let prep = codec.prepare_query(&v).expect("prepare_query"); + let got = codec.distance_prepared(&prep, &enc).expect("distance"); + + let PreparedQuery::Bytes(payload) = &prep else { + panic!("prepare_query must yield PreparedQuery::Bytes"); + }; + + // Wire layout: 32-byte `QuantHeader` (residual_norm at bytes 8..12), + // then `dim.div_ceil(8)` sign-packed bytes; the prepared payload is the + // 4-byte alpha followed by `dim` centred f32s. The offsets are pinned on + // purpose: a change on either side of the seam must fail this test. + let residual_norm = f32::from_le_bytes([enc[8], enc[9], enc[10], enc[11]]); + let packed = &enc[32..32 + DIM.div_ceil(8)]; + let centered: Vec = payload[4..] + .as_chunks::<4>() + .0 + .iter() + .map(|b| f32::from_le_bytes(*b) as f64) + .collect(); + assert_eq!(centered.len(), DIM, "centred query must carry `dim` lanes"); + + let scale = residual_norm as f64 / (DIM as f64).sqrt(); + let expected = centered + .iter() + .enumerate() + .map(|(i, q)| { + let bit = (packed[i / 8] >> (7 - (i % 8))) & 1; + let recon = if bit != 0 { scale } else { -scale }; + (q - recon).powi(2) + }) + .sum::() + .sqrt(); + + // A degenerate reference would make the comparison vacuous. + assert!( + expected > 1e-3, + "reference distance collapsed to {expected}" + ); + + let rel = 1e-4f64; + let abs = 1e-6f64; + assert!( + ((got as f64) - expected).abs() <= abs.max(rel * expected.abs()), + "distance_prepared = {got}, unfused L2 = {expected}" + ); + } + #[test] fn encode_before_train_returns_not_trained() { let codec = BbqRerank::new(DIM, DEFAULT_OVERSAMPLE); @@ -329,6 +377,49 @@ mod tests { ); } + /// A candidate encoded by a codec of another dimension parses against this + /// codec's packed length when its buffer is long enough, so only the header + /// check can reject it. Scoring it would silently return a wrong distance. + #[test] + fn candidate_encoded_for_another_dim_is_rejected() { + let codec = trained(); + let prep = codec + .prepare_query(&det_vec(0, DIM)) + .expect("prepare_query"); + + let other_dim = DIM * 2; + let vecs: Vec> = (0..N).map(|i| det_vec(i, other_dim)).collect(); + let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect(); + let mut other = BbqRerank::new(other_dim, DEFAULT_OVERSAMPLE); + other.train(&refs).expect("train must succeed"); + let enc = other.encode(&det_vec(0, other_dim)).expect("encode"); + + let err = codec.distance_prepared(&prep, &enc).unwrap_err(); + let msg = format!("{err}"); + assert!( + msg.contains("candidate dim"), + "expected a dimension mismatch, got: {msg}" + ); + } + + #[test] + fn candidate_with_another_quant_mode_is_rejected() { + let codec = trained(); + let v = det_vec(0, DIM); + let prep = codec.prepare_query(&v).expect("prepare_query"); + let mut enc = codec.encode(&v).expect("encode"); + + // Header bytes 0..2 carry the quant mode; Sq8 has the same header size. + enc[0..2].copy_from_slice(&(QuantMode::Sq8 as u16).to_le_bytes()); + + let err = codec.distance_prepared(&prep, &enc).unwrap_err(); + let msg = format!("{err}"); + assert!( + msg.contains("quant mode"), + "expected a quant-mode mismatch, got: {msg}" + ); + } + #[test] fn prepare_query_wrong_dim_fails() { let codec = trained(); diff --git a/nodedb-vector/tests/vector_suite/cases/simd_length_safety.rs b/nodedb-vector/tests/vector_suite/cases/simd_length_safety.rs index 83033451c..e69003ce9 100644 --- a/nodedb-vector/tests/vector_suite/cases/simd_length_safety.rs +++ b/nodedb-vector/tests/vector_suite/cases/simd_length_safety.rs @@ -58,3 +58,29 @@ fn l2_rejects_swapped_mismatch() { "distance() must reject length mismatch in either argument order" ); } + +/// The BBQ fused kernel is the same contract on byte slices: it reads the +/// centered query straight out of the prepared payload, so `centered` must +/// carry `dim * 4` bytes and `packed` must carry `dim.div_ceil(8)` sign bytes. +/// Every tier's safe entry validates both before any vector load, so an +/// external safe caller gets a panic — never a read past the slice. +#[test] +fn bbq_rejects_short_slices() { + let kernel = nodedb_vector::distance::simd::runtime::runtime(); + let dim = 16usize; + + let short_centered = std::panic::catch_unwind(|| (kernel.l2_bbq)(&[], &[], 1.0, dim)); + assert!( + short_centered.is_err(), + "dispatched bbq kernel ({}) must reject empty slices", + kernel.name + ); + + let centered = vec![0u8; dim * 4]; + let short_packed = std::panic::catch_unwind(|| (kernel.l2_bbq)(¢ered, &[], 1.0, dim)); + assert!( + short_packed.is_err(), + "dispatched bbq kernel ({}) must reject a short packed slice", + kernel.name + ); +}