#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
#[cfg(test)]
use std::sync::Arc;
use wide::f32x8;
use crate::superfile::vector::rerank_codec::{RerankCodec, SQ16_FIXED_OFFSET, SQ16_FIXED_SCALE};
#[cfg(target_arch = "x86_64")]
use crate::superfile::vector::simd_dispatch::{avx2_enabled, avx512_enabled};
pub(crate) const SQ8_RESIDUAL_DIVISOR: f32 = 16.0;
const F32X8_LANES: usize = 8;
#[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
const AVX512_F32_LANES: usize = 16;
pub(crate) const CENTROID_BATCH_LANES: usize = AVX512_F32_LANES;
const F32_BYTES: usize = 4;
pub(crate) const COSINE_DISTANCE_BASE: f32 = 1.0;
pub(crate) const L2_CROSS_TERM_COEFF: f32 = 2.0;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Metric {
Cosine,
L2Sq,
NegDot,
}
#[inline]
pub fn distance(metric: Metric, a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
match metric {
Metric::Cosine => COSINE_DISTANCE_BASE - dot(a, b),
Metric::L2Sq => l2_sq(a, b),
Metric::NegDot => -dot(a, b),
}
}
#[inline]
pub fn dot(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
#[cfg(target_arch = "x86_64")]
if avx512_enabled() {
return unsafe { dot_avx512(a, b) };
}
dot_wide(a, b)
}
#[inline]
pub(crate) fn l2_sq(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
#[cfg(target_arch = "x86_64")]
if avx512_enabled() {
return unsafe { l2_sq_avx512(a, b) };
}
l2_sq_wide(a, b)
}
pub(crate) fn transpose_centroids_cluster_major(
centroids: &[f32],
n_cent: usize,
dim: usize,
) -> Vec<f32> {
debug_assert_eq!(centroids.len(), n_cent * dim);
let n_blocks = n_cent.div_ceil(CENTROID_BATCH_LANES);
let mut transposed = vec![0f32; n_blocks * dim * CENTROID_BATCH_LANES];
for block in 0..n_blocks {
let centroid_base = block * CENTROID_BATCH_LANES;
let block_base = block * dim * CENTROID_BATCH_LANES;
for d in 0..dim {
let dst = block_base + d * CENTROID_BATCH_LANES;
for lane in 0..CENTROID_BATCH_LANES {
let centroid = centroid_base + lane;
if centroid < n_cent {
transposed[dst + lane] = centroids[centroid * dim + d];
}
}
}
}
transposed
}
#[inline]
pub(crate) fn relative_score_window(base: f32, slack: f32) -> f32 {
base + base.abs().max(f32::EPSILON) * slack.max(0.0)
}
#[inline]
pub(crate) fn insert_ranked(top: &mut Vec<(u32, f32)>, k: usize, centroid: u32, score: f32) {
if top.len() == k && score >= top[k - 1].1 {
return;
}
let pos = top
.iter()
.position(|&(_, s)| score < s)
.unwrap_or(top.len());
top.insert(pos, (centroid, score));
top.truncate(k);
}
#[inline]
fn for_each_centroid_block_scores(
metric: Metric,
query: &[f32],
transposed: &[f32],
n_cent: usize,
dim: usize,
mut reduce: impl FnMut(usize, &[f32]),
) {
debug_assert_eq!(query.len(), dim);
debug_assert_eq!(
transposed.len(),
n_cent.div_ceil(CENTROID_BATCH_LANES) * dim * CENTROID_BATCH_LANES
);
let n_blocks = n_cent.div_ceil(CENTROID_BATCH_LANES);
#[cfg(target_arch = "x86_64")]
if avx512_enabled() {
for block in 0..n_blocks {
let scores = unsafe {
score_centroid_block16_transposed_avx512(metric, query, transposed, dim, block)
};
reduce(block * CENTROID_BATCH_LANES, &scores);
}
return;
}
for block in 0..n_blocks {
let base_centroid = block * CENTROID_BATCH_LANES;
for half in 0..CENTROID_BATCH_LANES / F32X8_LANES {
let lane_offset = half * F32X8_LANES;
let scores = score_centroid_block8_transposed_wide(
metric,
query,
transposed,
dim,
block,
lane_offset,
);
reduce(base_centroid + lane_offset, &scores);
}
}
}
pub(crate) fn nearest_centroid_transposed(
metric: Metric,
query: &[f32],
transposed: &[f32],
n_cent: usize,
dim: usize,
) -> (u32, f32) {
debug_assert!(n_cent > 0);
let mut best = (0u32, f32::INFINITY);
for_each_centroid_block_scores(metric, query, transposed, n_cent, dim, |base, scores| {
for (lane, &score) in scores.iter().enumerate() {
let centroid = base + lane;
if centroid < n_cent && score < best.1 {
best = (centroid as u32, score);
}
}
});
best
}
#[cfg(test)]
pub(crate) fn nearest_two_centroids_transposed(
metric: Metric,
query: &[f32],
transposed: &[f32],
n_cent: usize,
dim: usize,
counts: Option<&[u32]>,
) -> Option<((u32, f32), Option<(u32, f32)>)> {
let top = nearest_k_centroids_transposed(metric, query, transposed, n_cent, dim, counts, 2);
let mut it = top.into_iter();
it.next().map(|best| (best, it.next()))
}
pub(crate) fn nearest_k_centroids_transposed(
metric: Metric,
query: &[f32],
transposed: &[f32],
n_cent: usize,
dim: usize,
counts: Option<&[u32]>,
k: usize,
) -> Vec<(u32, f32)> {
debug_assert!(counts.is_none_or(|counts| counts.len() >= n_cent));
let mut top: Vec<(u32, f32)> = Vec::with_capacity(k.saturating_add(1));
if k == 0 {
return top;
}
for_each_centroid_block_scores(metric, query, transposed, n_cent, dim, |base, scores| {
for (lane, &score) in scores.iter().enumerate() {
let centroid = base + lane;
if centroid < n_cent && centroid_included(counts, centroid) {
insert_ranked(&mut top, k, centroid as u32, score);
}
}
});
top
}
#[inline]
fn centroid_included(counts: Option<&[u32]>, centroid: usize) -> bool {
counts.is_none_or(|counts| counts[centroid] != 0)
}
pub(crate) fn all_centroid_scores_transposed(
metric: Metric,
query: &[f32],
transposed: &[f32],
n_cent: usize,
dim: usize,
) -> Vec<f32> {
let mut out = vec![0f32; n_cent];
for_each_centroid_block_scores(metric, query, transposed, n_cent, dim, |base, scores| {
for (lane, &score) in scores.iter().enumerate() {
let centroid = base + lane;
if centroid < n_cent {
out[centroid] = score;
}
}
});
out
}
pub(crate) fn nearest_k_centroids_bytes(
metric: Metric,
query: &[f32],
centroids_bytes: &[u8],
n_cent: usize,
dim: usize,
k: usize,
) -> Vec<(u32, f32)> {
let stride = dim * 4;
debug_assert!(centroids_bytes.len() >= n_cent * stride);
let mut top: Vec<(u32, f32)> = Vec::with_capacity(k.saturating_add(1));
if k == 0 {
return top;
}
for c in 0..n_cent {
let bytes = ¢roids_bytes[c * stride..(c + 1) * stride];
insert_ranked(&mut top, k, c as u32, distance_bytes(metric, query, bytes));
}
top
}
#[inline]
fn score_centroid_block8_transposed_wide(
metric: Metric,
query: &[f32],
transposed: &[f32],
dim: usize,
block: usize,
lane_offset: usize,
) -> [f32; F32X8_LANES] {
let mut acc = f32x8::ZERO;
let block_base = block * dim * CENTROID_BATCH_LANES;
for (d, &q_scalar) in query[..dim].iter().enumerate() {
let q = f32x8::splat(q_scalar);
let row = block_base + d * CENTROID_BATCH_LANES + lane_offset;
let c = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&transposed[row..row + F32X8_LANES])
.expect("transposed centroid row has 8-lane half"),
);
match metric {
Metric::L2Sq => {
let diff = q - c;
acc += diff * diff;
}
Metric::Cosine | Metric::NegDot => {
acc += q * c;
}
}
}
let mut scores = acc.to_array();
match metric {
Metric::Cosine => {
for score in &mut scores {
*score = COSINE_DISTANCE_BASE - *score;
}
}
Metric::NegDot => {
for score in &mut scores {
*score = -*score;
}
}
Metric::L2Sq => {}
}
scores
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn score_centroid_block16_transposed_avx512(
metric: Metric,
query: &[f32],
transposed: &[f32],
dim: usize,
block: usize,
) -> [f32; AVX512_F32_LANES] {
unsafe {
let mut acc = _mm512_setzero_ps();
let block_base = block * dim * CENTROID_BATCH_LANES;
for (d, &q_scalar) in query[..dim].iter().enumerate() {
let q = _mm512_set1_ps(q_scalar);
let row = block_base + d * CENTROID_BATCH_LANES;
let c = _mm512_loadu_ps(transposed.as_ptr().add(row));
match metric {
Metric::L2Sq => {
let diff = _mm512_sub_ps(q, c);
acc = _mm512_fmadd_ps(diff, diff, acc);
}
Metric::Cosine | Metric::NegDot => {
acc = _mm512_fmadd_ps(q, c, acc);
}
}
}
let mut scores = [0f32; AVX512_F32_LANES];
_mm512_storeu_ps(scores.as_mut_ptr(), acc);
match metric {
Metric::Cosine => {
for score in &mut scores {
*score = COSINE_DISTANCE_BASE - *score;
}
}
Metric::NegDot => {
for score in &mut scores {
*score = -*score;
}
}
Metric::L2Sq => {}
}
scores
}
}
#[inline]
pub(crate) fn sum_f32(a: &[f32]) -> f32 {
#[cfg(target_arch = "x86_64")]
if avx512_enabled() {
return unsafe { sum_f32_avx512(a) };
}
sum_f32_wide(a)
}
#[inline]
fn dot_wide(a: &[f32], b: &[f32]) -> f32 {
let chunks_a = a.chunks_exact(F32X8_LANES);
let chunks_b = b.chunks_exact(F32X8_LANES);
let tail_a = chunks_a.remainder();
let tail_b = chunks_b.remainder();
let mut acc = f32x8::ZERO;
for (ca, cb) in chunks_a.zip(chunks_b) {
let va = f32x8::from(
<[f32; F32X8_LANES]>::try_from(ca).expect("chunks_exact(8) yields slices of length 8"),
);
let vb = f32x8::from(
<[f32; F32X8_LANES]>::try_from(cb).expect("chunks_exact(8) yields slices of length 8"),
);
acc += va * vb;
}
let mut sum: f32 = acc.reduce_add();
for (x, y) in tail_a.iter().zip(tail_b.iter()) {
sum += x * y;
}
sum
}
#[inline]
fn l2_sq_wide(a: &[f32], b: &[f32]) -> f32 {
let chunks_a = a.chunks_exact(F32X8_LANES);
let chunks_b = b.chunks_exact(F32X8_LANES);
let tail_a = chunks_a.remainder();
let tail_b = chunks_b.remainder();
let mut acc = f32x8::ZERO;
for (ca, cb) in chunks_a.zip(chunks_b) {
let va = f32x8::from(
<[f32; F32X8_LANES]>::try_from(ca).expect("chunks_exact(8) yields slices of length 8"),
);
let vb = f32x8::from(
<[f32; F32X8_LANES]>::try_from(cb).expect("chunks_exact(8) yields slices of length 8"),
);
let d = va - vb;
acc += d * d;
}
let mut sum: f32 = acc.reduce_add();
for (x, y) in tail_a.iter().zip(tail_b.iter()) {
let d = x - y;
sum += d * d;
}
sum
}
#[inline]
fn sum_f32_wide(a: &[f32]) -> f32 {
let chunks = a.chunks_exact(F32X8_LANES);
let tail = chunks.remainder();
let mut acc = f32x8::ZERO;
for c in chunks {
let va = f32x8::from(
<[f32; F32X8_LANES]>::try_from(c).expect("chunks_exact(8) yields slices of length 8"),
);
acc += va;
}
let mut sum: f32 = acc.reduce_add();
for &x in tail {
sum += x;
}
sum
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn dot_avx512(a: &[f32], b: &[f32]) -> f32 {
let n = a.len();
unsafe {
let mut acc = _mm512_setzero_ps();
let mut i = 0;
while i + AVX512_F32_LANES <= n {
let va = _mm512_loadu_ps(a.as_ptr().add(i));
let vb = _mm512_loadu_ps(b.as_ptr().add(i));
acc = _mm512_fmadd_ps(va, vb, acc);
i += AVX512_F32_LANES;
}
let mut sum = _mm512_reduce_add_ps(acc);
while i < n {
sum += a[i] * b[i];
i += 1;
}
sum
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn l2_sq_avx512(a: &[f32], b: &[f32]) -> f32 {
let n = a.len();
unsafe {
let mut acc = _mm512_setzero_ps();
let mut i = 0;
while i + AVX512_F32_LANES <= n {
let va = _mm512_loadu_ps(a.as_ptr().add(i));
let vb = _mm512_loadu_ps(b.as_ptr().add(i));
let d = _mm512_sub_ps(va, vb);
acc = _mm512_fmadd_ps(d, d, acc);
i += AVX512_F32_LANES;
}
let mut sum = _mm512_reduce_add_ps(acc);
while i < n {
let d = a[i] - b[i];
sum += d * d;
i += 1;
}
sum
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn sum_f32_avx512(a: &[f32]) -> f32 {
let n = a.len();
unsafe {
let mut acc = _mm512_setzero_ps();
let mut i = 0;
while i + AVX512_F32_LANES <= n {
let va = _mm512_loadu_ps(a.as_ptr().add(i));
acc = _mm512_add_ps(va, acc);
i += AVX512_F32_LANES;
}
let mut sum = _mm512_reduce_add_ps(acc);
while i < n {
sum += a[i];
i += 1;
}
sum
}
}
#[inline]
pub fn distance_bytes(metric: Metric, query: &[f32], bytes: &[u8]) -> f32 {
debug_assert_eq!(query.len() * F32_BYTES, bytes.len());
match metric {
Metric::Cosine => COSINE_DISTANCE_BASE - dot_bytes(query, bytes),
Metric::L2Sq => l2_sq_bytes(query, bytes),
Metric::NegDot => -dot_bytes(query, bytes),
}
}
#[inline]
pub fn dot_bytes(query: &[f32], bytes: &[u8]) -> f32 {
if let Ok(v) = bytemuck::try_cast_slice::<u8, f32>(bytes) {
return dot(query, v);
}
dot_le_bytes_unaligned(query, bytes)
}
#[inline]
pub fn l2_sq_bytes(query: &[f32], bytes: &[u8]) -> f32 {
if let Ok(v) = bytemuck::try_cast_slice::<u8, f32>(bytes) {
return l2_sq(query, v);
}
l2_sq_le_bytes_unaligned(query, bytes)
}
#[inline]
fn dot_le_bytes_unaligned(query: &[f32], bytes: &[u8]) -> f32 {
let mut acc = f32x8::ZERO;
let mut i = 0;
while i + F32X8_LANES <= query.len() {
let qc: [f32; F32X8_LANES] = query[i..i + F32X8_LANES]
.try_into()
.expect("slice [i..i+8] has length 8");
let mut bc = [0f32; F32X8_LANES];
for (j, slot) in bc.iter_mut().enumerate() {
let off = (i + j) * F32_BYTES;
*slot =
f32::from_le_bytes([bytes[off], bytes[off + 1], bytes[off + 2], bytes[off + 3]]);
}
let qv = f32x8::from(qc);
let bv = f32x8::from(bc);
acc += qv * bv;
i += F32X8_LANES;
}
let mut sum = acc.reduce_add();
while i < query.len() {
let off = i * F32_BYTES;
let b = f32::from_le_bytes([bytes[off], bytes[off + 1], bytes[off + 2], bytes[off + 3]]);
sum += query[i] * b;
i += 1;
}
sum
}
#[inline]
fn l2_sq_le_bytes_unaligned(query: &[f32], bytes: &[u8]) -> f32 {
let mut acc = f32x8::ZERO;
let mut i = 0;
while i + F32X8_LANES <= query.len() {
let qc: [f32; F32X8_LANES] = query[i..i + F32X8_LANES]
.try_into()
.expect("slice [i..i+8] has length 8");
let mut bc = [0f32; F32X8_LANES];
for (j, slot) in bc.iter_mut().enumerate() {
let off = (i + j) * F32_BYTES;
*slot =
f32::from_le_bytes([bytes[off], bytes[off + 1], bytes[off + 2], bytes[off + 3]]);
}
let qv = f32x8::from(qc);
let bv = f32x8::from(bc);
let d = qv - bv;
acc += d * d;
i += F32X8_LANES;
}
let mut sum = acc.reduce_add();
while i < query.len() {
let off = i * F32_BYTES;
let b = f32::from_le_bytes([bytes[off], bytes[off + 1], bytes[off + 2], bytes[off + 3]]);
let d = query[i] - b;
sum += d * d;
i += 1;
}
sum
}
#[inline]
pub(crate) fn distance_bytes_codec(
metric: Metric,
codec: RerankCodec,
query: &[f32],
bytes: &[u8],
) -> f32 {
match codec {
RerankCodec::Fp32 => distance_bytes(metric, query, bytes),
RerankCodec::Sq8Residual | RerankCodec::Sq8FixedResidual => {
unreachable!(
"distance_bytes_codec called with residual-family codec — rerank goes \
through dedicated kernels (need per-column scale/offset + per-doc \
norm context)"
)
}
RerankCodec::Sq16 => {
unreachable!(
"distance_bytes_codec called with Sq16 — Sq16 rerank goes through \
Sq16Kernel (u16 → f32 dequant front on the fp32 distance path)"
)
}
RerankCodec::RabitqOnly => {
unreachable!(
"distance_bytes_codec called with RabitqOnly — RabitqOnly columns \
carry no full[] region to score against"
)
}
}
}
#[cfg(test)]
pub(crate) struct Sq8Kernel {
metric: Metric,
dim: usize,
q_prime: Vec<f32>,
q_dot_offset: f32,
q_norm_sq: f32,
per_doc_norms: Option<Arc<[f32]>>,
}
#[cfg(test)]
impl Sq8Kernel {
pub fn new(
metric: Metric,
query: &[f32],
scale: &[f32],
offset: &[f32],
per_doc_norms: Option<Arc<[f32]>>,
) -> Self {
let dim = query.len();
debug_assert_eq!(scale.len(), dim);
debug_assert_eq!(offset.len(), dim);
let mut q_prime = vec![0.0f32; dim];
let mut q_dot_offset_acc = f32x8::ZERO;
let mut i = 0;
while i + F32X8_LANES <= dim {
let qc = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&query[i..i + F32X8_LANES]).expect("len-8 slice"),
);
let sc = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&scale[i..i + F32X8_LANES]).expect("len-8 slice"),
);
let oc = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&offset[i..i + F32X8_LANES]).expect("len-8 slice"),
);
let qp = qc * sc;
q_prime[i..i + F32X8_LANES].copy_from_slice(&qp.to_array());
q_dot_offset_acc += qc * oc;
i += F32X8_LANES;
}
let mut q_dot_offset: f32 = q_dot_offset_acc.reduce_add();
while i < dim {
q_prime[i] = query[i] * scale[i];
q_dot_offset += query[i] * offset[i];
i += 1;
}
let q_norm_sq = match metric {
Metric::L2Sq => dot(query, query),
Metric::Cosine | Metric::NegDot => 0.0,
};
Self {
metric,
dim,
q_prime,
q_dot_offset,
q_norm_sq,
per_doc_norms,
}
}
#[inline]
pub fn distance_at(&self, pos: u32, code_bytes: &[u8]) -> f32 {
let norm = self.per_doc_norms.as_ref().map(|norms| norms[pos as usize]);
self.distance_with_norm(code_bytes, norm)
}
#[inline]
pub fn distance_with_norm(&self, code_bytes: &[u8], norm: Option<f32>) -> f32 {
debug_assert_eq!(code_bytes.len(), self.dim);
let qp_code_dot = sq8_dot(&self.q_prime, code_bytes, self.dim);
let dot = qp_code_dot + self.q_dot_offset;
match self.metric {
Metric::Cosine => {
let x_norm = norm
.expect("Sq8Kernel + Cosine requires per_doc_norms")
.sqrt();
if x_norm > 0.0 {
COSINE_DISTANCE_BASE - dot / x_norm
} else {
COSINE_DISTANCE_BASE - dot
}
}
Metric::NegDot => -dot,
Metric::L2Sq => {
let x_norm_sq = norm.expect("Sq8Kernel + L2Sq requires per_doc_norms");
self.q_norm_sq - L2_CROSS_TERM_COEFF * dot + x_norm_sq
}
}
}
}
pub(crate) struct Sq8ResidualKernel {
metric: Metric,
dim: usize,
q_code: Vec<f32>,
q_residual: Vec<f32>,
q_dot_offset: f32,
q_norm_sq: f32,
}
impl Sq8ResidualKernel {
pub fn new(
metric: Metric,
query: &[f32],
scale: &[f32],
offset: &[f32],
residual_divisor: f32,
) -> Self {
let dim = query.len();
debug_assert_eq!(scale.len(), dim);
debug_assert_eq!(offset.len(), dim);
debug_assert!(residual_divisor > 0.0);
let mut q_code = vec![0.0f32; dim];
let mut q_residual = vec![0.0f32; dim];
let inv_residual_divisor = 1.0 / residual_divisor;
let mut q_dot_offset_acc = f32x8::ZERO;
let mut i = 0;
while i + F32X8_LANES <= dim {
let qc = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&query[i..i + F32X8_LANES]).expect("len-8 slice"),
);
let sc = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&scale[i..i + F32X8_LANES]).expect("len-8 slice"),
);
let oc = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&offset[i..i + F32X8_LANES]).expect("len-8 slice"),
);
let q_code_v = qc * sc;
let q_residual_v = q_code_v * f32x8::splat(inv_residual_divisor);
q_code[i..i + F32X8_LANES].copy_from_slice(&q_code_v.to_array());
q_residual[i..i + F32X8_LANES].copy_from_slice(&q_residual_v.to_array());
q_dot_offset_acc += qc * oc;
i += F32X8_LANES;
}
let mut q_dot_offset = q_dot_offset_acc.reduce_add();
while i < dim {
let q_scale = query[i] * scale[i];
q_code[i] = q_scale;
q_residual[i] = q_scale * inv_residual_divisor;
q_dot_offset += query[i] * offset[i];
i += 1;
}
let q_norm_sq = match metric {
Metric::L2Sq => dot(query, query),
Metric::Cosine | Metric::NegDot => 0.0,
};
Self {
metric,
dim,
q_code,
q_residual,
q_dot_offset,
q_norm_sq,
}
}
#[inline]
pub fn distance_with_norm(
&self,
code_bytes: &[u8],
residual_bytes: &[u8],
norm: Option<f32>,
) -> f32 {
debug_assert_eq!(code_bytes.len(), self.dim);
debug_assert_eq!(residual_bytes.len(), self.dim);
let mut acc = f32x8::ZERO;
let mut i = 0;
while i + F32X8_LANES <= self.dim {
let qc: [f32; F32X8_LANES] = self.q_code[i..i + F32X8_LANES]
.try_into()
.expect("q_code[i..i+8] len 8");
let qr: [f32; F32X8_LANES] = self.q_residual[i..i + F32X8_LANES]
.try_into()
.expect("q_residual[i..i+8] len 8");
let mut code = [0f32; F32X8_LANES];
let mut residual = [0f32; F32X8_LANES];
for j in 0..F32X8_LANES {
code[j] = code_bytes[i + j] as f32;
residual[j] = i8::from_le_bytes([residual_bytes[i + j]]) as f32;
}
acc += f32x8::from(qc) * f32x8::from(code);
acc += f32x8::from(qr) * f32x8::from(residual);
i += F32X8_LANES;
}
let mut cross = acc.reduce_add();
while i < self.dim {
cross += self.q_code[i] * (code_bytes[i] as f32);
cross += self.q_residual[i] * (i8::from_le_bytes([residual_bytes[i]]) as f32);
i += 1;
}
let dot = cross + self.q_dot_offset;
match self.metric {
Metric::Cosine => {
let x_norm = norm
.expect("Sq8ResidualKernel + Cosine requires per_doc_norms")
.sqrt();
if x_norm > 0.0 {
COSINE_DISTANCE_BASE - dot / x_norm
} else {
COSINE_DISTANCE_BASE - dot
}
}
Metric::NegDot => -dot,
Metric::L2Sq => {
let x_norm_sq = norm.expect("Sq8ResidualKernel + L2Sq requires per_doc_norms");
self.q_norm_sq - L2_CROSS_TERM_COEFF * dot + x_norm_sq
}
}
}
}
pub(crate) struct Sq16Kernel {
metric: Metric,
dim: usize,
q_prime: Vec<f32>,
q_dot_offset: f32,
q_norm_sq: f32,
}
impl Sq16Kernel {
pub fn new(metric: Metric, query: &[f32]) -> Self {
let dim = query.len();
let mut q_prime = vec![0.0f32; dim];
let scale_v = f32x8::splat(SQ16_FIXED_SCALE);
let mut q_sum_acc = f32x8::ZERO;
let mut i = 0;
while i + F32X8_LANES <= dim {
let qc = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&query[i..i + F32X8_LANES]).expect("len-8 slice"),
);
q_prime[i..i + F32X8_LANES].copy_from_slice(&(qc * scale_v).to_array());
q_sum_acc += qc;
i += F32X8_LANES;
}
let mut q_sum = q_sum_acc.reduce_add();
while i < dim {
q_prime[i] = query[i] * SQ16_FIXED_SCALE;
q_sum += query[i];
i += 1;
}
let q_dot_offset = SQ16_FIXED_OFFSET * q_sum;
let q_norm_sq = match metric {
Metric::L2Sq => dot(query, query),
Metric::Cosine | Metric::NegDot => 0.0,
};
Self {
metric,
dim,
q_prime,
q_dot_offset,
q_norm_sq,
}
}
#[inline]
pub fn distance_with_norm(&self, code_bytes: &[u8], norm: Option<f32>) -> f32 {
debug_assert_eq!(code_bytes.len(), self.dim * 2);
let mut acc = f32x8::ZERO;
let mut i = 0;
while i + F32X8_LANES <= self.dim {
let qp: [f32; F32X8_LANES] = self.q_prime[i..i + F32X8_LANES]
.try_into()
.expect("q_prime[i..i+8] len 8");
let mut code = [0f32; F32X8_LANES];
for (j, lane) in code.iter_mut().enumerate() {
let b = 2 * (i + j);
*lane = u16::from_le_bytes([code_bytes[b], code_bytes[b + 1]]) as f32;
}
acc += f32x8::from(qp) * f32x8::from(code);
i += F32X8_LANES;
}
let mut cross = acc.reduce_add();
while i < self.dim {
let b = 2 * i;
let code = u16::from_le_bytes([code_bytes[b], code_bytes[b + 1]]) as f32;
cross += self.q_prime[i] * code;
i += 1;
}
let dot = cross + self.q_dot_offset;
match self.metric {
Metric::Cosine => {
let x_norm = norm
.expect("Sq16Kernel + Cosine requires per_doc_norms")
.sqrt();
if x_norm > 0.0 {
COSINE_DISTANCE_BASE - dot / x_norm
} else {
COSINE_DISTANCE_BASE - dot
}
}
Metric::NegDot => -dot,
Metric::L2Sq => {
let x_norm_sq = norm.expect("Sq16Kernel + L2Sq requires per_doc_norms");
self.q_norm_sq - L2_CROSS_TERM_COEFF * dot + x_norm_sq
}
}
}
}
#[inline]
pub(crate) fn encode_sq16_row(src: &[f32], out: &mut [u8]) {
debug_assert_eq!(out.len(), src.len() * 2);
let inv_scale = 1.0 / SQ16_FIXED_SCALE;
for (d, &v) in src.iter().enumerate() {
let code = (((v - SQ16_FIXED_OFFSET) * inv_scale).round()).clamp(0.0, 65535.0) as u16;
let b = d * 2;
out[b..b + 2].copy_from_slice(&code.to_le_bytes());
}
}
#[inline]
pub(crate) fn dequantize_sq16_into(code: &[u8], out: &mut [f32]) {
let dim = out.len();
debug_assert_eq!(code.len(), dim * 2);
for (d, slot) in out.iter_mut().enumerate() {
let b = d * 2;
let c = u16::from_le_bytes([code[b], code[b + 1]]) as f32;
*slot = c * SQ16_FIXED_SCALE + SQ16_FIXED_OFFSET;
}
}
#[inline]
pub(crate) fn sq16_decoded_norm_sq(code_bytes: &[u8], dim: usize) -> f32 {
debug_assert_eq!(code_bytes.len(), dim * 2);
let mut acc = f32x8::ZERO;
let off_v = f32x8::splat(SQ16_FIXED_OFFSET);
let scale_v = f32x8::splat(SQ16_FIXED_SCALE);
let mut i = 0;
while i + F32X8_LANES <= dim {
let mut code = [0f32; F32X8_LANES];
for (j, lane) in code.iter_mut().enumerate() {
let b = 2 * (i + j);
*lane = u16::from_le_bytes([code_bytes[b], code_bytes[b + 1]]) as f32;
}
let x = f32x8::from(code) * scale_v + off_v;
acc += x * x;
i += F32X8_LANES;
}
let mut s = acc.reduce_add();
while i < dim {
let b = 2 * i;
let code = u16::from_le_bytes([code_bytes[b], code_bytes[b + 1]]) as f32;
let x = code * SQ16_FIXED_SCALE + SQ16_FIXED_OFFSET;
s += x * x;
i += 1;
}
s
}
#[cfg(test)]
#[inline]
pub(crate) fn sq8_dot(q_prime: &[f32], code_bytes: &[u8], dim: usize) -> f32 {
#[cfg(target_arch = "x86_64")]
{
if avx512_enabled() {
return unsafe { sq8_dot_avx512(q_prime, code_bytes, dim) };
}
if avx2_enabled() {
return unsafe { sq8_dot_avx2(q_prime, code_bytes, dim) };
}
}
sq8_dot_wide(q_prime, code_bytes, dim)
}
#[cfg(test)]
#[inline]
fn sq8_dot_wide(q_prime: &[f32], code_bytes: &[u8], dim: usize) -> f32 {
let mut acc = f32x8::ZERO;
let mut i = 0;
while i + F32X8_LANES <= dim {
let qc: [f32; F32X8_LANES] = q_prime[i..i + F32X8_LANES]
.try_into()
.expect("q_prime[i..i+8] len 8");
let mut bc = [0f32; F32X8_LANES];
for (j, slot) in bc.iter_mut().enumerate() {
*slot = code_bytes[i + j] as f32;
}
let qv = f32x8::from(qc);
let bv = f32x8::from(bc);
acc += qv * bv;
i += F32X8_LANES;
}
let mut dot = acc.reduce_add();
while i < dim {
dot += q_prime[i] * (code_bytes[i] as f32);
i += 1;
}
dot
}
#[cfg(all(test, target_arch = "x86_64"))]
#[target_feature(enable = "avx2")]
unsafe fn sq8_dot_avx2(q_prime: &[f32], code_bytes: &[u8], dim: usize) -> f32 {
debug_assert_eq!(q_prime.len(), dim);
debug_assert_eq!(code_bytes.len(), dim);
unsafe {
let mut acc = _mm256_setzero_ps();
let mut i = 0;
while i + F32X8_LANES <= dim {
let codes_u8 = _mm_loadl_epi64(code_bytes.as_ptr().add(i) as *const __m128i);
let codes_i32 = _mm256_cvtepu8_epi32(codes_u8);
let codes_f32 = _mm256_cvtepi32_ps(codes_i32);
let q = _mm256_loadu_ps(q_prime.as_ptr().add(i));
acc = _mm256_fmadd_ps(q, codes_f32, acc);
i += F32X8_LANES;
}
let lo = _mm256_castps256_ps128(acc);
let hi = _mm256_extractf128_ps(acc, 1);
let sum128 = _mm_add_ps(lo, hi);
let shuf = _mm_movehdup_ps(sum128);
let sums = _mm_add_ps(sum128, shuf);
let shuf2 = _mm_movehl_ps(sums, sums);
let sums2 = _mm_add_ss(sums, shuf2);
let mut dot = _mm_cvtss_f32(sums2);
while i < dim {
dot += q_prime[i] * (code_bytes[i] as f32);
i += 1;
}
dot
}
}
#[cfg(all(test, target_arch = "x86_64"))]
#[target_feature(enable = "avx512f")]
unsafe fn sq8_dot_avx512(q_prime: &[f32], code_bytes: &[u8], dim: usize) -> f32 {
debug_assert_eq!(q_prime.len(), dim);
debug_assert_eq!(code_bytes.len(), dim);
unsafe {
let mut acc = _mm512_setzero_ps();
let mut i = 0;
while i + AVX512_F32_LANES <= dim {
let codes = _mm_loadu_si128(code_bytes.as_ptr().add(i) as *const __m128i);
let codes_i32 = _mm512_cvtepu8_epi32(codes);
let codes_f32 = _mm512_cvtepi32_ps(codes_i32);
let q = _mm512_loadu_ps(q_prime.as_ptr().add(i));
acc = _mm512_fmadd_ps(q, codes_f32, acc);
i += AVX512_F32_LANES;
}
let mut dot = _mm512_reduce_add_ps(acc);
while i < dim {
dot += q_prime[i] * (code_bytes[i] as f32);
i += 1;
}
dot
}
}
#[inline]
pub(crate) fn dequantize_sq8_residual_into(
scale: &[f32],
offset: &[f32],
codes: &[u8],
residuals: &[u8],
residual_divisor: f32,
out: &mut [f32],
) {
let dim = out.len();
debug_assert_eq!(scale.len(), dim);
debug_assert_eq!(offset.len(), dim);
debug_assert_eq!(codes.len(), dim);
debug_assert_eq!(residuals.len(), dim);
#[cfg(target_arch = "x86_64")]
{
if avx512_enabled() {
unsafe {
dequantize_sq8_residual_avx512(
scale,
offset,
codes,
residuals,
residual_divisor,
out,
dim,
);
}
return;
}
if avx2_enabled() {
unsafe {
dequantize_sq8_residual_avx2(
scale,
offset,
codes,
residuals,
residual_divisor,
out,
dim,
);
}
return;
}
}
dequantize_sq8_residual_wide(scale, offset, codes, residuals, residual_divisor, out, dim);
}
#[inline]
fn sq8_residual_component_scalar(
scale: f32,
offset: f32,
code: u8,
residual_byte: u8,
residual_divisor: f32,
) -> f32 {
let inv_div = 1.0 / residual_divisor;
offset + scale * (code as f32 + (i8::from_le_bytes([residual_byte]) as f32) * inv_div)
}
#[inline]
fn dequantize_sq8_residual_wide(
scale: &[f32],
offset: &[f32],
codes: &[u8],
residuals: &[u8],
residual_divisor: f32,
out: &mut [f32],
dim: usize,
) {
let inv_div = 1.0 / residual_divisor;
let inv_v = f32x8::splat(inv_div);
let mut i = 0;
while i + F32X8_LANES <= dim {
let off = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&offset[i..i + F32X8_LANES])
.expect("offset[i..i+8] len 8"),
);
let sc = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&scale[i..i + F32X8_LANES])
.expect("scale[i..i+8] len 8"),
);
let mut code_bc = [0f32; F32X8_LANES];
let mut res_bc = [0f32; F32X8_LANES];
for j in 0..F32X8_LANES {
code_bc[j] = codes[i + j] as f32;
res_bc[j] = i8::from_le_bytes([residuals[i + j]]) as f32;
}
let codes_v = f32x8::from(code_bc);
let res_v = f32x8::from(res_bc);
let term = codes_v + res_v * inv_v;
let decoded = off + sc * term;
out[i..i + F32X8_LANES].copy_from_slice(&decoded.to_array());
i += F32X8_LANES;
}
while i < dim {
out[i] = sq8_residual_component_scalar(
scale[i],
offset[i],
codes[i],
residuals[i],
residual_divisor,
);
i += 1;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn dequantize_sq8_residual_avx2(
scale: &[f32],
offset: &[f32],
codes: &[u8],
residuals: &[u8],
residual_divisor: f32,
out: &mut [f32],
dim: usize,
) {
unsafe {
let inv_div = _mm256_set1_ps(1.0 / residual_divisor);
let mut i = 0;
while i + F32X8_LANES <= dim {
let codes_u8 = _mm_loadl_epi64(codes.as_ptr().add(i) as *const __m128i);
let codes_f32 = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(codes_u8));
let res_u8 = _mm_loadl_epi64(residuals.as_ptr().add(i) as *const __m128i);
let res_f32 = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(res_u8));
let term = _mm256_fmadd_ps(res_f32, inv_div, codes_f32);
let off = _mm256_loadu_ps(offset.as_ptr().add(i));
let sc = _mm256_loadu_ps(scale.as_ptr().add(i));
let decoded = _mm256_fmadd_ps(sc, term, off);
_mm256_storeu_ps(out.as_mut_ptr().add(i), decoded);
i += F32X8_LANES;
}
while i < dim {
out[i] = sq8_residual_component_scalar(
scale[i],
offset[i],
codes[i],
residuals[i],
residual_divisor,
);
i += 1;
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn dequantize_sq8_residual_avx512(
scale: &[f32],
offset: &[f32],
codes: &[u8],
residuals: &[u8],
residual_divisor: f32,
out: &mut [f32],
dim: usize,
) {
unsafe {
let inv_div = _mm512_set1_ps(1.0 / residual_divisor);
let mut i = 0;
while i + AVX512_F32_LANES <= dim {
let codes_u8 = _mm_loadu_si128(codes.as_ptr().add(i) as *const __m128i);
let codes_f32 = _mm512_cvtepi32_ps(_mm512_cvtepu8_epi32(codes_u8));
let res_u8 = _mm_loadu_si128(residuals.as_ptr().add(i) as *const __m128i);
let res_f32 = _mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(res_u8));
let term = _mm512_fmadd_ps(res_f32, inv_div, codes_f32);
let off = _mm512_loadu_ps(offset.as_ptr().add(i));
let sc = _mm512_loadu_ps(scale.as_ptr().add(i));
let decoded = _mm512_fmadd_ps(sc, term, off);
_mm512_storeu_ps(out.as_mut_ptr().add(i), decoded);
i += AVX512_F32_LANES;
}
while i < dim {
out[i] = sq8_residual_component_scalar(
scale[i],
offset[i],
codes[i],
residuals[i],
residual_divisor,
);
i += 1;
}
}
}
#[inline]
pub(crate) fn sq8_residual_norm_sq(
scale: &[f32],
offset: &[f32],
codes: &[u8],
residuals: &[u8],
residual_divisor: f32,
) -> f32 {
let dim = scale.len();
debug_assert_eq!(offset.len(), dim);
debug_assert_eq!(codes.len(), dim);
debug_assert_eq!(residuals.len(), dim);
#[cfg(target_arch = "x86_64")]
{
if avx512_enabled() {
return unsafe {
sq8_residual_norm_sq_avx512(scale, offset, codes, residuals, residual_divisor, dim)
};
}
if avx2_enabled() {
return unsafe {
sq8_residual_norm_sq_avx2(scale, offset, codes, residuals, residual_divisor, dim)
};
}
}
sq8_residual_norm_sq_wide(scale, offset, codes, residuals, residual_divisor, dim)
}
#[inline]
fn sq8_residual_norm_sq_wide(
scale: &[f32],
offset: &[f32],
codes: &[u8],
residuals: &[u8],
residual_divisor: f32,
dim: usize,
) -> f32 {
let inv_div = 1.0 / residual_divisor;
let inv_v = f32x8::splat(inv_div);
let mut acc = f32x8::ZERO;
let mut i = 0;
while i + F32X8_LANES <= dim {
let off = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&offset[i..i + F32X8_LANES])
.expect("offset[i..i+8] len 8"),
);
let sc = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&scale[i..i + F32X8_LANES])
.expect("scale[i..i+8] len 8"),
);
let mut code_bc = [0f32; F32X8_LANES];
let mut res_bc = [0f32; F32X8_LANES];
for j in 0..F32X8_LANES {
code_bc[j] = codes[i + j] as f32;
res_bc[j] = i8::from_le_bytes([residuals[i + j]]) as f32;
}
let codes_v = f32x8::from(code_bc);
let res_v = f32x8::from(res_bc);
let term = codes_v + res_v * inv_v;
let decoded = off + sc * term;
acc += decoded * decoded;
i += F32X8_LANES;
}
let mut sum = acc.reduce_add();
while i < dim {
let v = sq8_residual_component_scalar(
scale[i],
offset[i],
codes[i],
residuals[i],
residual_divisor,
);
sum += v * v;
i += 1;
}
sum
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn sq8_residual_norm_sq_avx2(
scale: &[f32],
offset: &[f32],
codes: &[u8],
residuals: &[u8],
residual_divisor: f32,
dim: usize,
) -> f32 {
unsafe {
let inv_div = _mm256_set1_ps(1.0 / residual_divisor);
let mut acc = _mm256_setzero_ps();
let mut i = 0;
while i + F32X8_LANES <= dim {
let codes_u8 = _mm_loadl_epi64(codes.as_ptr().add(i) as *const __m128i);
let codes_f32 = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(codes_u8));
let res_u8 = _mm_loadl_epi64(residuals.as_ptr().add(i) as *const __m128i);
let res_f32 = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(res_u8));
let term = _mm256_fmadd_ps(res_f32, inv_div, codes_f32);
let off = _mm256_loadu_ps(offset.as_ptr().add(i));
let sc = _mm256_loadu_ps(scale.as_ptr().add(i));
let decoded = _mm256_fmadd_ps(sc, term, off);
acc = _mm256_fmadd_ps(decoded, decoded, acc);
i += F32X8_LANES;
}
let mut sum = horizontal_sum_avx256(acc);
while i < dim {
let v = sq8_residual_component_scalar(
scale[i],
offset[i],
codes[i],
residuals[i],
residual_divisor,
);
sum += v * v;
i += 1;
}
sum
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn sq8_residual_norm_sq_avx512(
scale: &[f32],
offset: &[f32],
codes: &[u8],
residuals: &[u8],
residual_divisor: f32,
dim: usize,
) -> f32 {
unsafe {
let inv_div = _mm512_set1_ps(1.0 / residual_divisor);
let mut acc = _mm512_setzero_ps();
let mut i = 0;
while i + AVX512_F32_LANES <= dim {
let codes_u8 = _mm_loadu_si128(codes.as_ptr().add(i) as *const __m128i);
let codes_f32 = _mm512_cvtepi32_ps(_mm512_cvtepu8_epi32(codes_u8));
let res_u8 = _mm_loadu_si128(residuals.as_ptr().add(i) as *const __m128i);
let res_f32 = _mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(res_u8));
let term = _mm512_fmadd_ps(res_f32, inv_div, codes_f32);
let off = _mm512_loadu_ps(offset.as_ptr().add(i));
let sc = _mm512_loadu_ps(scale.as_ptr().add(i));
let decoded = _mm512_fmadd_ps(sc, term, off);
acc = _mm512_fmadd_ps(decoded, decoded, acc);
i += AVX512_F32_LANES;
}
let mut sum = _mm512_reduce_add_ps(acc);
while i < dim {
let v = sq8_residual_component_scalar(
scale[i],
offset[i],
codes[i],
residuals[i],
residual_divisor,
);
sum += v * v;
i += 1;
}
sum
}
}
#[inline]
pub(crate) fn decode_f32_le_into(bytes: &[u8], out: &mut [f32]) {
debug_assert_eq!(bytes.len(), out.len() * F32_BYTES);
if let Ok(decoded) = bytemuck::try_cast_slice::<u8, f32>(bytes) {
out.copy_from_slice(decoded);
return;
}
let n = out.len();
#[cfg(target_arch = "x86_64")]
{
if avx512_enabled() {
unsafe {
decode_f32_le_avx512(bytes, out, n);
}
return;
}
if avx2_enabled() {
unsafe {
decode_f32_le_avx2(bytes, out, n);
}
return;
}
}
decode_f32_le_wide(bytes, out, n);
}
#[inline]
fn decode_f32_le_wide(bytes: &[u8], out: &mut [f32], n: usize) {
let mut i = 0;
while i + F32X8_LANES <= n {
let mut lane = [0f32; F32X8_LANES];
for (j, slot) in lane.iter_mut().enumerate() {
let b = (i + j) * F32_BYTES;
*slot = f32::from_le_bytes([bytes[b], bytes[b + 1], bytes[b + 2], bytes[b + 3]]);
}
out[i..i + F32X8_LANES].copy_from_slice(&lane);
i += F32X8_LANES;
}
while i < n {
let b = i * F32_BYTES;
out[i] = f32::from_le_bytes([bytes[b], bytes[b + 1], bytes[b + 2], bytes[b + 3]]);
i += 1;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn decode_f32_le_avx2(bytes: &[u8], out: &mut [f32], n: usize) {
unsafe {
let src = bytes.as_ptr();
let dst = out.as_mut_ptr();
let mut i = 0;
while i + F32X8_LANES <= n {
let v = _mm256_loadu_ps(src.add(i * F32_BYTES) as *const f32);
_mm256_storeu_ps(dst.add(i), v);
i += F32X8_LANES;
}
while i < n {
let b = i * F32_BYTES;
*dst.add(i) = f32::from_le_bytes([
*src.add(b),
*src.add(b + 1),
*src.add(b + 2),
*src.add(b + 3),
]);
i += 1;
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn decode_f32_le_avx512(bytes: &[u8], out: &mut [f32], n: usize) {
unsafe {
let src = bytes.as_ptr();
let dst = out.as_mut_ptr();
let mut i = 0;
while i + AVX512_F32_LANES <= n {
let v = _mm512_loadu_ps(src.add(i * F32_BYTES) as *const f32);
_mm512_storeu_ps(dst.add(i), v);
i += AVX512_F32_LANES;
}
while i < n {
let b = i * F32_BYTES;
*dst.add(i) = f32::from_le_bytes([
*src.add(b),
*src.add(b + 1),
*src.add(b + 2),
*src.add(b + 3),
]);
i += 1;
}
}
}
#[cfg(target_arch = "x86_64")]
#[inline]
unsafe fn horizontal_sum_avx256(v: __m256) -> f32 {
unsafe {
let hi = _mm256_extractf128_ps(v, 1);
let lo = _mm256_castps256_ps128(v);
let sum128 = _mm_add_ps(lo, hi);
let shuf = _mm_movehdup_ps(sum128);
let sums = _mm_add_ps(sum128, shuf);
let shuf2 = _mm_movehl_ps(shuf, sums);
let sums2 = _mm_add_ss(sums, shuf2);
_mm_cvtss_f32(sums2)
}
}
#[inline]
pub(crate) fn decode_f32_le_vec(bytes: &[u8]) -> Vec<f32> {
debug_assert_eq!(bytes.len() % F32_BYTES, 0);
let mut out = vec![0f32; bytes.len() / F32_BYTES];
decode_f32_le_into(bytes, &mut out);
out
}
#[inline]
pub(crate) fn add_f32_to_f64_acc(acc: &mut [f64], row: &[f32]) {
debug_assert_eq!(acc.len(), row.len());
#[cfg(target_arch = "x86_64")]
if avx2_enabled() {
unsafe {
add_f32_to_f64_acc_avx2(acc, row);
}
return;
}
add_f32_to_f64_acc_scalar(acc, row);
}
#[inline]
pub(crate) fn add_weighted_f32_to_f64_acc(acc: &mut [f64], row: &[f32], weight: f64) {
debug_assert_eq!(acc.len(), row.len());
#[cfg(target_arch = "x86_64")]
if avx2_enabled() {
unsafe {
add_weighted_f32_to_f64_acc_avx2(acc, row, weight);
}
return;
}
for (a, &x) in acc.iter_mut().zip(row.iter()) {
*a += x as f64 * weight;
}
}
#[inline]
pub(crate) fn f64_acc_mean_into_f32(acc: &[f64], inv: f64, out: &mut [f32]) {
debug_assert_eq!(acc.len(), out.len());
for (o, &a) in out.iter_mut().zip(acc.iter()) {
*o = (a * inv) as f32;
}
}
#[inline]
pub(crate) fn mean_f32_cluster_major(vectors: &[f32], dim: usize, n: usize) -> Vec<f32> {
debug_assert_eq!(vectors.len(), n * dim);
let mut acc = vec![0f64; dim];
for i in 0..n {
add_f32_to_f64_acc(&mut acc, &vectors[i * dim..(i + 1) * dim]);
}
let mut out = vec![0f32; dim];
if n > 0 {
f64_acc_mean_into_f32(&acc, 1.0 / n as f64, &mut out);
}
out
}
#[inline]
fn add_f32_to_f64_acc_scalar(acc: &mut [f64], row: &[f32]) {
for (a, &x) in acc.iter_mut().zip(row.iter()) {
*a += x as f64;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn add_f32_to_f64_acc_avx2(acc: &mut [f64], row: &[f32]) {
unsafe {
let mut i = 0;
let n = acc.len();
while i + F32X8_LANES <= n {
let vf = _mm256_loadu_ps(row.as_ptr().add(i));
let lo = _mm256_cvtps_pd(_mm256_castps256_ps128(vf));
let hi = _mm256_cvtps_pd(_mm256_extractf128_ps(vf, 1));
let alo = _mm256_loadu_pd(acc.as_mut_ptr().add(i));
let ahi = _mm256_loadu_pd(acc.as_mut_ptr().add(i + 4));
_mm256_storeu_pd(acc.as_mut_ptr().add(i), _mm256_add_pd(alo, lo));
_mm256_storeu_pd(acc.as_mut_ptr().add(i + 4), _mm256_add_pd(ahi, hi));
i += F32X8_LANES;
}
while i < n {
acc[i] += row[i] as f64;
i += 1;
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn add_weighted_f32_to_f64_acc_avx2(acc: &mut [f64], row: &[f32], weight: f64) {
unsafe {
let w = _mm256_set1_pd(weight);
let mut i = 0;
let n = acc.len();
while i + F32X8_LANES <= n {
let vf = _mm256_loadu_ps(row.as_ptr().add(i));
let lo = _mm256_cvtps_pd(_mm256_castps256_ps128(vf));
let hi = _mm256_cvtps_pd(_mm256_extractf128_ps(vf, 1));
let wlo = _mm256_mul_pd(lo, w);
let whi = _mm256_mul_pd(hi, w);
let alo = _mm256_loadu_pd(acc.as_mut_ptr().add(i));
let ahi = _mm256_loadu_pd(acc.as_mut_ptr().add(i + 4));
_mm256_storeu_pd(acc.as_mut_ptr().add(i), _mm256_add_pd(alo, wlo));
_mm256_storeu_pd(acc.as_mut_ptr().add(i + 4), _mm256_add_pd(ahi, whi));
i += F32X8_LANES;
}
while i < n {
acc[i] += row[i] as f64 * weight;
i += 1;
}
}
}
pub fn normalize(v: &mut [f32]) {
let mag = {
let mut acc = f32x8::ZERO;
let mut tail_acc: f32 = 0.0;
let chunks = v.chunks_exact(F32X8_LANES);
let tail = chunks.remainder();
for c in chunks {
let lane = f32x8::from(
<[f32; F32X8_LANES]>::try_from(c)
.expect("chunks_exact(8) yields slices of length 8"),
);
acc += lane * lane;
}
for &x in tail {
tail_acc += x * x;
}
(acc.reduce_add() + tail_acc).sqrt()
};
if mag.is_normal() {
let inv = 1.0 / mag;
let inv_v = f32x8::splat(inv);
let mut chunks = v.chunks_exact_mut(F32X8_LANES);
for c in chunks.by_ref() {
let lane = f32x8::from(
<[f32; F32X8_LANES]>::try_from(&*c)
.expect("chunks_exact_mut(8) yields slices of length 8"),
);
let scaled = lane * inv_v;
c.copy_from_slice(&scaled.to_array());
}
for x in chunks.into_remainder() {
*x *= inv;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn approx(a: f32, b: f32, eps: f32) -> bool {
(a - b).abs() < eps
}
fn scalar_nearest_two_centroids(
metric: Metric,
query: &[f32],
centroids: &[f32],
n_cent: usize,
dim: usize,
counts: Option<&[u32]>,
) -> Option<((u32, f32), Option<(u32, f32)>)> {
let mut best: Option<(u32, f32)> = None;
let mut second: Option<(u32, f32)> = None;
for c in 0..n_cent {
if counts.is_some_and(|counts| counts[c] == 0) {
continue;
}
let score = distance(metric, query, ¢roids[c * dim..(c + 1) * dim]);
match best {
None => best = Some((c as u32, score)),
Some((_, best_score)) if score < best_score => {
second = best;
best = Some((c as u32, score));
}
_ => {
if second.is_none_or(|(_, second_score)| score < second_score) {
second = Some((c as u32, score));
}
}
}
}
best.map(|best| (best, second))
}
fn assert_nearest_two_matches(
got: Option<((u32, f32), Option<(u32, f32)>)>,
expected: Option<((u32, f32), Option<(u32, f32)>)>,
) {
let (got_best, got_second) = got.expect("nearest result");
let (expected_best, expected_second) = expected.expect("scalar nearest result");
assert_eq!(got_best.0, expected_best.0);
assert!(
approx(got_best.1, expected_best.1, 1e-3),
"best score got {} expected {}",
got_best.1,
expected_best.1
);
let got_second = got_second.expect("second result");
let expected_second = expected_second.expect("scalar second result");
assert_eq!(got_second.0, expected_second.0);
assert!(
approx(got_second.1, expected_second.1, 1e-3),
"second score got {} expected {}",
got_second.1,
expected_second.1
);
}
#[test]
fn nearest_centroid_transposed_matches_naive_scan() {
for n_cent in [128usize, 130, 144, 160] {
let dim = 33;
let mut centroids = Vec::with_capacity(n_cent * dim);
for c in 0..n_cent {
for d in 0..dim {
centroids.push(((c * 31 + d * 17) % 29) as f32 * 0.04 - 0.5);
}
}
let dup = centroids[3 * dim..4 * dim].to_vec();
let last = (n_cent - 1) * dim;
centroids[last..last + dim].copy_from_slice(&dup);
let transposed = transpose_centroids_cluster_major(¢roids, n_cent, dim);
for probe in 0..64 {
let query: Vec<f32> = (0..dim)
.map(|d| ((probe * 13 + d * 7) % 23) as f32 * 0.05 - 0.4)
.collect();
let mut naive = (0u32, f32::INFINITY);
for c in 0..n_cent {
let dist = l2_sq(&query, ¢roids[c * dim..(c + 1) * dim]);
if dist < naive.1 {
naive = (c as u32, dist);
}
}
let blocked =
nearest_centroid_transposed(Metric::L2Sq, &query, &transposed, n_cent, dim);
assert_eq!(
blocked.0, naive.0,
"n_cent {n_cent} probe {probe}: blocked argmin diverged from naive"
);
}
let tie_query = dup.clone();
let blocked =
nearest_centroid_transposed(Metric::L2Sq, &tie_query, &transposed, n_cent, dim);
assert_eq!(blocked.0, 3, "tie must resolve to the lowest index");
}
}
#[test]
fn nearest_two_centroids_transposed_matches_scalar_reference() {
let dim = 17;
let n_cent = 19;
let query: Vec<f32> = (0..dim)
.map(|d| ((d * 37 % 23) as f32 - 11.0) * 0.031)
.collect();
let mut centroids = Vec::with_capacity(n_cent * dim);
for c in 0..n_cent {
for d in 0..dim {
centroids.push(((c * 13 + d * 7) % 31) as f32 * 0.02 - 0.3 + c as f32 * 0.001);
}
}
let transposed = transpose_centroids_cluster_major(¢roids, n_cent, dim);
for metric in [Metric::L2Sq, Metric::Cosine, Metric::NegDot] {
let got =
nearest_two_centroids_transposed(metric, &query, &transposed, n_cent, dim, None);
let expected =
scalar_nearest_two_centroids(metric, &query, ¢roids, n_cent, dim, None);
assert_nearest_two_matches(got, expected);
}
}
#[test]
fn nearest_two_centroids_transposed_honors_zero_counts() {
let dim = 17;
let n_cent = 19;
let query: Vec<f32> = (0..dim).map(|d| d as f32 * 0.01).collect();
let mut centroids = Vec::with_capacity(n_cent * dim);
for c in 0..n_cent {
for d in 0..dim {
centroids.push(c as f32 + d as f32 * 0.001);
}
}
let mut counts = vec![1u32; n_cent];
counts[0] = 0;
let transposed = transpose_centroids_cluster_major(¢roids, n_cent, dim);
let got = nearest_two_centroids_transposed(
Metric::L2Sq,
&query,
&transposed,
n_cent,
dim,
Some(&counts),
);
let expected = scalar_nearest_two_centroids(
Metric::L2Sq,
&query,
¢roids,
n_cent,
dim,
Some(&counts),
);
assert_nearest_two_matches(got, expected);
assert_ne!(got.expect("nearest result").0.0, 0);
}
#[test]
fn dot_zero_vectors() {
let a = vec![0.0; 16];
let b = vec![0.0; 16];
assert_eq!(dot(&a, &b), 0.0);
}
#[test]
fn dot_orthogonal_basis_vectors() {
let mut a = vec![0.0; 16];
let mut b = vec![0.0; 16];
a[0] = 1.0;
b[1] = 1.0;
assert_eq!(dot(&a, &b), 0.0);
}
#[test]
fn dot_self_is_squared_norm() {
let v: Vec<f32> = (1..=16).map(|i| i as f32).collect();
let want: f32 = (1..=16).map(|i| (i * i) as f32).sum();
assert!(approx(dot(&v, &v), want, 1e-3));
}
#[test]
fn sum_f32_matches_scalar_reference() {
for len in [1, 7, 8, 15, 16, 17, 384] {
let v: Vec<f32> = (0..len).map(|i| (i as f32) * 0.25 - 1.0).collect();
let expected: f32 = v.iter().sum();
assert!(
approx(sum_f32(&v), expected, 1e-4),
"len={len}: got {} expected {expected}",
sum_f32(&v)
);
}
}
#[test]
fn sq8_residual_norm_sq_matches_dequant_dot_self() {
let dim = 17;
let scale: Vec<f32> = (0..dim).map(|i| 0.01 * (i as f32 + 1.0)).collect();
let offset: Vec<f32> = (0..dim).map(|i| -0.5 + 0.03 * i as f32).collect();
let codes: Vec<u8> = (0..dim).map(|i| (i * 17 % 256) as u8).collect();
let residuals: Vec<u8> = (0..dim)
.map(|i| ((i as i8).wrapping_mul(3)).to_le_bytes()[0])
.collect();
let mut decoded = vec![0f32; dim];
dequantize_sq8_residual_into(
&scale,
&offset,
&codes,
&residuals,
SQ8_RESIDUAL_DIVISOR,
&mut decoded,
);
let expected = dot(&decoded, &decoded);
let got = sq8_residual_norm_sq(&scale, &offset, &codes, &residuals, SQ8_RESIDUAL_DIVISOR);
assert!(
(got - expected).abs() <= 1e-5 * (1.0 + expected.abs()),
"norm {got} vs dequant-dot-self {expected}"
);
}
#[test]
fn decode_f32_le_into_round_trip() {
let values: Vec<f32> = (0..19).map(|i| i as f32 * 0.125 - 2.0).collect();
let bytes: Vec<u8> = values.iter().flat_map(|v| v.to_le_bytes()).collect();
let mut out = vec![0f32; values.len()];
decode_f32_le_into(&bytes, &mut out);
assert_eq!(out, values);
}
#[test]
fn dot_handles_tail_not_multiple_of_8() {
let a: Vec<f32> = vec![1.0; 11];
let b: Vec<f32> = vec![2.0; 11];
assert!(approx(dot(&a, &b), 22.0, 1e-5));
}
#[test]
fn dot_short_input() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![4.0, 5.0, 6.0];
assert!(approx(dot(&a, &b), 32.0, 1e-5));
}
#[test]
fn l2_sq_identical_inputs_zero() {
let v = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
assert_eq!(l2_sq(&v, &v), 0.0);
}
#[test]
fn l2_sq_unit_offset_per_dim() {
let a = vec![0.0; 16];
let b = vec![1.0; 16];
assert!(approx(l2_sq(&a, &b), 16.0, 1e-5));
}
#[test]
fn l2_sq_handles_tail() {
let a = vec![0.0; 11];
let b = vec![3.0; 11];
assert!(approx(l2_sq(&a, &b), 99.0, 1e-5));
}
#[test]
fn normalize_unit_vector_stays_unit() {
let mut v = vec![1.0, 0.0, 0.0, 0.0];
normalize(&mut v);
assert_eq!(v, vec![1.0, 0.0, 0.0, 0.0]);
}
#[test]
fn normalize_scales_magnitude_to_one() {
let mut v = vec![3.0, 4.0]; normalize(&mut v);
assert!(approx(v[0], 0.6, 1e-5));
assert!(approx(v[1], 0.8, 1e-5));
}
#[test]
fn normalize_zero_vector_left_alone() {
let mut v = vec![0.0; 16];
normalize(&mut v);
for &x in &v {
assert_eq!(x, 0.0);
}
}
#[test]
fn normalize_then_self_dot_is_one() {
let mut v: Vec<f32> = (1..=16).map(|i| i as f32).collect();
normalize(&mut v);
assert!(approx(dot(&v, &v), 1.0, 1e-5));
}
#[test]
fn normalize_degenerate_magnitude_never_produces_inf() {
let mut v = vec![f32::from_bits(1); 16]; let before = v.clone();
normalize(&mut v);
assert_eq!(v, before);
for &x in &v {
assert!(x.is_finite());
}
}
#[test]
fn distance_cosine_uses_one_minus_dot() {
let a = vec![1.0, 0.0, 0.0, 0.0];
let b = vec![1.0, 0.0, 0.0, 0.0];
assert!(approx(distance(Metric::Cosine, &a, &b), 0.0, 1e-5));
let c = vec![0.0, 1.0, 0.0, 0.0];
assert!(approx(distance(Metric::Cosine, &a, &c), 1.0, 1e-5));
}
#[test]
fn distance_l2sq_zero_for_identical() {
let v = vec![1.0, 2.0, 3.0, 4.0];
assert_eq!(distance(Metric::L2Sq, &v, &v), 0.0);
}
#[test]
fn distance_negdot_inverts_dot() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let b = vec![4.0, 3.0, 2.0, 1.0];
assert!(approx(distance(Metric::NegDot, &a, &b), -20.0, 1e-5));
}
#[test]
fn distance_smaller_is_closer_for_every_metric() {
let q = vec![1.0, 0.0, 0.0, 0.0];
let near = vec![1.0, 0.0, 0.0, 0.0];
let far = vec![-1.0, 0.0, 0.0, 0.0];
for m in [Metric::Cosine, Metric::L2Sq, Metric::NegDot] {
let d_near = distance(m, &q, &near);
let d_far = distance(m, &q, &far);
assert!(
d_near < d_far,
"metric {m:?}: near {d_near} should be < far {d_far}"
);
}
}
fn encode_sq8(values: &[f32], dim: usize, scale: &[f32], offset: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(values.len());
for row in values.chunks_exact(dim) {
for d in 0..dim {
let q = ((row[d] - offset[d]) / scale[d]).round().clamp(0.0, 255.0) as u8;
out.push(q);
}
}
out
}
fn decode_sq8(codes: &[u8], dim: usize, scale: &[f32], offset: &[f32]) -> Vec<f32> {
codes
.iter()
.enumerate()
.map(|(i, &c)| (c as f32) * scale[i % dim] + offset[i % dim])
.collect()
}
fn decode_sq8_residual(
codes: &[u8],
residuals: &[u8],
dim: usize,
scale: &[f32],
offset: &[f32],
residual_divisor: f32,
) -> Vec<f32> {
codes
.iter()
.zip(residuals.iter())
.enumerate()
.map(|(i, (&c, &r))| {
let d = i % dim;
(c as f32) * scale[d]
+ offset[d]
+ (i8::from_le_bytes([r]) as f32) * scale[d] / residual_divisor
})
.collect()
}
#[test]
fn sq8_residual_kernel_matches_corrected_reference() {
let dim = 24usize;
let residual_divisor = SQ8_RESIDUAL_DIVISOR;
let query: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.04 - 0.2).collect();
let scale: Vec<f32> = (0..dim).map(|i| 0.01 + (i as f32) * 0.001).collect();
let offset: Vec<f32> = (0..dim).map(|i| -0.4 + (i as f32) * 0.03).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 29 + 7) % 256) as u8).collect();
let residuals: Vec<u8> = (0..dim)
.map(|i| (((i * 17 + 3) % 63) as i8 - 31).to_le_bytes()[0])
.collect();
let corrected =
decode_sq8_residual(&codes, &residuals, dim, &scale, &offset, residual_divisor);
let corrected_norm: f32 = corrected.iter().map(|x| x * x).sum();
let norms = [corrected_norm];
for metric in [Metric::Cosine, Metric::L2Sq, Metric::NegDot] {
let norms_arg = match metric {
Metric::Cosine | Metric::L2Sq => Some(&norms[..]),
Metric::NegDot => None,
};
let kernel = Sq8ResidualKernel::new(metric, &query, &scale, &offset, residual_divisor);
let got =
kernel.distance_with_norm(&codes, &residuals, norms_arg.map(|norms| norms[0]));
let want = match metric {
Metric::Cosine => 1.0 - dot(&query, &corrected) / corrected_norm.sqrt(),
_ => distance(metric, &query, &corrected),
};
assert!(
(want - got).abs() <= 1e-4,
"metric {metric:?}: residual kernel {got} vs corrected ref {want}"
);
}
}
#[test]
fn sq8_residual_kernel_handles_tail_dim_not_multiple_of_8() {
let dim = 13usize;
let residual_divisor = SQ8_RESIDUAL_DIVISOR;
let query: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.03 + 0.1).collect();
let scale: Vec<f32> = (0..dim).map(|i| 0.02 + (i as f32) * 0.001).collect();
let offset: Vec<f32> = (0..dim).map(|i| -0.2 + (i as f32) * 0.02).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 11 + 5) % 256) as u8).collect();
let residuals: Vec<u8> = (0..dim)
.map(|i| (((i * 23 + 9) % 47) as i8 - 23).to_le_bytes()[0])
.collect();
let corrected =
decode_sq8_residual(&codes, &residuals, dim, &scale, &offset, residual_divisor);
let kernel =
Sq8ResidualKernel::new(Metric::NegDot, &query, &scale, &offset, residual_divisor);
let got = kernel.distance_with_norm(&codes, &residuals, None);
let want = distance(Metric::NegDot, &query, &corrected);
assert!(
(want - got).abs() <= 1e-4,
"tail-dim residual kernel: got {got} vs corrected ref {want}"
);
}
#[test]
fn sq8_kernel_dot_matches_decoded_reference() {
let dim = 16usize;
let query: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.05 - 0.3).collect();
let scale: Vec<f32> = (0..dim).map(|i| 0.01 + (i as f32) * 0.002).collect();
let offset: Vec<f32> = (0..dim).map(|i| -1.0 + (i as f32) * 0.1).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 17 + 3) % 256) as u8).collect();
let decoded = decode_sq8(&codes, dim, &scale, &offset);
for m in [Metric::Cosine, Metric::NegDot] {
let norms = if m == Metric::Cosine {
Some(vec![decoded.iter().map(|x| x * x).sum::<f32>()])
} else {
None
};
let want = match m {
Metric::Cosine => {
let x_norm = decoded.iter().map(|x| x * x).sum::<f32>().sqrt();
if x_norm > 0.0 {
1.0 - dot(&query, &decoded) / x_norm
} else {
1.0 - dot(&query, &decoded)
}
}
Metric::NegDot => distance(m, &query, &decoded),
Metric::L2Sq => unreachable!(),
};
let kernel = Sq8Kernel::new(m, &query, &scale, &offset, norms.clone().map(Arc::from));
let got = kernel.distance_at(0, &codes);
let err = (want - got).abs();
assert!(
err <= 1e-4,
"metric {m:?}: kernel {got} vs decoded ref {want} (err {err})"
);
}
}
#[test]
fn sq8_kernel_l2sq_matches_decoded_reference() {
let dim = 24usize;
let query: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.07 - 0.1).collect();
let scale: Vec<f32> = (0..dim).map(|i| 0.02 + (i as f32) * 0.003).collect();
let offset: Vec<f32> = (0..dim).map(|i| 0.5 - (i as f32) * 0.05).collect();
let codes_doc0: Vec<u8> = (0..dim).map(|i| ((i * 7) % 256) as u8).collect();
let codes_doc1: Vec<u8> = (0..dim).map(|i| ((i * 31 + 12) % 256) as u8).collect();
let decoded0 = decode_sq8(&codes_doc0, dim, &scale, &offset);
let decoded1 = decode_sq8(&codes_doc1, dim, &scale, &offset);
let norm0: f32 = decoded0.iter().map(|x| x * x).sum();
let norm1: f32 = decoded1.iter().map(|x| x * x).sum();
let per_doc_norms = vec![norm0, norm1];
let kernel = Sq8Kernel::new(
Metric::L2Sq,
&query,
&scale,
&offset,
Some(Arc::from(per_doc_norms.clone())),
);
let got0 = kernel.distance_at(0, &codes_doc0);
let want0 = distance(Metric::L2Sq, &query, &decoded0);
assert!(
(want0 - got0).abs() <= 1e-3,
"doc0: kernel {got0} vs decoded ref {want0}"
);
let got1 = kernel.distance_at(1, &codes_doc1);
let want1 = distance(Metric::L2Sq, &query, &decoded1);
assert!(
(want1 - got1).abs() <= 1e-3,
"doc1: kernel {got1} vs decoded ref {want1}"
);
}
#[test]
fn sq8_kernel_handles_tail_dim_not_multiple_of_8() {
let dim = 13usize;
let query: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.03 + 0.1).collect();
let scale: Vec<f32> = (0..dim).map(|i| 0.01 + (i as f32) * 0.001).collect();
let offset: Vec<f32> = (0..dim).map(|i| -0.1 + (i as f32) * 0.02).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 11 + 5) % 256) as u8).collect();
let decoded = decode_sq8(&codes, dim, &scale, &offset);
let kernel = Sq8Kernel::new(Metric::NegDot, &query, &scale, &offset, None);
let got = kernel.distance_at(0, &codes);
let want = distance(Metric::NegDot, &query, &decoded);
assert!(
(want - got).abs() <= 1e-4,
"tail-dim Sq8 kernel: got {got} vs decoded ref {want}"
);
}
#[test]
fn sq16_round_trip_within_16bit_tolerance_of_fp32() {
fn normalize(v: &mut [f32]) {
let n = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if n > 0.0 {
for x in v.iter_mut() {
*x /= n;
}
}
}
for &dim in &[8usize, 13, 100, 384] {
let mut query: Vec<f32> = (0..dim).map(|i| ((i as f32) * 0.017 - 0.4).sin()).collect();
let mut vec: Vec<f32> = (0..dim)
.map(|i| ((i as f32) * 0.023 + 0.11).cos())
.collect();
normalize(&mut query);
normalize(&mut vec);
let mut bytes = vec![0u8; dim * 2];
encode_sq16_row(&vec, &mut bytes);
let norm_sq = sq16_decoded_norm_sq(&bytes, dim);
let kernel = Sq16Kernel::new(Metric::Cosine, &query);
let got = kernel.distance_with_norm(&bytes, Some(norm_sq));
let want = distance(Metric::Cosine, &query, &vec);
let rel = (got - want).abs() / want.abs();
assert!(
rel < 1e-4,
"dim {dim}: Sq16 cosine {got} vs fp32 ref {want} (rel err {rel:e})"
);
let decoded: Vec<f32> = (0..dim)
.map(|d| {
let b = d * 2;
let code = u16::from_le_bytes([bytes[b], bytes[b + 1]]);
code as f32 * SQ16_FIXED_SCALE + SQ16_FIXED_OFFSET
})
.collect();
let decoded_norm = decoded.iter().map(|x| x * x).sum::<f32>().sqrt();
let decoded_ref = COSINE_DISTANCE_BASE - dot(&query, &decoded) / decoded_norm;
assert!(
(got - decoded_ref).abs() <= 1e-4,
"dim {dim}: kernel {got} disagrees with decoded ref {decoded_ref}"
);
}
}
#[test]
fn sq16_norm_correction_ranks_planted_case_at_least_as_well() {
let dim = 8usize;
let inv = 1.0 / (dim as f32).sqrt();
let query = vec![inv; dim];
let true_nn = query.clone();
let mut e = vec![0.0f32; dim];
for (i, ei) in e.iter_mut().enumerate() {
*ei = if i % 2 == 0 { inv } else { -inv };
}
let c = 0.9f32;
let s = (1.0 - c * c).sqrt();
let d_unit: Vec<f32> = (0..dim).map(|i| c * query[i] + s * e[i]).collect();
let distractor: Vec<f32> = d_unit.iter().map(|v| v * 1.2).collect();
for &v in true_nn.iter().chain(distractor.iter()) {
assert!(v.abs() <= 1.0, "component {v} outside Sq16 grid");
}
let encode = |v: &[f32]| {
let mut b = vec![0u8; dim * 2];
encode_sq16_row(v, &mut b);
b
};
let nn_b = encode(&true_nn);
let dis_b = encode(&distractor);
let kernel = Sq16Kernel::new(Metric::Cosine, &query);
let nn_corr = kernel.distance_with_norm(&nn_b, Some(sq16_decoded_norm_sq(&nn_b, dim)));
let dis_corr = kernel.distance_with_norm(&dis_b, Some(sq16_decoded_norm_sq(&dis_b, dim)));
assert!(
nn_corr < dis_corr,
"norm-corrected: true NN {nn_corr} must rank before distractor {dis_corr}"
);
let decode = |b: &[u8]| -> Vec<f32> {
(0..dim)
.map(|d| {
let o = d * 2;
u16::from_le_bytes([b[o], b[o + 1]]) as f32 * SQ16_FIXED_SCALE
+ SQ16_FIXED_OFFSET
})
.collect()
};
let nn_raw = COSINE_DISTANCE_BASE - dot(&query, &decode(&nn_b));
let dis_raw = COSINE_DISTANCE_BASE - dot(&query, &decode(&dis_b));
assert!(
dis_raw < nn_raw,
"sanity: uncorrected 1-dot should (wrongly) rank the distractor first \
(nn_raw {nn_raw}, dis_raw {dis_raw})"
);
}
#[test]
fn sq16_vs_sq8residual_kernel_recall() {
use std::collections::HashSet;
use crate::superfile::vector::{
rerank_codec::{SQ8_FIXED_OFFSET, SQ8_FIXED_RESIDUAL_DIVISOR, SQ8_FIXED_SCALE},
sq8_simd::{Sq8EncodeConsts, encode_sq8_residual_row},
};
fn next_u64(state: &mut u64) -> u64 {
*state = state.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = *state;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
fn next_f32(state: &mut u64) -> f32 {
let u = (next_u64(state) >> 40) as f32 / (1u64 << 24) as f32;
u * 2.0 - 1.0
}
fn fill_unit(state: &mut u64, out: &mut [f32]) {
for x in out.iter_mut() {
*x = next_f32(state);
}
let n = out.iter().map(|v| v * v).sum::<f32>().sqrt();
if n > 0.0 {
for x in out.iter_mut() {
*x /= n;
}
}
}
fn topk_asc(dists: &[f32], k: usize) -> Vec<usize> {
let mut idx: Vec<usize> = (0..dists.len()).collect();
idx.sort_by(|&a, &b| dists[a].total_cmp(&dists[b]));
idx.truncate(k);
idx
}
const K: usize = 10;
let dim = 768usize;
let mut rng = 0xDEAD_BEEF_1234_5678u64;
let scale = vec![SQ8_FIXED_SCALE; dim];
let offset = vec![SQ8_FIXED_OFFSET; dim];
let divisor = SQ8_FIXED_RESIDUAL_DIVISOR;
let consts = Sq8EncodeConsts::from_scale_offset(&scale, &offset);
#[allow(non_snake_case)]
let N = 800usize;
#[allow(non_snake_case)]
let Q = 80usize;
let mut corpus: Vec<Vec<f32>> = Vec::with_capacity(N);
let mut sq16_buf = vec![0u8; N * dim * 2];
let mut sq8_code = vec![0u8; N * dim];
let mut sq8_res = vec![0u8; N * dim];
let mut sq16_norms = vec![0.0f32; N];
let mut sq8_norms = vec![0.0f32; N];
let mut recon = vec![0.0f32; dim];
for i in 0..N {
let mut c = vec![0.0f32; dim];
fill_unit(&mut rng, &mut c);
let sc = &mut sq16_buf[i * dim * 2..(i + 1) * dim * 2];
encode_sq16_row(&c, sc);
sq16_norms[i] = sq16_decoded_norm_sq(sc, dim);
let n = encode_sq8_residual_row(
&c,
&consts,
&scale,
&offset,
&mut sq8_code[i * dim..(i + 1) * dim],
&mut sq8_res[i * dim..(i + 1) * dim],
&mut recon,
true,
divisor,
)
.expect("store_norm=true yields a per-doc norm");
sq8_norms[i] = n;
corpus.push(c);
}
let (mut r_sq16, mut r_sq16_nn, mut r_sq8, mut total) = (0usize, 0usize, 0usize, 0usize);
let mut rng_q = 0x0BAD_F00D_CAFE_BABEu64;
for _ in 0..Q {
let mut q = vec![0.0f32; dim];
fill_unit(&mut rng_q, &mut q);
let truth_scores: Vec<f32> = corpus.iter().map(|c| dot(&q, c)).collect();
let mut ti: Vec<usize> = (0..N).collect();
ti.sort_by(|&a, &b| truth_scores[b].total_cmp(&truth_scores[a]));
let truth: HashSet<usize> = ti.into_iter().take(K).collect();
let sq16_kernel = Sq16Kernel::new(Metric::Cosine, &q);
let sq8_kernel = Sq8ResidualKernel::new(Metric::Cosine, &q, &scale, &offset, divisor);
let sq16_d: Vec<f32> = (0..N)
.map(|i| {
sq16_kernel.distance_with_norm(
&sq16_buf[i * dim * 2..(i + 1) * dim * 2],
Some(sq16_norms[i]),
)
})
.collect();
let sq8_d: Vec<f32> = (0..N)
.map(|i| {
sq8_kernel.distance_with_norm(
&sq8_code[i * dim..(i + 1) * dim],
&sq8_res[i * dim..(i + 1) * dim],
Some(sq8_norms[i]),
)
})
.collect();
let sq16_nn_d: Vec<f32> = (0..N)
.map(|i| {
let base = i * dim * 2;
let mut d = 0.0f32;
for (j, &qj) in q.iter().enumerate().take(dim) {
let b = base + j * 2;
let code = u16::from_le_bytes([sq16_buf[b], sq16_buf[b + 1]]) as f32;
d += qj * (code * SQ16_FIXED_SCALE + SQ16_FIXED_OFFSET);
}
COSINE_DISTANCE_BASE - d
})
.collect();
for &i in topk_asc(&sq16_d, K).iter() {
if truth.contains(&i) {
r_sq16 += 1;
}
}
for &i in topk_asc(&sq16_nn_d, K).iter() {
if truth.contains(&i) {
r_sq16_nn += 1;
}
}
for &i in topk_asc(&sq8_d, K).iter() {
if truth.contains(&i) {
r_sq8 += 1;
}
}
total += K;
}
let rec = |x: usize| x as f64 / total as f64;
eprintln!(
"\n### Kernel-level recall@{K} vs fp32 truth (synthetic, N={N}, Q={Q}, dim={dim})"
);
eprintln!("Sq16 (norm) : {:.4}", rec(r_sq16));
eprintln!("Sq16 (no-norm) : {:.4}", rec(r_sq16_nn));
eprintln!("Sq8FixedResidual : {:.4}", rec(r_sq8));
eprintln!("Sq16norm - C1 : {:+.4}", rec(r_sq16) - rec(r_sq8));
eprintln!("Sq16norm - nonorm: {:+.4}", rec(r_sq16) - rec(r_sq16_nn));
assert!(
rec(r_sq16) >= rec(r_sq8) - 0.002,
"Sq16 kernel recall {:.4} trails C1 {:.4} by >0.002 — codec-level bug reproduced",
rec(r_sq16),
rec(r_sq8)
);
}
#[test]
#[ignore = "microbench: run explicitly in release with --nocapture"]
fn rerank_kernel_leg_cost_microbench() {
use std::{hint::black_box, time::Instant};
use crate::superfile::vector::{
rerank_codec::{SQ8_FIXED_OFFSET, SQ8_FIXED_RESIDUAL_DIVISOR, SQ8_FIXED_SCALE},
sq8_simd::{Sq8EncodeConsts, encode_sq8_residual_row},
};
fn next_u64(state: &mut u64) -> u64 {
*state = state.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = *state;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
fn next_f32(state: &mut u64) -> f32 {
let u = (next_u64(state) >> 40) as f32 / (1u64 << 24) as f32; u * 2.0 - 1.0
}
fn fill_unit(state: &mut u64, out: &mut [f32]) {
for x in out.iter_mut() {
*x = next_f32(state);
}
let n = out.iter().map(|v| v * v).sum::<f32>().sqrt();
if n > 0.0 {
for x in out.iter_mut() {
*x /= n;
}
}
}
const N: usize = 100_000;
const PASSES: usize = 7;
let mut rng = 0x1234_5678_9ABC_DEF0u64;
eprintln!("\n### Rerank kernel per-candidate cost — Sq8Residual (2 legs) vs Sq16 (1 leg)");
eprintln!("N = {N} candidates, cosine/fixed-grid, median of {PASSES} timed passes\n");
eprintln!(
"{:>6} {:>16} {:>16} {:>10}",
"dim", "sq8_residual ns", "sq16 ns", "sq16/sq8"
);
for &dim in &[768usize, 1536] {
let scale = vec![SQ8_FIXED_SCALE; dim];
let offset = vec![SQ8_FIXED_OFFSET; dim];
let divisor = SQ8_FIXED_RESIDUAL_DIVISOR;
let consts = Sq8EncodeConsts::from_scale_offset(&scale, &offset);
let mut query = vec![0.0f32; dim];
fill_unit(&mut rng, &mut query);
let sq16_kernel = Sq16Kernel::new(Metric::Cosine, &query);
let sq8_kernel =
Sq8ResidualKernel::new(Metric::Cosine, &query, &scale, &offset, divisor);
let mut sq16_buf = vec![0u8; N * dim * 2];
let mut sq8_code = vec![0u8; N * dim];
let mut sq8_res = vec![0u8; N * dim];
let mut sq8_norms = vec![0.0f32; N];
let mut sq16_norms = vec![0.0f32; N];
let mut cand = vec![0.0f32; dim];
let mut recon = vec![0.0f32; dim];
for i in 0..N {
fill_unit(&mut rng, &mut cand);
let sq16_code = &mut sq16_buf[i * dim * 2..(i + 1) * dim * 2];
encode_sq16_row(&cand, sq16_code);
sq16_norms[i] = sq16_decoded_norm_sq(sq16_code, dim);
let norm = encode_sq8_residual_row(
&cand,
&consts,
&scale,
&offset,
&mut sq8_code[i * dim..(i + 1) * dim],
&mut sq8_res[i * dim..(i + 1) * dim],
&mut recon,
true,
divisor,
)
.expect("store_norm=true yields a per-doc norm");
sq8_norms[i] = norm;
}
let time_sq16 = || {
let mut sink = 0.0f32;
for i in 0..N {
let code = black_box(&sq16_buf[i * dim * 2..(i + 1) * dim * 2]);
sink += sq16_kernel.distance_with_norm(code, Some(black_box(sq16_norms[i])));
}
black_box(sink);
};
let time_sq8 = || {
let mut sink = 0.0f32;
for i in 0..N {
let code = black_box(&sq8_code[i * dim..(i + 1) * dim]);
let res = black_box(&sq8_res[i * dim..(i + 1) * dim]);
sink += sq8_kernel.distance_with_norm(code, res, Some(black_box(sq8_norms[i])));
}
black_box(sink);
};
let median = |f: &mut dyn FnMut()| -> f64 {
f(); let mut samples: Vec<f64> = (0..PASSES)
.map(|_| {
let t = Instant::now();
f();
t.elapsed().as_secs_f64()
})
.collect();
samples.sort_by(|a, b| a.total_cmp(b));
samples[PASSES / 2]
};
let mut sq16_fn = time_sq16;
let mut sq8_fn = time_sq8;
let sq16_secs = median(&mut sq16_fn);
let sq8_secs = median(&mut sq8_fn);
let sq16_ns = sq16_secs / N as f64 * 1e9;
let sq8_ns = sq8_secs / N as f64 * 1e9;
eprintln!(
"{dim:>6} {sq8_ns:>16.2} {sq16_ns:>16.2} {:>10.3}",
sq16_ns / sq8_ns
);
}
eprintln!();
}
#[test]
fn sq8_full_round_trip_within_recall_tolerance_of_fp32() {
let dim = 16usize;
let n_docs = 32usize;
let query: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.5).collect();
let corpus: Vec<f32> = (0..n_docs)
.flat_map(|i| (0..dim).map(move |j| ((i * 7 + j * 3) as f32 % 32.0) - 8.0))
.collect();
let mut min_v = vec![f32::INFINITY; dim];
let mut max_v = vec![f32::NEG_INFINITY; dim];
for row in corpus.chunks_exact(dim) {
for (d, &x) in row.iter().enumerate() {
min_v[d] = min_v[d].min(x);
max_v[d] = max_v[d].max(x);
}
}
for d in 0..dim {
assert!(
max_v[d] - min_v[d] > 0.0,
"test corpus must span each dim: dim {d} has min == max"
);
}
let mut scale = vec![0.0f32; dim];
let mut offset = vec![0.0f32; dim];
for d in 0..dim {
offset[d] = min_v[d];
scale[d] = (max_v[d] - min_v[d]) / 255.0;
}
let codes_all = encode_sq8(&corpus, dim, &scale, &offset);
let decoded_all = decode_sq8(&codes_all, dim, &scale, &offset);
let per_doc_norms: Vec<f32> = decoded_all
.chunks_exact(dim)
.map(|row| row.iter().map(|x| x * x).sum::<f32>())
.collect();
for m in [Metric::Cosine, Metric::L2Sq, Metric::NegDot] {
let norms_arg: Option<Arc<[f32]>> = match m {
Metric::L2Sq | Metric::Cosine => Some(Arc::from(per_doc_norms.clone())),
Metric::NegDot => None,
};
let kernel = Sq8Kernel::new(m, &query, &scale, &offset, norms_arg);
for pos in [0u32, 1, 5, 17, 31] {
let codes_doc = &codes_all[(pos as usize) * dim..(pos as usize + 1) * dim];
let decoded_doc = &decoded_all[(pos as usize) * dim..(pos as usize + 1) * dim];
let got = kernel.distance_at(pos, codes_doc);
let want_fp32 = distance(
m,
&query,
&corpus[(pos as usize) * dim..(pos as usize + 1) * dim],
);
let want_decoded = match m {
Metric::Cosine => {
let x_norm = per_doc_norms[pos as usize].sqrt();
if x_norm > 0.0 {
1.0 - dot(&query, decoded_doc) / x_norm
} else {
1.0 - dot(&query, decoded_doc)
}
}
_ => distance(m, &query, decoded_doc),
};
assert!(
(got - want_decoded).abs() <= 1e-3,
"metric {m:?} pos {pos}: kernel {got} vs decoded ref {want_decoded}"
);
if m != Metric::Cosine {
let rel = (got - want_fp32).abs() / want_fp32.abs().max(1e-2);
assert!(
rel <= 0.1 || (got - want_fp32).abs() <= 1.0,
"metric {m:?} pos {pos}: Sq8 {got} vs fp32 {want_fp32} (rel {rel})"
);
}
}
}
}
#[cfg(target_arch = "x86_64")]
fn fake_vec(dim: usize, seed: u32) -> Vec<f32> {
(0..dim)
.map(|i| {
let x = ((i as u32).wrapping_mul(2654435761).wrapping_add(seed)) as i32;
(x as f32) * 1e-9
})
.collect()
}
#[test]
#[cfg(target_arch = "x86_64")]
fn dot_avx512_matches_wide_across_lengths() {
if !avx512_enabled() {
eprintln!("dot_avx512_matches_wide_across_lengths: skipped, no AVX-512");
return;
}
for dim in 1..=64 {
let a = fake_vec(dim, 0xA5A5);
let b = fake_vec(dim, 0x5A5A);
let want = dot_wide(&a, &b);
let got = unsafe { dot_avx512(&a, &b) };
let tol = 1e-5 * want.abs().max(1.0);
assert!(
(want - got).abs() <= tol,
"dim {dim}: avx512 {got} vs wide {want} (tol {tol})"
);
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn l2_sq_avx512_matches_wide_across_lengths() {
if !avx512_enabled() {
eprintln!("l2_sq_avx512_matches_wide_across_lengths: skipped, no AVX-512");
return;
}
for dim in 1..=64 {
let a = fake_vec(dim, 0xDEAD);
let b = fake_vec(dim, 0xBEEF);
let want = l2_sq_wide(&a, &b);
let got = unsafe { l2_sq_avx512(&a, &b) };
let tol = 1e-5 * want.abs().max(1.0);
assert!(
(want - got).abs() <= tol,
"dim {dim}: avx512 {got} vs wide {want} (tol {tol})"
);
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn dot_avx512_matches_wide_at_embedding_dims() {
if !avx512_enabled() {
eprintln!("dot_avx512_matches_wide_at_embedding_dims: skipped, no AVX-512");
return;
}
for &dim in &[128usize, 384, 768, 1024, 1536] {
let a: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.001 - 0.5).collect();
let b: Vec<f32> = (0..dim).map(|i| ((i + 7) as f32) * 0.0017 - 0.3).collect();
let want = dot_wide(&a, &b);
let got = unsafe { dot_avx512(&a, &b) };
let tol = 1e-4 * want.abs().max(1.0);
assert!(
(want - got).abs() <= tol,
"dim {dim}: avx512 {got} vs wide {want} (tol {tol})"
);
}
}
#[test]
fn public_dot_dispatches_consistently() {
for &dim in &[7usize, 16, 17, 384] {
let a: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.01).collect();
let b: Vec<f32> = (0..dim).map(|i| ((i * 3) as f32) * 0.02 - 0.1).collect();
let public_result = dot(&a, &b);
let wide_result = dot_wide(&a, &b);
let tol = 1e-4 * wide_result.abs().max(1.0);
assert!(
(public_result - wide_result).abs() <= tol,
"dim {dim}: dot() {public_result} vs dot_wide() {wide_result} (tol {tol})"
);
}
}
#[test]
fn disable_env_var_parses_truthy_values() {
fn parse(v: &str) -> bool {
v == "1" || v.eq_ignore_ascii_case("true")
}
assert!(parse("1"));
assert!(parse("true"));
assert!(parse("TRUE"));
assert!(parse("True"));
assert!(!parse("0"));
assert!(!parse("false"));
assert!(!parse(""));
assert!(!parse("yes")); }
#[test]
#[cfg(target_arch = "x86_64")]
fn sq8_dot_avx512_matches_wide_across_lengths() {
if !avx512_enabled() {
eprintln!("sq8_dot_avx512_matches_wide_across_lengths: skipped, no AVX-512");
return;
}
for dim in [1usize, 7, 15, 16, 17, 31, 32, 33, 64, 96, 128, 384, 768] {
let q_prime: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.013 - 0.4).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 17 + 3) % 256) as u8).collect();
let want = sq8_dot_wide(&q_prime, &codes, dim);
let got = unsafe { sq8_dot_avx512(&q_prime, &codes, dim) };
let tol = 1e-5 * want.abs().max(1.0);
assert!(
(want - got).abs() <= tol,
"dim {dim}: sq8 avx512 {got} vs sq8 wide {want} (tol {tol})"
);
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn sq8_dot_avx2_matches_wide_across_lengths() {
if !avx2_enabled() {
eprintln!("sq8_dot_avx2_matches_wide_across_lengths: skipped, no AVX2");
return;
}
for dim in [
1usize, 7, 8, 9, 15, 16, 17, 31, 32, 33, 64, 96, 128, 384, 768,
] {
let q_prime: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.013 - 0.4).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 17 + 3) % 256) as u8).collect();
let want = sq8_dot_wide(&q_prime, &codes, dim);
let got = unsafe { sq8_dot_avx2(&q_prime, &codes, dim) };
let tol = 1e-5 * want.abs().max(1.0);
assert!(
(want - got).abs() <= tol,
"dim {dim}: sq8 avx2 {got} vs wide {want} (tol {tol})"
);
}
}
#[cfg(target_arch = "x86_64")]
fn time_ns<R, F: FnMut() -> R>(iters: u32, mut f: F) -> f64 {
use std::{hint::black_box, time::Instant};
for _ in 0..(iters / 10).max(64) {
black_box(f());
}
let t = Instant::now();
for _ in 0..iters {
black_box(f());
}
let dt = t.elapsed();
dt.as_secs_f64() * 1e9 / (iters as f64)
}
#[cfg(target_arch = "x86_64")]
fn realistic_dims() -> &'static [usize] {
&[128, 384, 768, 1024, 1536]
}
#[test]
#[ignore]
#[cfg(target_arch = "x86_64")]
fn avx512_microbench_distance_kernels() {
if !avx512_enabled() {
eprintln!("avx512_microbench: skipped, no AVX-512 on this host");
return;
}
eprintln!();
eprintln!(
"### distance kernel — AVX-512 vs wide (ns per call, single thread, release build)\n"
);
eprintln!("| kernel | dim | wide ns | avx512 ns | speedup |");
eprintln!("|--------|----:|--------:|----------:|--------:|");
use std::hint::black_box;
for &dim in realistic_dims() {
let a: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.001 - 0.5).collect();
let b: Vec<f32> = (0..dim).map(|i| ((i + 7) as f32) * 0.0017 - 0.3).collect();
let iters: u32 = (10_000_000u64 / (dim as u64).max(1)).max(50_000) as u32;
let wide_ns = time_ns(iters, || dot_wide(black_box(&a), black_box(&b)));
let avx_ns = time_ns(iters, || unsafe {
dot_avx512(black_box(&a), black_box(&b))
});
eprintln!(
"| `distance::dot` | {dim} | {:>7.1} | {:>7.1} | {:>5.2}× |",
wide_ns,
avx_ns,
wide_ns / avx_ns,
);
let wide_ns = time_ns(iters, || l2_sq_wide(black_box(&a), black_box(&b)));
let avx_ns = time_ns(iters, || unsafe {
l2_sq_avx512(black_box(&a), black_box(&b))
});
eprintln!(
"| `distance::l2_sq` | {dim} | {:>7.1} | {:>7.1} | {:>5.2}× |",
wide_ns,
avx_ns,
wide_ns / avx_ns,
);
}
}
#[test]
#[ignore]
#[cfg(target_arch = "x86_64")]
fn avx512_microbench_sq8_kernel() {
if !avx512_enabled() {
eprintln!("avx512_microbench: skipped, no AVX-512 on this host");
return;
}
eprintln!();
eprintln!(
"### Sq8 cross-product kernel — AVX-512 (vpmovzxbd widen) vs wide (ns per call)\n"
);
eprintln!("| kernel | dim | wide ns | avx512 ns | speedup |");
eprintln!("|--------|----:|--------:|----------:|--------:|");
use std::hint::black_box;
for &dim in realistic_dims() {
let q_prime: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.013 - 0.4).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 17 + 3) % 256) as u8).collect();
let iters: u32 = (10_000_000u64 / (dim as u64).max(1)).max(50_000) as u32;
let wide_ns = time_ns(iters, || {
sq8_dot_wide(black_box(&q_prime), black_box(&codes), black_box(dim))
});
let avx_ns = time_ns(iters, || unsafe {
sq8_dot_avx512(black_box(&q_prime), black_box(&codes), black_box(dim))
});
eprintln!(
"| `Sq8Kernel::distance_at` (dot) | {dim} | {:>7.1} | {:>7.1} | {:>5.2}× |",
wide_ns,
avx_ns,
wide_ns / avx_ns,
);
}
}
#[cfg(target_arch = "x86_64")]
fn dot_scalar(a: &[f32], b: &[f32]) -> f32 {
let mut s = 0.0f32;
for i in 0..a.len() {
s += a[i] * b[i];
}
s
}
#[cfg(target_arch = "x86_64")]
fn l2_sq_scalar(a: &[f32], b: &[f32]) -> f32 {
let mut s = 0.0f32;
for i in 0..a.len() {
let d = a[i] - b[i];
s += d * d;
}
s
}
#[cfg(target_arch = "x86_64")]
fn sq8_dot_scalar(q_prime: &[f32], code_bytes: &[u8], dim: usize) -> f32 {
let mut s = 0.0f32;
for d in 0..dim {
s += q_prime[d] * (code_bytes[d] as f32);
}
s
}
#[test]
#[ignore = "perf microbench, not a correctness gate"]
#[cfg(target_arch = "x86_64")]
fn simd_microbench_all_tiers() {
use std::hint::black_box;
let avx2 = avx2_enabled();
let avx512 = avx512_enabled();
eprintln!();
eprintln!(
"### vector distance kernels — per-tier ns / call on this host (single thread, release)\n"
);
eprintln!("host caps: avx2={avx2}, avx512f={avx512}");
eprintln!(
"build: `target-cpu=x86-64-v3` (Haswell+AVX2+FMA baseline) from .cargo/config.toml\n"
);
eprintln!("| kernel | dim | scalar ns | wide ns | avx2 ns | avx512 ns |");
eprintln!("|--------|----:|----------:|--------:|--------:|----------:|");
fn avx2_cell(v: Option<f64>, wide_ns: f64) -> String {
match v {
Some(x) => format!("{:>7.1}", x),
None => format!("wide(={:>5.1})", wide_ns),
}
}
fn avx512_cell(v: Option<f64>) -> String {
match v {
Some(x) => format!("{:>7.1}", x),
None => " —".to_string(),
}
}
for &dim in realistic_dims() {
let a: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.001 - 0.5).collect();
let b: Vec<f32> = (0..dim).map(|i| ((i + 7) as f32) * 0.0017 - 0.3).collect();
let q_prime: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.013 - 0.4).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 17 + 3) % 256) as u8).collect();
let iters: u32 = (10_000_000u64 / (dim as u64).max(1)).max(50_000) as u32;
let s = time_ns(iters, || dot_scalar(black_box(&a), black_box(&b)));
let w = time_ns(iters, || dot_wide(black_box(&a), black_box(&b)));
let a2 = None::<f64>;
let a5 = if avx512 {
Some(time_ns(iters, || unsafe {
dot_avx512(black_box(&a), black_box(&b))
}))
} else {
None
};
eprintln!(
"| `distance::dot` (fp32) | {dim} | {:>9.1} | {:>7.1} | {} | {} |",
s,
w,
avx2_cell(a2, w),
avx512_cell(a5),
);
let s = time_ns(iters, || l2_sq_scalar(black_box(&a), black_box(&b)));
let w = time_ns(iters, || l2_sq_wide(black_box(&a), black_box(&b)));
let a2 = None::<f64>;
let a5 = if avx512 {
Some(time_ns(iters, || unsafe {
l2_sq_avx512(black_box(&a), black_box(&b))
}))
} else {
None
};
eprintln!(
"| `distance::l2_sq` (fp32) | {dim} | {:>9.1} | {:>7.1} | {} | {} |",
s,
w,
avx2_cell(a2, w),
avx512_cell(a5),
);
let s = time_ns(iters, || {
sq8_dot_scalar(black_box(&q_prime), black_box(&codes), black_box(dim))
});
let w = time_ns(iters, || {
sq8_dot_wide(black_box(&q_prime), black_box(&codes), black_box(dim))
});
let a2 = if avx2 {
Some(time_ns(iters, || unsafe {
sq8_dot_avx2(black_box(&q_prime), black_box(&codes), black_box(dim))
}))
} else {
None
};
let a5 = if avx512 {
Some(time_ns(iters, || unsafe {
sq8_dot_avx512(black_box(&q_prime), black_box(&codes), black_box(dim))
}))
} else {
None
};
eprintln!(
"| `Sq8Kernel::distance_at` (dot) | {dim} | {:>9.1} | {:>7.1} | {} | {} |",
s,
w,
avx2_cell(a2, w),
avx512_cell(a5),
);
}
eprintln!();
eprintln!(
"Notes: `wide(=N.N)` in the AVX2 column means there is no \
dedicated AVX2 kernel — the dispatch on an AVX2-only host \
actually runs the wide kernel at that timing. This applies to \
the fp32 `dot` / `l2_sq` kernels because `wide::f32x8` on \
`target-cpu=x86-64-v3` lowers to `__m256` + `vfmadd*ps`, \
which is what a hand-written AVX2 kernel would emit. The \
Sq8 widen kernel has a dedicated AVX2 path (visible \
above) because the wide path previously did per-lane scalar \
widening; the dedicated AVX2 path replaces that with \
VPMOVZXBD / VPMOVZXWD + shift."
);
}
#[test]
#[ignore]
#[cfg(target_arch = "x86_64")]
fn avx2_microbench_widen_kernels() {
if !avx2_enabled() {
eprintln!("avx2_microbench: skipped, no AVX2 on this host");
return;
}
eprintln!();
eprintln!("### AVX2 widen + FMA vs portable scalar-widen wide path (ns per call)\n");
eprintln!("| kernel | dim | wide ns | avx2 ns | speedup |");
eprintln!("|--------|----:|--------:|--------:|--------:|");
use std::hint::black_box;
for &dim in realistic_dims() {
let q_prime: Vec<f32> = (0..dim).map(|i| (i as f32) * 0.013 - 0.4).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 17 + 3) % 256) as u8).collect();
let iters: u32 = (10_000_000u64 / (dim as u64).max(1)).max(50_000) as u32;
let wide_sq8_ns = time_ns(iters, || {
sq8_dot_wide(black_box(&q_prime), black_box(&codes), black_box(dim))
});
let avx2_sq8_ns = time_ns(iters, || unsafe {
sq8_dot_avx2(black_box(&q_prime), black_box(&codes), black_box(dim))
});
eprintln!(
"| `Sq8Kernel::distance_at` (dot) | {dim} | {:>7.1} | {:>7.1} | {:>5.2}× |",
wide_sq8_ns,
avx2_sq8_ns,
wide_sq8_ns / avx2_sq8_ns,
);
}
}
}