use std::sync::atomic::{AtomicU64, Ordering};
use rayon::prelude::*;
pub const SINGLE_QUERY_PARALLEL_MIN_BLOCKS: usize = 1024;
pub fn single_query_parallelizes(n_vectors: usize) -> bool {
n_vectors.div_ceil(crate::BLOCK) >= SINGLE_QUERY_PARALLEL_MIN_BLOCKS
}
pub(crate) const MIN_TILE_BLOCKS: usize = 1024;
#[cfg(target_arch = "x86_64")]
#[inline]
fn serial_required(mask_present: bool, simd_ok: bool, force_scalar_any: bool) -> bool {
mask_present || !simd_ok || force_scalar_any
}
#[inline]
fn n_block_ranges(
nq: usize,
n_quads: usize,
n_blocks: usize,
n_vectors: usize,
k: usize,
n_threads: usize,
tiles_per_thread: usize,
min_tile_blocks: usize,
serial: bool,
) -> usize {
if n_threads == 1 || serial || (nq == 1 && !single_query_parallelizes(n_vectors)) {
return 1;
}
(n_threads * tiles_per_thread)
.div_ceil(n_quads)
.min(n_blocks.div_ceil(min_tile_blocks))
.min(range_cap_for_k(n_vectors, k))
.max(1)
}
const TILES_PER_THREAD: usize = 32;
#[cfg(target_arch = "aarch64")]
const TILES_PER_THREAD_NEON: usize = TILES_PER_THREAD * 2;
#[cfg(target_arch = "aarch64")]
const MIN_TILE_BLOCKS_NEON: usize = MIN_TILE_BLOCKS / 2;
#[cfg(target_arch = "x86_64")]
const MIN_TILE_BLOCKS_X86: usize = MIN_TILE_BLOCKS * 3;
#[inline(always)]
fn rescan_min(hs: &[f32], hi: &[u64], k: usize) -> (f32, usize) {
let mut mi = 0usize;
for h in 1..k {
if hs[h] < hs[mi] || (hs[h] == hs[mi] && hi[h] > hi[mi]) {
mi = h;
}
}
(hs[mi], mi)
}
#[inline]
fn range_cap_for_k(n_vectors: usize, k: usize) -> usize {
const MIN_VECTORS_PER_RANGE_PER_K: usize = 512;
n_vectors
.div_ceil(MIN_VECTORS_PER_RANGE_PER_K * k.max(1))
.max(1)
}
#[inline]
fn smooth_tile_count(n_ranges: usize, n_quads: usize, n_threads: usize) -> usize {
let tiles = n_quads * n_ranges;
if tiles > n_threads && tiles < 2 * n_threads {
(n_threads / n_quads).max(1)
} else {
n_ranges
}
}
use crate::rotation::Rotation;
use crate::{BLOCK, FLUSH_EVERY};
pub(crate) static BLOCKS_SKIPPED_BY_MASK: AtomicU64 = AtomicU64::new(0);
#[cfg(test)]
#[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
pub(crate) static FORCE_SCALAR_FALLBACK: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
pub fn blocks_skipped_by_mask() -> Option<u64> {
#[cfg(feature = "mask-skip-counter")]
{
Some(BLOCKS_SKIPPED_BY_MASK.load(Ordering::Relaxed))
}
#[cfg(not(feature = "mask-skip-counter"))]
{
None
}
}
pub fn reset_blocks_skipped_by_mask() {
BLOCKS_SKIPPED_BY_MASK.store(0, Ordering::Relaxed);
}
#[cfg(target_arch = "aarch64")]
pub(crate) unsafe fn score_4bit_block_neon(
blocked_codes: &[u8],
uint8_luts: &[u8],
block_offset: usize,
n_byte_groups: usize,
scale: f32,
bias: f32,
vec_scales: &[f32],
base_vec: usize,
n_vectors: usize,
out: &mut [f32; BLOCK],
) {
use std::arch::aarch64::*;
let mask = vdupq_n_u8(0x0F);
let v_scale = vdupq_n_f32(scale);
let n_batches = (n_byte_groups + FLUSH_EVERY - 1) / FLUSH_EVERY;
let mut fa = [vdupq_n_f32(bias); 8];
let codes_base = blocked_codes.as_ptr().add(block_offset);
let luts_base = uint8_luts.as_ptr();
for batch in 0..n_batches {
let g_start = batch * FLUSH_EVERY;
let g_end = (g_start + FLUSH_EVERY).min(n_byte_groups);
let mut accum = [vdupq_n_u16(0); 4];
let mut g = g_start;
while g + 3 < g_end {
let lp0 = luts_base.add(g * 32);
let lp1 = luts_base.add((g + 1) * 32);
let lp2 = luts_base.add((g + 2) * 32);
let lp3 = luts_base.add((g + 3) * 32);
let cp0 = codes_base.add(g * BLOCK);
let cp1 = codes_base.add((g + 1) * BLOCK);
let cp2 = codes_base.add((g + 2) * BLOCK);
let cp3 = codes_base.add((g + 3) * BLOCK);
for (lp, cp) in [(lp0, cp0), (lp1, cp1), (lp2, cp2), (lp3, cp3)] {
let lut_hi = vld1q_u8(lp);
let lut_lo = vld1q_u8(lp.add(16));
let c0 = vld1q_u8(cp);
let c1 = vld1q_u8(cp.add(16));
let s0 = vaddq_u8(vqtbl1q_u8(lut_lo, vandq_u8(c0, mask)), vqtbl1q_u8(lut_hi, vshrq_n_u8(c0, 4)));
let s1 = vaddq_u8(vqtbl1q_u8(lut_lo, vandq_u8(c1, mask)), vqtbl1q_u8(lut_hi, vshrq_n_u8(c1, 4)));
accum[0] = vaddw_u8(accum[0], vget_low_u8(s0));
accum[1] = vaddw_u8(accum[1], vget_high_u8(s0));
accum[2] = vaddw_u8(accum[2], vget_low_u8(s1));
accum[3] = vaddw_u8(accum[3], vget_high_u8(s1));
}
g += 4;
}
while g < g_end {
let lp = luts_base.add(g * 32);
let lut_hi = vld1q_u8(lp);
let lut_lo = vld1q_u8(lp.add(16));
let cp = codes_base.add(g * BLOCK);
let c0 = vld1q_u8(cp);
let c1 = vld1q_u8(cp.add(16));
let s0 = vaddq_u8(vqtbl1q_u8(lut_lo, vandq_u8(c0, mask)),
vqtbl1q_u8(lut_hi, vshrq_n_u8(c0, 4)));
let s1 = vaddq_u8(vqtbl1q_u8(lut_lo, vandq_u8(c1, mask)),
vqtbl1q_u8(lut_hi, vshrq_n_u8(c1, 4)));
accum[0] = vaddw_u8(accum[0], vget_low_u8(s0));
accum[1] = vaddw_u8(accum[1], vget_high_u8(s0));
accum[2] = vaddw_u8(accum[2], vget_low_u8(s1));
accum[3] = vaddw_u8(accum[3], vget_high_u8(s1));
g += 1;
}
for i in 0..4 {
let lo = vcvtq_f32_u32(vmovl_u16(vget_low_u16(accum[i])));
let hi = vcvtq_f32_u32(vmovl_u16(vget_high_u16(accum[i])));
fa[i * 2] = vfmaq_f32(fa[i * 2], v_scale, lo);
fa[i * 2 + 1] = vfmaq_f32(fa[i * 2 + 1], v_scale, hi);
}
}
let end = (base_vec + BLOCK).min(n_vectors);
let out_ptr = out.as_mut_ptr();
let vec_scales_ptr = vec_scales.as_ptr().add(base_vec);
if end - base_vec == BLOCK {
for i in 0..8 {
let n = vld1q_f32(vec_scales_ptr.add(i * 4));
vst1q_f32(out_ptr.add(i * 4), vmulq_f32(fa[i], n));
}
} else {
let mut float_accum = [0.0f32; BLOCK];
for i in 0..8 {
vst1q_f32(float_accum.as_mut_ptr().add(i * 4), fa[i]);
}
for lane in 0..BLOCK {
*out_ptr.add(lane) = if lane < end - base_vec {
float_accum[lane] * *vec_scales_ptr.add(lane)
} else {
f32::NEG_INFINITY
};
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn search_multi_query_avx2(
blocked_codes: &[u8],
luts: &[&[u8]],
scales: &[f32],
biases: &[f32],
n_byte_groups: usize,
vec_scales: &[f32],
n_vectors: usize,
nq: usize,
k: usize,
mask: Option<&[u64]>,
heap_scores: &mut [Vec<f32>],
heap_indices: &mut [Vec<u64>],
heap_sizes: &mut [usize],
heap_mins: &mut [f32],
heap_min_idxs: &mut [usize],
) {
use std::arch::x86_64::*;
let n_blocks = (n_vectors + BLOCK - 1) / BLOCK;
let nibble_mask = _mm256_set1_epi8(0x0F);
let codes_base = blocked_codes.as_ptr();
for b in 0..n_blocks {
let base_vec = b * BLOCK;
if !block_has_allowed(mask, base_vec) {
continue;
}
let v_scales: [__m256; 4] = [
_mm256_set1_ps(scales[0]),
_mm256_set1_ps(scales[1]),
_mm256_set1_ps(scales[2]),
_mm256_set1_ps(scales[3]),
];
let v_biases: [__m256; 4] = [
_mm256_set1_ps(biases[0]),
_mm256_set1_ps(biases[1]),
_mm256_set1_ps(biases[2]),
_mm256_set1_ps(biases[3]),
];
let mut fa = [
[v_biases[0]; 4],
[v_biases[1]; 4],
[v_biases[2]; 4],
[v_biases[3]; 4],
];
let n_batches = (n_byte_groups + FLUSH_EVERY - 1) / FLUSH_EVERY;
for batch in 0..n_batches {
let g_start = batch * FLUSH_EVERY;
let g_end = (g_start + FLUSH_EVERY).min(n_byte_groups);
let mut accus = [[_mm256_setzero_si256(); 4]; 4];
for g in g_start..g_end {
let cp = codes_base.add((b * n_byte_groups + g) * BLOCK);
let codes_v = _mm256_loadu_si256(cp as *const __m256i);
let clo = _mm256_and_si256(codes_v, nibble_mask);
let chi = _mm256_and_si256(_mm256_srli_epi16(codes_v, 4), nibble_mask);
for qi in 0..4 {
let lut = _mm256_loadu_si256(luts[qi].as_ptr().add(g * 32) as *const __m256i);
let res0 = _mm256_shuffle_epi8(lut, clo);
let res1 = _mm256_shuffle_epi8(lut, chi);
accus[qi][0] = _mm256_add_epi16(accus[qi][0], res0);
accus[qi][1] = _mm256_add_epi16(accus[qi][1], _mm256_srli_epi16(res0, 8));
accus[qi][2] = _mm256_add_epi16(accus[qi][2], res1);
accus[qi][3] = _mm256_add_epi16(accus[qi][3], _mm256_srli_epi16(res1, 8));
}
}
for qi in 0..4 {
let mut lo_a0 = accus[qi][0];
let lo_a1 = accus[qi][1];
let mut hi_a2 = accus[qi][2];
let hi_a3 = accus[qi][3];
lo_a0 = _mm256_sub_epi16(lo_a0, _mm256_slli_epi16(lo_a1, 8));
hi_a2 = _mm256_sub_epi16(hi_a2, _mm256_slli_epi16(hi_a3, 8));
let dis0 = _mm256_add_epi16(
_mm256_permute2x128_si256(lo_a0, lo_a1, 0x21),
_mm256_blend_epi32(lo_a0, lo_a1, 0xF0),
);
let dis1 = _mm256_add_epi16(
_mm256_permute2x128_si256(hi_a2, hi_a3, 0x21),
_mm256_blend_epi32(hi_a2, hi_a3, 0xF0),
);
let f0 = _mm256_cvtepi32_ps(_mm256_cvtepu16_epi32(_mm256_castsi256_si128(dis0)));
let f1 = _mm256_cvtepi32_ps(_mm256_cvtepu16_epi32(_mm256_extracti128_si256(dis0, 1)));
let f2 = _mm256_cvtepi32_ps(_mm256_cvtepu16_epi32(_mm256_castsi256_si128(dis1)));
let f3 = _mm256_cvtepi32_ps(_mm256_cvtepu16_epi32(_mm256_extracti128_si256(dis1, 1)));
fa[qi][0] = _mm256_fmadd_ps(v_scales[qi], f0, fa[qi][0]);
fa[qi][1] = _mm256_fmadd_ps(v_scales[qi], f1, fa[qi][1]);
fa[qi][2] = _mm256_fmadd_ps(v_scales[qi], f2, fa[qi][2]);
fa[qi][3] = _mm256_fmadd_ps(v_scales[qi], f3, fa[qi][3]);
}
}
let end = (base_vec + BLOCK).min(n_vectors);
let vec_scales_ptr = vec_scales.as_ptr().add(base_vec);
for qi in 0..nq {
let f0 = fa[qi][0];
let f1 = fa[qi][1];
let f2 = fa[qi][2];
let f3 = fa[qi][3];
let mut block_out = [0.0f32; BLOCK];
let bp = block_out.as_mut_ptr();
if end - base_vec == BLOCK {
for (i, f) in [f0, f1, f2, f3].iter().enumerate() {
let n = _mm256_loadu_ps(vec_scales_ptr.add(i * 8));
_mm256_storeu_ps(bp.add(i * 8), _mm256_mul_ps(*f, n));
}
} else {
for (i, f) in [f0, f1, f2, f3].iter().enumerate() {
_mm256_storeu_ps(bp.add(i * 8), *f);
}
for lane in 0..(end - base_vec) {
block_out[lane] *= *vec_scales_ptr.add(lane);
}
for lane in (end - base_vec)..BLOCK {
block_out[lane] = f32::NEG_INFINITY;
}
}
let hs = &mut heap_scores[qi];
let hi = &mut heap_indices[qi];
let sz = &mut heap_sizes[qi];
let hmin = &mut heap_mins[qi];
let hmi = &mut heap_min_idxs[qi];
if *sz < k {
for lane in 0..(end - base_vec) {
if let Some(m) = mask {
if !mask_allows(m, base_vec + lane) { continue; }
}
let score = block_out[lane];
if *sz < k {
hs[*sz] = score;
hi[*sz] = (base_vec + lane) as u64;
*sz += 1;
if *sz == k {
let (m, mi) = rescan_min(hs, hi, k);
*hmin = m;
*hmi = mi;
}
} else if score > *hmin {
hs[*hmi] = score;
hi[*hmi] = (base_vec + lane) as u64;
let (m, mi) = rescan_min(hs, hi, k);
*hmin = m;
*hmi = mi;
}
}
} else {
let v_hmin = _mm256_set1_ps(*hmin);
for chunk in 0..4 {
let chunk_start = chunk * 8;
if chunk_start >= end - base_vec { break; }
let scores_v = _mm256_loadu_ps(block_out.as_ptr().add(chunk_start));
let cmp = _mm256_cmp_ps(scores_v, v_hmin, _CMP_GT_OQ);
if _mm256_movemask_ps(cmp) == 0 { continue; }
let chunk_end = (chunk_start + 8).min(end - base_vec);
for lane in chunk_start..chunk_end {
if let Some(m) = mask {
if !mask_allows(m, base_vec + lane) { continue; }
}
let score = block_out[lane];
if score > *hmin {
hs[*hmi] = score;
hi[*hmi] = (base_vec + lane) as u64;
let (m, mi) = rescan_min(hs, hi, k);
*hmin = m;
*hmi = mi;
}
}
}
}
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma", enable = "avx512f", enable = "avx512bw")]
unsafe fn search_multi_query_avx512bw(
blocked_codes: &[u8],
luts: &[&[u8]],
scales: &[f32],
biases: &[f32],
n_byte_groups: usize,
vec_scales: &[f32],
n_vectors: usize,
nq: usize,
k: usize,
mask: Option<&[u64]>,
heap_scores: &mut [Vec<f32>],
heap_indices: &mut [Vec<u64>],
heap_sizes: &mut [usize],
heap_mins: &mut [f32],
heap_min_idxs: &mut [usize],
) {
use std::arch::x86_64::*;
let n_blocks = (n_vectors + BLOCK - 1) / BLOCK;
let n_block_pairs = n_blocks / 2;
let mask512 = _mm512_set1_epi8(0x0F);
let mask256 = _mm256_set1_epi8(0x0F);
let codes_base = blocked_codes.as_ptr();
let v_biases: [__m256; 4] = [
_mm256_set1_ps(biases[0]),
_mm256_set1_ps(biases[1]),
_mm256_set1_ps(biases[2]),
_mm256_set1_ps(biases[3]),
];
let v_scales: [__m256; 4] = [
_mm256_set1_ps(scales[0]),
_mm256_set1_ps(scales[1]),
_mm256_set1_ps(scales[2]),
_mm256_set1_ps(scales[3]),
];
for p in 0..n_block_pairs {
let b0 = p * 2;
let b1 = b0 + 1;
if !block_pair_has_allowed(mask, b0 * BLOCK) {
continue;
}
let mut fa_b0: [[__m256; 4]; 4] = [
[v_biases[0]; 4],
[v_biases[1]; 4],
[v_biases[2]; 4],
[v_biases[3]; 4],
];
let mut fa_b1 = fa_b0;
debug_assert!(FLUSH_EVERY % 2 == 0);
let n_batches = (n_byte_groups + FLUSH_EVERY - 1) / FLUSH_EVERY;
for batch in 0..n_batches {
let g_start = batch * FLUSH_EVERY;
let g_end = (g_start + FLUSH_EVERY).min(n_byte_groups);
let mut accus = [[_mm512_setzero_si512(); 4]; 4];
let mut g_pair = g_start;
while g_pair + 1 < g_end {
let g0 = g_pair;
let g1 = g0 + 1;
let cp0_a = codes_base.add((b0 * n_byte_groups + g0) * BLOCK);
let cp1_a = codes_base.add((b1 * n_byte_groups + g0) * BLOCK);
let codes_a = _mm512_inserti64x4(
_mm512_castsi256_si512(_mm256_loadu_si256(cp0_a as *const __m256i)),
_mm256_loadu_si256(cp1_a as *const __m256i),
1,
);
let cp0_b = codes_base.add((b0 * n_byte_groups + g1) * BLOCK);
let cp1_b = codes_base.add((b1 * n_byte_groups + g1) * BLOCK);
let codes_b = _mm512_inserti64x4(
_mm512_castsi256_si512(_mm256_loadu_si256(cp0_b as *const __m256i)),
_mm256_loadu_si256(cp1_b as *const __m256i),
1,
);
let clo_a = _mm512_and_si512(codes_a, mask512);
let chi_a = _mm512_and_si512(_mm512_srli_epi16(codes_a, 4), mask512);
let clo_b = _mm512_and_si512(codes_b, mask512);
let chi_b = _mm512_and_si512(_mm512_srli_epi16(codes_b, 4), mask512);
for qi in 0..4 {
let lut_a = _mm512_broadcast_i64x4(
_mm256_loadu_si256(luts[qi].as_ptr().add(g0 * 32) as *const __m256i),
);
let lut_b = _mm512_broadcast_i64x4(
_mm256_loadu_si256(luts[qi].as_ptr().add(g1 * 32) as *const __m256i),
);
let res0_a = _mm512_shuffle_epi8(lut_a, clo_a);
let res1_a = _mm512_shuffle_epi8(lut_a, chi_a);
let res0_b = _mm512_shuffle_epi8(lut_b, clo_b);
let res1_b = _mm512_shuffle_epi8(lut_b, chi_b);
accus[qi][0] = _mm512_add_epi16(accus[qi][0], _mm512_add_epi16(res0_a, res0_b));
accus[qi][1] = _mm512_add_epi16(
accus[qi][1],
_mm512_add_epi16(_mm512_srli_epi16(res0_a, 8), _mm512_srli_epi16(res0_b, 8)),
);
accus[qi][2] = _mm512_add_epi16(accus[qi][2], _mm512_add_epi16(res1_a, res1_b));
accus[qi][3] = _mm512_add_epi16(
accus[qi][3],
_mm512_add_epi16(_mm512_srli_epi16(res1_a, 8), _mm512_srli_epi16(res1_b, 8)),
);
}
g_pair += 2;
}
for g in g_pair..g_end {
let cp0 = codes_base.add((b0 * n_byte_groups + g) * BLOCK);
let cp1 = codes_base.add((b1 * n_byte_groups + g) * BLOCK);
let codes_low = _mm256_loadu_si256(cp0 as *const __m256i);
let codes_high = _mm256_loadu_si256(cp1 as *const __m256i);
let codes_v = _mm512_inserti64x4(
_mm512_castsi256_si512(codes_low),
codes_high,
1,
);
let clo = _mm512_and_si512(codes_v, mask512);
let chi = _mm512_and_si512(_mm512_srli_epi16(codes_v, 4), mask512);
for qi in 0..4 {
let lut_low =
_mm256_loadu_si256(luts[qi].as_ptr().add(g * 32) as *const __m256i);
let lut = _mm512_broadcast_i64x4(lut_low);
let res0 = _mm512_shuffle_epi8(lut, clo);
let res1 = _mm512_shuffle_epi8(lut, chi);
accus[qi][0] = _mm512_add_epi16(accus[qi][0], res0);
accus[qi][1] = _mm512_add_epi16(accus[qi][1], _mm512_srli_epi16(res0, 8));
accus[qi][2] = _mm512_add_epi16(accus[qi][2], res1);
accus[qi][3] = _mm512_add_epi16(accus[qi][3], _mm512_srli_epi16(res1, 8));
}
}
for qi in 0..4 {
let block_accus_b0: [__m256i; 4] = [
_mm512_castsi512_si256(accus[qi][0]),
_mm512_castsi512_si256(accus[qi][1]),
_mm512_castsi512_si256(accus[qi][2]),
_mm512_castsi512_si256(accus[qi][3]),
];
avx2_batch_flush_to_fa(block_accus_b0, v_scales[qi], &mut fa_b0[qi]);
let block_accus_b1: [__m256i; 4] = [
_mm512_extracti64x4_epi64(accus[qi][0], 1),
_mm512_extracti64x4_epi64(accus[qi][1], 1),
_mm512_extracti64x4_epi64(accus[qi][2], 1),
_mm512_extracti64x4_epi64(accus[qi][3], 1),
];
avx2_batch_flush_to_fa(block_accus_b1, v_scales[qi], &mut fa_b1[qi]);
}
}
for which_block in 0..2usize {
let b = b0 + which_block;
let base_vec = b * BLOCK;
if base_vec >= n_vectors { break; }
if !block_has_allowed(mask, base_vec) { continue; }
let end = (base_vec + BLOCK).min(n_vectors);
let vec_scales_ptr = vec_scales.as_ptr().add(base_vec);
let fa = if which_block == 0 { &fa_b0 } else { &fa_b1 };
for qi in 0..nq {
avx2_post_flush_heap_update(
&fa[qi],
base_vec,
end,
vec_scales_ptr,
qi,
k,
mask,
heap_scores,
heap_indices,
heap_sizes,
heap_mins,
heap_min_idxs,
);
}
}
}
let bulk_blocks = n_block_pairs * 2;
if bulk_blocks < n_blocks {
let b = bulk_blocks;
let base_vec = b * BLOCK;
if !block_has_allowed(mask, base_vec) {
return;
}
let mut fa: [[__m256; 4]; 4] = [
[v_biases[0]; 4],
[v_biases[1]; 4],
[v_biases[2]; 4],
[v_biases[3]; 4],
];
let n_batches = (n_byte_groups + FLUSH_EVERY - 1) / FLUSH_EVERY;
for batch in 0..n_batches {
let g_start = batch * FLUSH_EVERY;
let g_end = (g_start + FLUSH_EVERY).min(n_byte_groups);
let mut accus = [[_mm256_setzero_si256(); 4]; 4];
for g in g_start..g_end {
let cp = codes_base.add((b * n_byte_groups + g) * BLOCK);
let codes_v = _mm256_loadu_si256(cp as *const __m256i);
let clo = _mm256_and_si256(codes_v, mask256);
let chi = _mm256_and_si256(_mm256_srli_epi16(codes_v, 4), mask256);
for qi in 0..4 {
let lut = _mm256_loadu_si256(luts[qi].as_ptr().add(g * 32) as *const __m256i);
let res0 = _mm256_shuffle_epi8(lut, clo);
let res1 = _mm256_shuffle_epi8(lut, chi);
accus[qi][0] = _mm256_add_epi16(accus[qi][0], res0);
accus[qi][1] = _mm256_add_epi16(accus[qi][1], _mm256_srli_epi16(res0, 8));
accus[qi][2] = _mm256_add_epi16(accus[qi][2], res1);
accus[qi][3] = _mm256_add_epi16(accus[qi][3], _mm256_srli_epi16(res1, 8));
}
}
for qi in 0..4 {
avx2_batch_flush_to_fa(
[accus[qi][0], accus[qi][1], accus[qi][2], accus[qi][3]],
v_scales[qi],
&mut fa[qi],
);
}
}
let end = (base_vec + BLOCK).min(n_vectors);
let vec_scales_ptr = vec_scales.as_ptr().add(base_vec);
for qi in 0..nq {
avx2_post_flush_heap_update(
&fa[qi],
base_vec,
end,
vec_scales_ptr,
qi,
k,
mask,
heap_scores,
heap_indices,
heap_sizes,
heap_mins,
heap_min_idxs,
);
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f", enable = "avx512bw", enable = "avx512vnni", enable = "avx512vl", enable = "avx2", enable = "fma")]
#[allow(clippy::too_many_arguments)]
unsafe fn search_multi_query_vnni_dispatch(
blocked_codes: &[u8],
split_luts: &[&[u8]],
scales: &[f32],
biases: &[f32],
n_byte_groups: usize,
vec_scales: &[f32],
n_vectors: usize,
nq: usize,
k: usize,
mask: Option<&[u64]>,
heap_scores: &mut [Vec<f32>],
heap_indices: &mut [Vec<u64>],
heap_sizes: &mut [usize],
heap_mins: &mut [f32],
heap_min_idxs: &mut [usize],
) {
if nq == 1 {
search_single_query_vnni_blk2(
blocked_codes, split_luts, scales, biases, n_byte_groups, vec_scales,
n_vectors, k, mask, heap_scores, heap_indices, heap_sizes,
heap_mins, heap_min_idxs,
)
} else {
search_multi_query_vnni::<false>(
blocked_codes, split_luts, scales, biases, n_byte_groups, vec_scales,
n_vectors, nq, k, mask, heap_scores, heap_indices, heap_sizes,
heap_mins, heap_min_idxs,
)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(
enable = "avx2",
enable = "fma",
enable = "avx512f",
enable = "avx512bw",
enable = "avx512vbmi",
enable = "avx512vnni"
)]
#[allow(clippy::too_many_arguments)]
unsafe fn search_multi_query_vnni<const PF: bool>(
blocked_codes: &[u8],
split_luts: &[&[u8]],
scales: &[f32],
biases: &[f32],
n_byte_groups: usize,
vec_scales: &[f32],
n_vectors: usize,
nq: usize,
k: usize,
mask: Option<&[u64]>,
heap_scores: &mut [Vec<f32>],
heap_indices: &mut [Vec<u64>],
heap_sizes: &mut [usize],
heap_mins: &mut [f32],
heap_min_idxs: &mut [usize],
) {
use std::arch::x86_64::*;
debug_assert!(nq <= 8, "search_multi_query_vnni is 8-wide; got nq={nq}");
let n_blocks = n_vectors.div_ceil(BLOCK);
let m0f = _mm512_set1_epi8(0x0F);
let kpos = _mm512_set1_epi32(0x3020_1000u32 as i32);
let ones = _mm512_set1_epi8(1);
let quads = n_byte_groups / 4;
let block_bytes = n_byte_groups * BLOCK;
for b in 0..n_blocks {
let base_vec = b * BLOCK;
if !block_has_allowed(mask, base_vec) {
continue;
}
let block_base = b * block_bytes;
let mut acc = [[_mm512_setzero_si512(); 2]; 8];
for q4 in 0..quads {
for h in 0..2 {
if PF {
let pf = block_base + (q4 + 8) * 128 + h * 64;
if pf + 64 <= blocked_codes.len() {
_mm_prefetch(blocked_codes.as_ptr().add(pf) as *const i8, _MM_HINT_T0);
}
}
let c = _mm512_loadu_si512(
blocked_codes.as_ptr().add(block_base + q4 * 128 + h * 64) as *const __m512i,
);
let ilo = _mm512_or_si512(_mm512_and_si512(c, m0f), kpos);
let ihi = _mm512_or_si512(
_mm512_and_si512(_mm512_srli_epi16(c, 4), m0f),
kpos,
);
for qi in 0..nq.min(8) {
let tp = split_luts[qi].as_ptr().add(q4 * 128);
let tlo = _mm512_loadu_si512(tp as *const __m512i);
let thi = _mm512_loadu_si512(tp.add(64) as *const __m512i);
acc[qi][h] = _mm512_dpbusd_epi32(
acc[qi][h],
_mm512_permutexvar_epi8(ilo, tlo),
ones,
);
acc[qi][h] = _mm512_dpbusd_epi32(
acc[qi][h],
_mm512_permutexvar_epi8(ihi, thi),
ones,
);
}
}
}
let end = (base_vec + BLOCK).min(n_vectors);
for qi in 0..nq.min(8) {
let vs = _mm512_set1_ps(scales[qi]);
let vb = _mm512_set1_ps(biases[qi]);
let f0 = _mm512_add_ps(_mm512_mul_ps(_mm512_cvtepi32_ps(acc[qi][0]), vs), vb);
let f1 = _mm512_add_ps(_mm512_mul_ps(_mm512_cvtepi32_ps(acc[qi][1]), vs), vb);
avx512_post_flush_heap_update(
f0,
f1,
base_vec,
end,
vec_scales.as_ptr().add(base_vec),
qi,
k,
mask,
heap_scores,
heap_indices,
heap_sizes,
heap_mins,
heap_min_idxs,
);
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(
enable = "avx2",
enable = "fma",
enable = "avx512f",
enable = "avx512bw",
// vbmi is what makes `vpermb` a real instruction. Without it LLVM
// emulates `_mm512_permutexvar_epi8` and the kernel runs 3x slower —
// the first two builds of this hypothesis measured exactly that, and
// the feature list, not the loop structure, was the defect.
enable = "avx512vbmi",
enable = "avx512vnni"
)]
#[allow(clippy::too_many_arguments)]
unsafe fn search_single_query_vnni_blk2(
blocked_codes: &[u8],
split_luts: &[&[u8]],
scales: &[f32],
biases: &[f32],
n_byte_groups: usize,
vec_scales: &[f32],
n_vectors: usize,
k: usize,
mask: Option<&[u64]>,
heap_scores: &mut [Vec<f32>],
heap_indices: &mut [Vec<u64>],
heap_sizes: &mut [usize],
heap_mins: &mut [f32],
heap_min_idxs: &mut [usize],
) {
use std::arch::x86_64::*;
let n_blocks = n_vectors.div_ceil(BLOCK);
let m0f = _mm512_set1_epi8(0x0F);
let kpos = _mm512_set1_epi32(0x3020_1000u32 as i32);
let ones = _mm512_set1_epi8(1);
let quads = n_byte_groups / 4;
let block_bytes = n_byte_groups * BLOCK;
let n_pairs = n_blocks / 2;
for pb in 0..n_pairs {
let b = pb * 2;
if !block_has_allowed(mask, b * BLOCK) && !block_has_allowed(mask, (b + 1) * BLOCK) {
continue;
}
let base0 = b * block_bytes;
let base1 = base0 + block_bytes;
let mut a0 = [_mm512_setzero_si512(); 2];
let mut a1 = [_mm512_setzero_si512(); 2];
for q4 in 0..quads {
for h in 0..2 {
let pf = base0 + (q4 + 8) * 128 + h * 64;
if pf + 64 <= blocked_codes.len() {
_mm_prefetch(blocked_codes.as_ptr().add(pf) as *const i8, _MM_HINT_T0);
}
let tp = split_luts[0].as_ptr().add(q4 * 128);
let tlo = _mm512_loadu_si512(tp as *const __m512i);
let thi = _mm512_loadu_si512(tp.add(64) as *const __m512i);
let c0 = _mm512_loadu_si512(
blocked_codes.as_ptr().add(base0 + q4 * 128 + h * 64) as *const __m512i);
let c1 = _mm512_loadu_si512(
blocked_codes.as_ptr().add(base1 + q4 * 128 + h * 64) as *const __m512i);
let i0lo = _mm512_or_si512(_mm512_and_si512(c0, m0f), kpos);
let i0hi = _mm512_or_si512(
_mm512_and_si512(_mm512_srli_epi16(c0, 4), m0f), kpos);
let i1lo = _mm512_or_si512(_mm512_and_si512(c1, m0f), kpos);
let i1hi = _mm512_or_si512(
_mm512_and_si512(_mm512_srli_epi16(c1, 4), m0f), kpos);
a0[h] = _mm512_dpbusd_epi32(a0[h], _mm512_permutexvar_epi8(i0lo, tlo), ones);
a0[h] = _mm512_dpbusd_epi32(a0[h], _mm512_permutexvar_epi8(i0hi, thi), ones);
a1[h] = _mm512_dpbusd_epi32(a1[h], _mm512_permutexvar_epi8(i1lo, tlo), ones);
a1[h] = _mm512_dpbusd_epi32(a1[h], _mm512_permutexvar_epi8(i1hi, thi), ones);
}
}
let vs = _mm512_set1_ps(scales[0]);
let vb = _mm512_set1_ps(biases[0]);
for (i, a) in [a0, a1].iter().enumerate() {
let base_vec = (b + i) * BLOCK;
let end = (base_vec + BLOCK).min(n_vectors);
let f0 = _mm512_add_ps(_mm512_mul_ps(_mm512_cvtepi32_ps(a[0]), vs), vb);
let f1 = _mm512_add_ps(_mm512_mul_ps(_mm512_cvtepi32_ps(a[1]), vs), vb);
avx512_post_flush_heap_update(
f0, f1, base_vec, end, vec_scales.as_ptr().add(base_vec), 0, k, mask,
heap_scores, heap_indices, heap_sizes, heap_mins, heap_min_idxs,
);
}
}
for b in (n_pairs * 2)..n_blocks {
if block_has_allowed(mask, b * BLOCK) {
let base = b * block_bytes;
let mut a = [_mm512_setzero_si512(); 2];
for q4 in 0..quads {
for h in 0..2 {
let tp = split_luts[0].as_ptr().add(q4 * 128);
let tlo = _mm512_loadu_si512(tp as *const __m512i);
let thi = _mm512_loadu_si512(tp.add(64) as *const __m512i);
let c = _mm512_loadu_si512(
blocked_codes.as_ptr().add(base + q4 * 128 + h * 64) as *const __m512i);
let ilo = _mm512_or_si512(_mm512_and_si512(c, m0f), kpos);
let ihi = _mm512_or_si512(
_mm512_and_si512(_mm512_srli_epi16(c, 4), m0f), kpos);
a[h] = _mm512_dpbusd_epi32(a[h], _mm512_permutexvar_epi8(ilo, tlo), ones);
a[h] = _mm512_dpbusd_epi32(a[h], _mm512_permutexvar_epi8(ihi, thi), ones);
}
}
let base_vec = b * BLOCK;
let end = (base_vec + BLOCK).min(n_vectors);
let vs = _mm512_set1_ps(scales[0]);
let vb = _mm512_set1_ps(biases[0]);
let f0 = _mm512_add_ps(_mm512_mul_ps(_mm512_cvtepi32_ps(a[0]), vs), vb);
let f1 = _mm512_add_ps(_mm512_mul_ps(_mm512_cvtepi32_ps(a[1]), vs), vb);
avx512_post_flush_heap_update(
f0, f1, base_vec, end, vec_scales.as_ptr().add(base_vec), 0, k, mask,
heap_scores, heap_indices, heap_sizes, heap_mins, heap_min_idxs,
);
}
}
}
#[cfg(target_arch = "x86_64")]
pub(crate) fn have_gfni() -> bool {
static G: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*G.get_or_init(|| {
std::env::var_os("TURBOVEC_NO_GFNI").is_none() && is_x86_feature_detected!("gfni")
})
}
#[cfg(target_arch = "x86_64")]
macro_rules! define_permute_dot {
($name:ident, [$($feature:literal),*], |$c:ident, $mask:ident| $hi:expr) => {
#[target_feature($(enable = $feature),*)]
#[allow(clippy::too_many_arguments)]
unsafe fn $name<const NQ: usize, const BLK: usize>(
blocked_codes: &[u8],
pds: &[&QueryPermuteDot; NQ],
n_byte_groups: usize,
vec_scales: &[f32],
n_vectors: usize,
nq: usize,
k: usize,
mask: Option<&[u64]>,
heap_scores: &mut [Vec<f32>],
heap_indices: &mut [Vec<u64>],
heap_sizes: &mut [usize],
heap_mins: &mut [f32],
heap_min_idxs: &mut [usize],
) {
use std::arch::x86_64::*;
let n_blocks = n_vectors.div_ceil(BLOCK);
let m0f = _mm512_set1_epi8(0x0F);
let quads = n_byte_groups / 4;
let block_bytes = n_byte_groups * BLOCK;
let nqm = nq.min(NQ);
let levels = _mm512_xor_si512(
_mm512_broadcast_i32x4(_mm_loadu_si128(pds[0].levels.as_ptr() as *const __m128i)),
_mm512_set1_epi8(0x80u8 as i8),
);
for b in (0..n_blocks).step_by(BLK) {
let base_vec = b * BLOCK;
let group_blocks = BLK.min(n_blocks - b);
if !(0..group_blocks).any(|s| block_has_allowed(mask, base_vec + s * BLOCK)) {
continue;
}
let block_base = b * block_bytes;
let nb = BLK.min(n_blocks - b);
let mut acc = [[[_mm512_setzero_si512(); 2]; NQ]; BLK];
for sub in 0..nb {
for (a, pd) in acc[sub].iter_mut().zip(pds.iter()) {
let z = _mm512_set1_epi32(pd.zero);
a[0] = z;
a[1] = z;
}
}
for q4 in 0..quads {
for h in 0..2 {
for sub in 0..nb {
let block_base = block_base + sub * block_bytes;
let acc = &mut acc[sub];
let pf_quads = if NQ == 1 { 8 } else { 32 };
{
let pf = block_base + (q4 + pf_quads) * 128 + h * 64;
if pf + 64 <= blocked_codes.len() {
_mm_prefetch(blocked_codes.as_ptr().add(pf) as *const i8, _MM_HINT_T0);
}
}
let c = _mm512_loadu_si512(
blocked_codes.as_ptr().add(block_base + q4 * 128 + h * 64) as *const __m512i,
);
let vlo = _mm512_shuffle_epi8(levels, _mm512_and_si512(c, m0f));
let vhi = _mm512_shuffle_epi8(levels, { let ($c, $mask) = (c, m0f); $hi });
for qi in 0..NQ {
let wp = pds[qi].weights.as_ptr().add(q4 * 8);
let wlo = _mm512_set1_epi32((wp as *const i32).read_unaligned());
let whi = _mm512_set1_epi32((wp.add(4) as *const i32).read_unaligned());
acc[qi][h] = _mm512_dpbusd_epi32(acc[qi][h], vlo, wlo);
acc[qi][h] = _mm512_dpbusd_epi32(acc[qi][h], vhi, whi);
}
}
}
}
for sub in 0..nb {
let base_vec = base_vec + sub * BLOCK;
let acc = &acc[sub];
let end = (base_vec + BLOCK).min(n_vectors);
for qi in 0..nqm {
let vs = _mm512_set1_ps(pds[qi].scale);
let vb = _mm512_set1_ps(pds[qi].bias);
let f0 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(acc[qi][0]), vs, vb);
let f1 = _mm512_fmadd_ps(_mm512_cvtepi32_ps(acc[qi][1]), vs, vb);
avx512_post_flush_heap_update(
f0,
f1,
base_vec,
end,
vec_scales.as_ptr().add(base_vec),
qi,
k,
mask,
heap_scores,
heap_indices,
heap_sizes,
heap_mins,
heap_min_idxs,
);
}
}
}
}
};
}
#[cfg(target_arch = "x86_64")]
define_permute_dot!(
search_multi_query_permute_dot,
["avx2", "fma", "avx512f", "avx512bw", "avx512vnni"],
|c, m0f| _mm512_and_si512(_mm512_srli_epi16(c, 4), m0f)
);
#[cfg(target_arch = "x86_64")]
define_permute_dot!(
search_multi_query_permute_dot_gfni,
["avx2", "fma", "avx512f", "avx512bw", "avx512vnni", "gfni"],
|c, _m0f| _mm512_gf2p8affine_epi64_epi8(c, _mm512_set1_epi64(0x1020408000000000u64 as i64), 0)
);
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn avx2_batch_flush_to_fa(
accus: [std::arch::x86_64::__m256i; 4],
v_scale: std::arch::x86_64::__m256,
fa: &mut [std::arch::x86_64::__m256; 4],
) {
use std::arch::x86_64::*;
let a0 = _mm256_sub_epi16(accus[0], _mm256_slli_epi16(accus[1], 8));
let a1 = accus[1];
let a2 = _mm256_sub_epi16(accus[2], _mm256_slli_epi16(accus[3], 8));
let a3 = accus[3];
let dis0 = _mm256_add_epi16(
_mm256_permute2x128_si256(a0, a1, 0x21),
_mm256_blend_epi32(a0, a1, 0xF0),
);
let dis1 = _mm256_add_epi16(
_mm256_permute2x128_si256(a2, a3, 0x21),
_mm256_blend_epi32(a2, a3, 0xF0),
);
let f0 = _mm256_cvtepi32_ps(_mm256_cvtepu16_epi32(_mm256_castsi256_si128(dis0)));
let f1 = _mm256_cvtepi32_ps(_mm256_cvtepu16_epi32(_mm256_extracti128_si256(dis0, 1)));
let f2 = _mm256_cvtepi32_ps(_mm256_cvtepu16_epi32(_mm256_castsi256_si128(dis1)));
let f3 = _mm256_cvtepi32_ps(_mm256_cvtepu16_epi32(_mm256_extracti128_si256(dis1, 1)));
fa[0] = _mm256_fmadd_ps(v_scale, f0, fa[0]);
fa[1] = _mm256_fmadd_ps(v_scale, f1, fa[1]);
fa[2] = _mm256_fmadd_ps(v_scale, f2, fa[2]);
fa[3] = _mm256_fmadd_ps(v_scale, f3, fa[3]);
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn avx2_post_flush_heap_update(
fa: &[std::arch::x86_64::__m256; 4],
base_vec: usize,
end: usize,
vec_scales_ptr: *const f32,
qi: usize,
k: usize,
mask: Option<&[u64]>,
heap_scores: &mut [Vec<f32>],
heap_indices: &mut [Vec<u64>],
heap_sizes: &mut [usize],
heap_mins: &mut [f32],
heap_min_idxs: &mut [usize],
) {
use std::arch::x86_64::*;
let end_lane = end - base_vec;
let (s0, s1, s2, s3) = if end_lane == BLOCK {
(
_mm256_mul_ps(fa[0], _mm256_loadu_ps(vec_scales_ptr)),
_mm256_mul_ps(fa[1], _mm256_loadu_ps(vec_scales_ptr.add(8))),
_mm256_mul_ps(fa[2], _mm256_loadu_ps(vec_scales_ptr.add(16))),
_mm256_mul_ps(fa[3], _mm256_loadu_ps(vec_scales_ptr.add(24))),
)
} else {
(fa[0], fa[1], fa[2], fa[3])
};
let hs = &mut heap_scores[qi];
let hi = &mut heap_indices[qi];
let sz = &mut heap_sizes[qi];
let hmin = &mut heap_mins[qi];
let hmi = &mut heap_min_idxs[qi];
if *sz >= k && end_lane == BLOCK {
let thr = _mm256_set1_ps(*hmin);
let m0 = _mm256_movemask_ps(_mm256_cmp_ps(s0, thr, _CMP_GT_OQ)) as u32;
let m1 = _mm256_movemask_ps(_mm256_cmp_ps(s1, thr, _CMP_GT_OQ)) as u32;
let m2 = _mm256_movemask_ps(_mm256_cmp_ps(s2, thr, _CMP_GT_OQ)) as u32;
let m3 = _mm256_movemask_ps(_mm256_cmp_ps(s3, thr, _CMP_GT_OQ)) as u32;
if (m0 | m1 | m2 | m3) == 0 {
return;
}
let mut block_out = [0.0f32; BLOCK];
let bp = block_out.as_mut_ptr();
if m0 != 0 { _mm256_storeu_ps(bp, s0); }
if m1 != 0 { _mm256_storeu_ps(bp.add(8), s1); }
if m2 != 0 { _mm256_storeu_ps(bp.add(16), s2); }
if m3 != 0 { _mm256_storeu_ps(bp.add(24), s3); }
for (chunk, &mask0) in [m0, m1, m2, m3].iter().enumerate() {
let mut m = mask0;
while m != 0 {
let bit = m.trailing_zeros() as usize;
m &= m - 1;
let lane = chunk * 8 + bit;
if let Some(am) = mask {
if !mask_allows(am, base_vec + lane) { continue; }
}
let score = block_out[lane];
if score > *hmin {
hs[*hmi] = score;
hi[*hmi] = (base_vec + lane) as u64;
let (m, mi) = rescan_min(hs, hi, k);
*hmin = m;
*hmi = mi;
}
}
}
return;
}
let mut block_out = [0.0f32; BLOCK];
let bp = block_out.as_mut_ptr();
_mm256_storeu_ps(bp, s0);
_mm256_storeu_ps(bp.add(8), s1);
_mm256_storeu_ps(bp.add(16), s2);
_mm256_storeu_ps(bp.add(24), s3);
if end_lane != BLOCK {
for lane in 0..end_lane {
block_out[lane] *= *vec_scales_ptr.add(lane);
}
for lane in end_lane..BLOCK {
block_out[lane] = f32::NEG_INFINITY;
}
}
if *sz < k {
for lane in 0..end_lane {
if let Some(am) = mask {
if !mask_allows(am, base_vec + lane) { continue; }
}
let score = block_out[lane];
if *sz < k {
hs[*sz] = score;
hi[*sz] = (base_vec + lane) as u64;
*sz += 1;
if *sz == k {
let (m, mi) = rescan_min(hs, hi, k);
*hmin = m;
*hmi = mi;
}
} else if score > *hmin {
hs[*hmi] = score;
hi[*hmi] = (base_vec + lane) as u64;
let (m, mi) = rescan_min(hs, hi, k);
*hmin = m;
*hmi = mi;
}
}
} else {
let v_hmin = _mm256_set1_ps(*hmin);
for chunk in 0..4 {
let chunk_start = chunk * 8;
if chunk_start >= end_lane { break; }
let scores_v = _mm256_loadu_ps(block_out.as_ptr().add(chunk_start));
let cmp = _mm256_cmp_ps(scores_v, v_hmin, _CMP_GT_OQ);
if _mm256_movemask_ps(cmp) == 0 { continue; }
let chunk_end = (chunk_start + 8).min(end_lane);
for lane in chunk_start..chunk_end {
if let Some(am) = mask {
if !mask_allows(am, base_vec + lane) { continue; }
}
let score = block_out[lane];
if score > *hmin {
hs[*hmi] = score;
hi[*hmi] = (base_vec + lane) as u64;
let (m, mi) = rescan_min(hs, hi, k);
*hmin = m;
*hmi = mi;
}
}
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f", enable = "avx512bw", enable = "avx2", enable = "fma")]
#[allow(clippy::too_many_arguments)]
unsafe fn avx512_post_flush_heap_update(
f0: std::arch::x86_64::__m512,
f1: std::arch::x86_64::__m512,
base_vec: usize,
end: usize,
vec_scales_ptr: *const f32,
qi: usize,
k: usize,
mask: Option<&[u64]>,
heap_scores: &mut [Vec<f32>],
heap_indices: &mut [Vec<u64>],
heap_sizes: &mut [usize],
heap_mins: &mut [f32],
heap_min_idxs: &mut [usize],
) {
use std::arch::x86_64::*;
let end_lane = end - base_vec;
let sz_now = heap_sizes[qi];
if sz_now >= k && end_lane == BLOCK {
let s0 = _mm512_mul_ps(f0, _mm512_loadu_ps(vec_scales_ptr));
let s1 = _mm512_mul_ps(f1, _mm512_loadu_ps(vec_scales_ptr.add(16)));
let thr = _mm512_set1_ps(heap_mins[qi]);
let m0 = _mm512_cmp_ps_mask(s0, thr, _CMP_GT_OQ) as u32;
let m1 = _mm512_cmp_ps_mask(s1, thr, _CMP_GT_OQ) as u32;
if (m0 | m1) == 0 {
return;
}
let mut block_out = [0.0f32; BLOCK];
let bp = block_out.as_mut_ptr();
if m0 != 0 {
_mm512_storeu_ps(bp, s0);
}
if m1 != 0 {
_mm512_storeu_ps(bp.add(16), s1);
}
let hs = &mut heap_scores[qi];
let hi = &mut heap_indices[qi];
let hmin = &mut heap_mins[qi];
let hmi = &mut heap_min_idxs[qi];
for (half, &mask0) in [m0, m1].iter().enumerate() {
let mut m = mask0;
while m != 0 {
let bit = m.trailing_zeros() as usize;
m &= m - 1;
let lane = half * 16 + bit;
if let Some(am) = mask {
if !mask_allows(am, base_vec + lane) {
continue;
}
}
let score = block_out[lane];
if score > *hmin {
hs[*hmi] = score;
hi[*hmi] = (base_vec + lane) as u64;
let (m2, mi) = rescan_min(hs, hi, k);
*hmin = m2;
*hmi = mi;
}
}
}
return;
}
let fa = [
_mm512_extractf32x8_ps(f0, 0),
_mm512_extractf32x8_ps(f0, 1),
_mm512_extractf32x8_ps(f1, 0),
_mm512_extractf32x8_ps(f1, 1),
];
avx2_post_flush_heap_update(
&fa, base_vec, end, vec_scales_ptr, qi, k, mask,
heap_scores, heap_indices, heap_sizes, heap_mins, heap_min_idxs,
);
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
unsafe fn scan_groups_neon(
codes_base: *const u8,
luts: [&[u8]; 4],
g0: usize,
g1: usize,
acc: &mut [[std::arch::aarch64::uint16x8_t; 4]; 4],
) {
use std::arch::aarch64::*;
let mask = vdupq_n_u8(0x0F);
for g in g0..g1 {
let cp = codes_base.add(g * BLOCK);
let c0 = vld1q_u8(cp);
let c1 = vld1q_u8(cp.add(16));
let lo0 = vandq_u8(c0, mask);
let lo1 = vandq_u8(c1, mask);
let hi0 = vshrq_n_u8(c0, 4);
let hi1 = vshrq_n_u8(c1, 4);
for q in 0..4 {
let lp = luts[q].as_ptr().add(g * 32);
let lut_hi = vld1q_u8(lp);
let lut_lo = vld1q_u8(lp.add(16));
let s0 = vaddq_u8(vqtbl1q_u8(lut_lo, lo0), vqtbl1q_u8(lut_hi, hi0));
let s1 = vaddq_u8(vqtbl1q_u8(lut_lo, lo1), vqtbl1q_u8(lut_hi, hi1));
acc[q][0] = vaddw_u8(acc[q][0], vget_low_u8(s0));
acc[q][1] = vaddw_u8(acc[q][1], vget_high_u8(s0));
acc[q][2] = vaddw_u8(acc[q][2], vget_low_u8(s1));
acc[q][3] = vaddw_u8(acc[q][3], vget_high_u8(s1));
}
}
}
#[cfg(target_arch = "aarch64")]
unsafe fn score_4query_block_neon(
blocked_codes: &[u8],
luts: [&[u8]; 4],
block_offset: usize,
n_byte_groups: usize,
scales: [f32; 4],
biases: [f32; 4],
vec_scales: &[f32],
base_vec: usize,
n_vectors: usize,
out: &mut [[f32; BLOCK]; 4],
) {
use std::arch::aarch64::*;
let n_batches = (n_byte_groups + FLUSH_EVERY - 1) / FLUSH_EVERY;
let mut fa: [[float32x4_t; 8]; 4] = [
[vdupq_n_f32(biases[0]); 8],
[vdupq_n_f32(biases[1]); 8],
[vdupq_n_f32(biases[2]); 8],
[vdupq_n_f32(biases[3]); 8],
];
let codes_base = blocked_codes.as_ptr().add(block_offset);
if n_batches == 1 {
let mut acc: [[uint16x8_t; 4]; 4] = [[vdupq_n_u16(0); 4]; 4];
scan_groups_neon(codes_base, luts, 0, n_byte_groups, &mut acc);
for q in 0..4 {
let v_scale = vdupq_n_f32(scales[q]);
let v_bias = vdupq_n_f32(biases[q]);
for i in 0..4 {
let lo = vcvtq_f32_u32(vmovl_u16(vget_low_u16(acc[q][i])));
let hi = vcvtq_f32_u32(vmovl_u16(vget_high_u16(acc[q][i])));
fa[q][i * 2] = vfmaq_f32(v_bias, v_scale, lo);
fa[q][i * 2 + 1] = vfmaq_f32(v_bias, v_scale, hi);
}
}
} else {
for batch in 0..n_batches {
let g_start = batch * FLUSH_EVERY;
let g_end = (g_start + FLUSH_EVERY).min(n_byte_groups);
let mut acc: [[uint16x8_t; 4]; 4] = [[vdupq_n_u16(0); 4]; 4];
scan_groups_neon(codes_base, luts, g_start, g_end, &mut acc);
for q in 0..4 {
let v_scale = vdupq_n_f32(scales[q]);
for i in 0..4 {
let lo = vcvtq_f32_u32(vmovl_u16(vget_low_u16(acc[q][i])));
let hi = vcvtq_f32_u32(vmovl_u16(vget_high_u16(acc[q][i])));
fa[q][i * 2] = vfmaq_f32(fa[q][i * 2], v_scale, lo);
fa[q][i * 2 + 1] = vfmaq_f32(fa[q][i * 2 + 1], v_scale, hi);
}
}
}
}
let end = (base_vec + BLOCK).min(n_vectors);
let vec_scales_ptr = vec_scales.as_ptr().add(base_vec);
for q in 0..4 {
let op = out[q].as_mut_ptr();
if end - base_vec == BLOCK {
for i in 0..8 {
let n = vld1q_f32(vec_scales_ptr.add(i * 4));
vst1q_f32(op.add(i * 4), vmulq_f32(fa[q][i], n));
}
} else {
let mut buf = [0.0f32; BLOCK];
for i in 0..8 {
vst1q_f32(buf.as_mut_ptr().add(i * 4), fa[q][i]);
}
for lane in 0..BLOCK {
*op.add(lane) = if lane < end - base_vec {
buf[lane] * *vec_scales_ptr.add(lane)
} else {
f32::NEG_INFINITY
};
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
unsafe fn sdot_lane<const IDX: i32>(
acc: std::arch::aarch64::int32x4_t,
a: std::arch::aarch64::int8x16_t,
b: std::arch::aarch64::int8x16_t,
) -> std::arch::aarch64::int32x4_t {
let mut o = acc;
std::arch::asm!(
".arch_extension dotprod",
"sdot {o:v}.4s, {a:v}.16b, {b:v}.4b[{idx}]",
o = inout(vreg) o,
a = in(vreg) a,
b = in(vreg) b,
idx = const IDX,
options(pure, nomem, nostack),
);
o
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
unsafe fn smmla(
acc: std::arch::aarch64::int32x4_t,
a: std::arch::aarch64::int8x16_t,
b: std::arch::aarch64::int8x16_t,
) -> std::arch::aarch64::int32x4_t {
let mut o = acc;
std::arch::asm!(
".arch_extension i8mm",
"smmla {o:v}.4s, {a:v}.16b, {b:v}.16b",
o = inout(vreg) o,
a = in(vreg) a,
b = in(vreg) b,
options(pure, nomem, nostack),
);
o
}
pub(crate) fn have_i8mm_layout() -> bool {
#[cfg(target_arch = "aarch64")]
{
have_i8mm()
}
#[cfg(not(target_arch = "aarch64"))]
{
false
}
}
#[cfg(target_arch = "aarch64")]
pub(crate) fn have_i8mm() -> bool {
static I8MM: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*I8MM.get_or_init(|| {
std::env::var_os("TURBOVEC_NO_I8MM").is_none()
&& std::arch::is_aarch64_feature_detected!("i8mm")
})
}
#[cfg(target_arch = "aarch64")]
fn build_smmla_a<const NQ: usize>(pds: &[&QueryPermuteDot; NQ], quads: usize) -> Vec<i8> {
let pairs = NQ / 2;
let mut a = vec![0i8; quads * pairs * 16];
for q4 in 0..quads {
for p in 0..pairs {
let dst = (q4 * pairs + p) * 16;
for (r, pd) in pds[2 * p..2 * p + 2].iter().enumerate() {
let w = &pd.weights[q4 * 8..q4 * 8 + 8];
for j in 0..4 {
a[dst + r * 8 + 2 * j] = w[4 + j];
a[dst + r * 8 + 2 * j + 1] = w[j];
}
}
}
}
a
}
#[cfg(target_arch = "aarch64")]
fn build_smmla_a_vm8<const NQ: usize>(pds: &[&QueryPermuteDot; NQ], octs: usize) -> Vec<i8> {
let pairs = NQ / 2;
let mut a = vec![0i8; octs * pairs * 32];
for q8 in 0..octs {
for p in 0..pairs {
let dst = (q8 * pairs + p) * 32;
for (r, pd) in pds[2 * p..2 * p + 2].iter().enumerate() {
for j in 0..8 {
let g = 8 * q8 + j;
let (q4, slot) = (g / 4, g % 4);
a[dst + r * 8 + j] = pd.weights[q4 * 8 + 4 + slot];
a[dst + 16 + r * 8 + j] = pd.weights[q4 * 8 + slot];
}
}
}
}
a
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments)]
unsafe fn score_block_vm8_single(
blocked_codes: &[u8],
pd: &QueryPermuteDot,
a_buf: &[i8],
block_offset: usize,
n_byte_groups: usize,
vec_scales: &[f32],
base_vec: usize,
n_vectors: usize,
out: &mut [[f32; BLOCK]; 1],
) {
use std::arch::aarch64::*;
let mask = vdupq_n_u8(0x0F);
let levels = vld1q_s8(pd.levels.as_ptr());
let octs = n_byte_groups / 8;
let codes_base = blocked_codes.as_ptr().add(block_offset);
let a_base = a_buf.as_ptr();
let mut acc = [vdupq_n_s32(0); 16];
for q8 in 0..octs {
let ap = a_base.add(q8 * 32);
let ae = vld1q_s8(ap);
let ao = vld1q_s8(ap.add(16));
for g in 0..4 {
let mut cs = [vdupq_n_u8(0); 4];
for (j, c) in cs.iter_mut().enumerate() {
*c = vld1q_u8(codes_base.add(q8 * 256 + (g * 4 + j) * 16));
}
for (j, &c) in cs.iter().enumerate() {
let r = g * 4 + j;
let bo = vqtbl1q_s8(levels, vandq_u8(c, mask));
let be = vqtbl1q_s8(levels, vshrq_n_u8(c, 4));
acc[r] = smmla(acc[r], ae, be);
acc[r] = smmla(acc[r], ao, bo);
}
}
}
let vs = vdupq_n_f32(pd.scale);
let vb = vdupq_n_f32(pd.bias);
let end = (base_vec + BLOCK).min(n_vectors);
let vec_scales_ptr = vec_scales.as_ptr().add(base_vec);
let mut raw = [0.0f32; BLOCK];
for i in 0..8 {
let t = vcombine_s32(vget_low_s32(acc[i * 2]), vget_low_s32(acc[i * 2 + 1]));
vst1q_f32(raw.as_mut_ptr().add(i * 4), vfmaq_f32(vb, vcvtq_f32_s32(t), vs));
}
let op = out[0].as_mut_ptr();
if end - base_vec == BLOCK {
for i in 0..8 {
let f = vld1q_f32(raw.as_ptr().add(i * 4));
let n = vld1q_f32(vec_scales_ptr.add(i * 4));
vst1q_f32(op.add(i * 4), vmulq_f32(f, n));
}
} else {
for lane in 0..BLOCK {
*op.add(lane) = if lane < end - base_vec {
raw[lane] * *vec_scales_ptr.add(lane)
} else {
f32::NEG_INFINITY
};
}
}
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments)]
unsafe fn score_block_smmla_vm8<const NQ: usize, const NP: usize>(
blocked_codes: &[u8],
pds: &[&QueryPermuteDot; NQ],
a_buf: &[i8],
block_offset: usize,
n_byte_groups: usize,
vec_scales: &[f32],
base_vec: usize,
n_vectors: usize,
out: &mut [[f32; BLOCK]; NQ],
) {
use std::arch::aarch64::*;
const { assert!(NQ == 2 * NP, "SMMLA tiles queries in pairs") };
let pairs = NP;
let mask = vdupq_n_u8(0x0F);
let levels = vld1q_s8(pds[0].levels.as_ptr());
let octs = n_byte_groups / 8;
let codes_base = blocked_codes.as_ptr().add(block_offset);
let a_base = a_buf.as_ptr();
let mut raw = [[0.0f32; BLOCK]; NQ];
for part in 0..8 {
let mut acc = [[vdupq_n_s32(0); 2]; NP];
for q8 in 0..octs {
let ap = a_base.add(q8 * pairs * 32);
let mut ae = [vdupq_n_s8(0); NP];
let mut ao = [vdupq_n_s8(0); NP];
for p in 0..pairs {
ae[p] = vld1q_s8(ap.add(p * 32));
ao[p] = vld1q_s8(ap.add(p * 32 + 16));
}
{
let pf = q8 * 256 + 32 * 256;
if block_offset + pf + 64 <= blocked_codes.len() {
std::arch::asm!(
"prfm pldl1keep, [{p}]",
p = in(reg) codes_base.add(pf),
options(nostack, readonly, preserves_flags),
);
}
}
for r in 0..2 {
let c = vld1q_u8(codes_base.add(q8 * 256 + (part * 2 + r) * 16));
let bo = vqtbl1q_s8(levels, vandq_u8(c, mask));
let be = vqtbl1q_s8(levels, vshrq_n_u8(c, 4));
for p in 0..pairs {
acc[p][r] = smmla(acc[p][r], ae[p], be);
acc[p][r] = smmla(acc[p][r], ao[p], bo);
}
}
}
for p in 0..pairs {
for (r, q) in [2 * p, 2 * p + 1].into_iter().enumerate() {
let vs = vdupq_n_f32(pds[q].scale);
let vb = vdupq_n_f32(pds[q].bias);
let (x, y) = (acc[p][0], acc[p][1]);
let t = if r == 0 {
vcombine_s32(vget_low_s32(x), vget_low_s32(y))
} else {
vcombine_s32(vget_high_s32(x), vget_high_s32(y))
};
let f = vfmaq_f32(vb, vcvtq_f32_s32(t), vs);
vst1q_f32(raw[q].as_mut_ptr().add(part * 4), f);
}
}
}
let end = (base_vec + BLOCK).min(n_vectors);
let vec_scales_ptr = vec_scales.as_ptr().add(base_vec);
for q in 0..NQ {
let op = out[q].as_mut_ptr();
if end - base_vec == BLOCK {
for i in 0..8 {
let f = vld1q_f32(raw[q].as_ptr().add(i * 4));
let n = vld1q_f32(vec_scales_ptr.add(i * 4));
vst1q_f32(op.add(i * 4), vmulq_f32(f, n));
}
} else {
for lane in 0..BLOCK {
*op.add(lane) = if lane < end - base_vec {
raw[q][lane] * *vec_scales_ptr.add(lane)
} else {
f32::NEG_INFINITY
};
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments)]
unsafe fn score_block_permute_smmla_neon<const NQ: usize, const NP: usize>(
blocked_codes: &[u8],
pds: &[&QueryPermuteDot; NQ],
a_buf: &[i8],
block_offset: usize,
n_byte_groups: usize,
vec_scales: &[f32],
base_vec: usize,
n_vectors: usize,
out: &mut [[f32; BLOCK]; NQ],
) {
use std::arch::aarch64::*;
const { assert!(NQ == 2 * NP, "SMMLA tiles queries in pairs") };
let pairs = NP;
let mask = vdupq_n_u8(0x0F);
let levels = vld1q_s8(pds[0].levels.as_ptr());
let quads = n_byte_groups / 4;
let codes_base = blocked_codes.as_ptr().add(block_offset);
let a_base = a_buf.as_ptr();
let mut raw = [[0.0f32; BLOCK]; NQ];
for part in 0..4 {
let mut acc = [[vdupq_n_s32(0); 4]; NP];
for q4 in 0..quads {
let ap = a_base.add(q4 * pairs * 16);
let mut a = [vdupq_n_s8(0); NP];
for (p, aq) in a.iter_mut().enumerate() {
*aq = vld1q_s8(ap.add(p * 16));
}
for i in 0..2 {
let c = vld1q_u8(codes_base.add(q4 * 128 + (part * 2 + i) * 16));
let vlo = vqtbl1q_s8(levels, vandq_u8(c, mask));
let vhi = vqtbl1q_s8(levels, vshrq_n_u8(c, 4));
let b0 = vzip1q_s8(vhi, vlo);
let b1 = vzip2q_s8(vhi, vlo);
for p in 0..pairs {
acc[p][i * 2] = smmla(acc[p][i * 2], a[p], b0);
acc[p][i * 2 + 1] = smmla(acc[p][i * 2 + 1], a[p], b1);
}
}
}
for p in 0..pairs {
for (r, q) in [2 * p, 2 * p + 1].into_iter().enumerate() {
let vs = vdupq_n_f32(pds[q].scale);
let vb = vdupq_n_f32(pds[q].bias);
for h in 0..2 {
let (x, y) = (acc[p][h * 2], acc[p][h * 2 + 1]);
let t = if r == 0 {
vcombine_s32(vget_low_s32(x), vget_low_s32(y))
} else {
vcombine_s32(vget_high_s32(x), vget_high_s32(y))
};
let f = vfmaq_f32(vb, vcvtq_f32_s32(t), vs);
vst1q_f32(raw[q].as_mut_ptr().add(part * 8 + h * 4), f);
}
}
}
}
let end = (base_vec + BLOCK).min(n_vectors);
let vec_scales_ptr = vec_scales.as_ptr().add(base_vec);
for q in 0..NQ {
let op = out[q].as_mut_ptr();
if end - base_vec == BLOCK {
for i in 0..8 {
let f = vld1q_f32(raw[q].as_ptr().add(i * 4));
let n = vld1q_f32(vec_scales_ptr.add(i * 4));
vst1q_f32(op.add(i * 4), vmulq_f32(f, n));
}
} else {
for lane in 0..BLOCK {
*op.add(lane) = if lane < end - base_vec {
raw[q][lane] * *vec_scales_ptr.add(lane)
} else {
f32::NEG_INFINITY
};
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "dotprod")]
#[allow(clippy::too_many_arguments)]
unsafe fn score_block_permute_dot_neon<const NQ: usize>(
blocked_codes: &[u8],
pds: &[&QueryPermuteDot; NQ],
block_offset: usize,
n_byte_groups: usize,
vec_scales: &[f32],
base_vec: usize,
n_vectors: usize,
out: &mut [[f32; BLOCK]; NQ],
) {
use std::arch::aarch64::*;
let mask = vdupq_n_u8(0x0F);
let levels = vld1q_s8(pds[0].levels.as_ptr());
let quads = n_byte_groups / 4;
let codes_base = blocked_codes.as_ptr().add(block_offset);
let mut raw = [[0.0f32; BLOCK]; NQ];
for part in 0..4 {
let mut acc = [[vdupq_n_s32(0); 2]; NQ];
let mut q4 = 0usize;
while q4 < quads {
let paired = q4 + 1 < quads;
let mut w = [vdupq_n_s8(0); NQ];
for (wq, pd) in w.iter_mut().zip(pds.iter()) {
let wp = pd.weights.as_ptr().add(q4 * 8);
*wq = if paired {
vld1q_s8(wp)
} else {
vcombine_s8(vld1_s8(wp), vdup_n_s8(0))
};
}
for i in 0..2 {
let c = vld1q_u8(codes_base.add(q4 * 128 + (part * 2 + i) * 16));
let vlo = vqtbl1q_s8(levels, vandq_u8(c, mask));
let vhi = vqtbl1q_s8(levels, vshrq_n_u8(c, 4));
for q in 0..NQ {
acc[q][i] = sdot_lane::<0>(acc[q][i], vlo, w[q]);
acc[q][i] = sdot_lane::<1>(acc[q][i], vhi, w[q]);
}
}
if paired {
for i in 0..2 {
let c = vld1q_u8(codes_base.add((q4 + 1) * 128 + (part * 2 + i) * 16));
let vlo = vqtbl1q_s8(levels, vandq_u8(c, mask));
let vhi = vqtbl1q_s8(levels, vshrq_n_u8(c, 4));
for q in 0..NQ {
acc[q][i] = sdot_lane::<2>(acc[q][i], vlo, w[q]);
acc[q][i] = sdot_lane::<3>(acc[q][i], vhi, w[q]);
}
}
}
q4 += 2;
}
for q in 0..NQ {
let vs = vdupq_n_f32(pds[q].scale);
let vb = vdupq_n_f32(pds[q].bias);
for i in 0..2 {
let f = vfmaq_f32(vb, vcvtq_f32_s32(acc[q][i]), vs);
vst1q_f32(raw[q].as_mut_ptr().add((part * 2 + i) * 4), f);
}
}
}
let end = (base_vec + BLOCK).min(n_vectors);
let vec_scales_ptr = vec_scales.as_ptr().add(base_vec);
for q in 0..NQ {
let op = out[q].as_mut_ptr();
if end - base_vec == BLOCK {
for i in 0..8 {
let f = vld1q_f32(raw[q].as_ptr().add(i * 4));
let n = vld1q_f32(vec_scales_ptr.add(i * 4));
vst1q_f32(op.add(i * 4), vmulq_f32(f, n));
}
} else {
for lane in 0..BLOCK {
*op.add(lane) = if lane < end - base_vec {
raw[q][lane] * *vec_scales_ptr.add(lane)
} else {
f32::NEG_INFINITY
};
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
unsafe fn neon_block_topk_update(
block_scores: &[f32; BLOCK],
base_vec: usize,
end_lane: usize,
mask: Option<&[u64]>,
k: usize,
hs: &mut [f32],
hi: &mut [u64],
sz: &mut usize,
hmin: &mut f32,
hmi: &mut usize,
) {
use std::arch::aarch64::*;
if *sz >= k {
let p = block_scores.as_ptr();
let mut m = vld1q_f32(p);
for i in 1..8 {
m = vmaxq_f32(m, vld1q_f32(p.add(i * 4)));
}
if vmaxvq_f32(m) <= *hmin {
return;
}
}
for (lane, &s) in block_scores.iter().enumerate().take(end_lane) {
if let Some(am) = mask {
if !mask_allows(am, base_vec + lane) {
continue;
}
}
if *sz < k {
hs[*sz] = s;
hi[*sz] = (base_vec + lane) as u64;
*sz += 1;
if *sz == k {
let (m, mi) = rescan_min(hs, hi, k);
*hmin = m;
*hmi = mi;
}
} else if s > *hmin {
hs[*hmi] = s;
hi[*hmi] = (base_vec + lane) as u64;
let (m, mi) = rescan_min(hs, hi, k);
*hmin = m;
*hmi = mi;
}
}
}
pub(crate) struct QueryNeonLut {
pub(crate) uint8_luts: Vec<u8>, #[cfg(target_arch = "x86_64")]
pub(crate) split: Vec<u8>,
pub(crate) pd: Option<QueryPermuteDot>,
pub(crate) scale: f32,
pub(crate) bias: f32,
}
#[cfg(target_arch = "x86_64")]
pub(crate) fn split_lut_for_vnni(uint8_luts: &[u8], n_byte_groups: usize) -> Vec<u8> {
debug_assert_eq!(uint8_luts.len(), n_byte_groups * 32);
debug_assert_eq!(n_byte_groups % 4, 0);
let mut out = vec![0u8; n_byte_groups * 32];
for g0 in (0..n_byte_groups).step_by(4) {
let c = (g0 / 4) * 128;
for j in 0..4 {
let src = (g0 + j) * 32;
out[c + j * 16..c + j * 16 + 16].copy_from_slice(&uint8_luts[src + 16..src + 32]);
out[c + 64 + j * 16..c + 64 + j * 16 + 16]
.copy_from_slice(&uint8_luts[src..src + 16]);
}
}
out
}
pub(crate) struct QueryPermuteDot {
pub(crate) levels: [i8; 16],
pub(crate) weights: Vec<i8>,
#[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
pub(crate) zero: i32,
pub(crate) scale: f32,
pub(crate) bias: f32,
}
fn build_permute_dot(q_rot_row: &[f32], centroids: &[f32], dim: usize) -> QueryPermuteDot {
debug_assert_eq!(centroids.len(), 16);
let n_byte_groups = dim / 2;
let cmax = centroids.iter().fold(0.0f32, |m, &c| m.max(c.abs()));
let cs = if cmax > 0.0 { cmax / 127.0 } else { 1.0 };
let mut levels = [0i8; 16];
for (l, &c) in levels.iter_mut().zip(centroids.iter()) {
*l = (c / cs).round().clamp(-127.0, 127.0) as i8;
}
let qmax = q_rot_row.iter().fold(0.0f32, |m, &v| m.max(v.abs()));
let qs = qmax / 127.0;
let (qs, inv_qs) = if qs >= f32::MIN_POSITIVE { (qs, 1.0 / qs) } else { (1.0, 1.0) };
let mut weights = vec![0i8; n_byte_groups * 2];
let mut wsum: i32 = 0;
for g in 0..n_byte_groups {
let (q4, j) = (g / 4, g % 4);
let lo = (q_rot_row[2 * g + 1] * inv_qs).round().clamp(-127.0, 127.0) as i8;
let hi = (q_rot_row[2 * g] * inv_qs).round().clamp(-127.0, 127.0) as i8;
weights[q4 * 8 + j] = lo;
weights[q4 * 8 + 4 + j] = hi;
wsum += lo as i32 + hi as i32;
}
QueryPermuteDot {
levels,
weights,
zero: -128 * wsum,
scale: cs * qs,
bias: 0.0,
}
}
pub(crate) fn build_query_neon_lut_from_slice(
q_rot_row: &[f32],
centroids: &[f32],
bits: usize,
dim: usize,
) -> QueryNeonLut {
let codes_per_byte = 8 / bits;
let codes_per_nibble = codes_per_byte / 2;
let n_byte_groups = dim / codes_per_byte;
let code_mask = (1u16 << bits) - 1;
let n_subs = n_byte_groups * 2;
let mut uint8_luts = vec![0u8; n_byte_groups * 32];
let mut float_vals = vec![0.0f32; n_byte_groups * 32];
let mut mins = vec![0.0f32; n_subs];
let mut max_span = 0.0f32;
let mut sum_spans = 0.0f32;
let mut bias = 0.0f32;
for g in 0..n_byte_groups {
let dim_start = g * codes_per_byte;
let mut lo_min = f32::MAX;
let mut lo_max = f32::MIN;
for nibble_val in 0u16..16 {
let mut s = 0.0f32;
for c in 0..codes_per_nibble {
let shift = (codes_per_nibble - 1 - c) * bits;
let code = (nibble_val >> shift) & code_mask;
s += q_rot_row[dim_start + c] * centroids[code as usize];
}
float_vals[g * 32 + nibble_val as usize] = s;
if s < lo_min { lo_min = s; }
if s > lo_max { lo_max = s; }
}
let mut hi_min = f32::MAX;
let mut hi_max = f32::MIN;
for nibble_val in 0u16..16 {
let mut s = 0.0f32;
for c in 0..codes_per_nibble {
let shift = (codes_per_nibble - 1 - c) * bits;
let code = (nibble_val >> shift) & code_mask;
s += q_rot_row[dim_start + codes_per_nibble + c] * centroids[code as usize];
}
float_vals[g * 32 + 16 + nibble_val as usize] = s;
if s < hi_min { hi_min = s; }
if s > hi_max { hi_max = s; }
}
mins[g * 2] = lo_min;
mins[g * 2 + 1] = hi_min;
bias += lo_min + hi_min;
let lo_span = lo_max - lo_min;
let hi_span = hi_max - hi_min;
if lo_span > max_span { max_span = lo_span; }
if hi_span > max_span { max_span = hi_span; }
sum_spans += lo_span + hi_span;
}
let _ = sum_spans; let max_lut: f32 = 127.0;
let scale = if max_span > 0.0 { max_span / max_lut } else { 1.0 };
let (scale, inv_scale) = if scale >= f32::MIN_POSITIVE {
(scale, 1.0 / scale)
} else {
(1.0, 1.0)
};
for g in 0..n_byte_groups {
let lo_min = mins[g * 2];
let hi_min = mins[g * 2 + 1];
for i in 0..16 {
let j_lo = g * 32 + i;
let j_hi = g * 32 + 16 + i;
uint8_luts[j_lo] =
((float_vals[j_lo] - lo_min) * inv_scale).round().clamp(0.0, max_lut) as u8;
uint8_luts[j_hi] =
((float_vals[j_hi] - hi_min) * inv_scale).round().clamp(0.0, max_lut) as u8;
}
}
let vm = crate::pack::vector_major_for(bits, n_byte_groups);
let pd = if vm && bits == 4 {
Some(build_permute_dot(q_rot_row, centroids, dim))
} else {
None
};
QueryNeonLut {
#[cfg(target_arch = "x86_64")]
split: if vm && pd.is_none() {
split_lut_for_vnni(&uint8_luts, n_byte_groups)
} else {
Vec::new()
},
pd,
uint8_luts,
scale,
bias,
}
}
#[inline(always)]
pub(crate) fn mask_allows(mask: &[u64], slot: usize) -> bool {
(mask[slot >> 6] >> (slot & 63)) & 1 != 0
}
#[inline(always)]
pub(crate) fn block_has_allowed(mask: Option<&[u64]>, base_vec: usize) -> bool {
match mask {
None => true,
Some(m) => {
let word = m[base_vec >> 6];
let bit_offset = base_vec & 63;
let allowed = ((word >> bit_offset) & 0xFFFF_FFFF) != 0;
#[cfg(feature = "mask-skip-counter")]
if !allowed {
BLOCKS_SKIPPED_BY_MASK.fetch_add(1, Ordering::Relaxed);
}
allowed
}
}
}
#[inline]
pub(crate) fn block_range_stride(n_blocks: usize, n_threads: usize) -> usize {
(n_blocks.div_ceil(n_threads)).max(64).next_multiple_of(2)
}
#[cfg(target_arch = "x86_64")]
#[inline(always)]
pub(crate) fn block_pair_has_allowed(mask: Option<&[u64]>, base_vec_pair: usize) -> bool {
match mask {
None => true,
Some(m) => {
let allowed = m[base_vec_pair >> 6] != 0;
#[cfg(feature = "mask-skip-counter")]
if !allowed {
BLOCKS_SKIPPED_BY_MASK.fetch_add(2, Ordering::Relaxed);
}
allowed
}
}
}
#[cfg(not(target_arch = "aarch64"))]
#[allow(clippy::too_many_arguments)]
fn score_query_into_heap(
qlut_uint8: &[u8],
qlut_scale: f32,
qlut_bias: f32,
blocked_codes: &[u8],
vec_scales: &[f32],
bits: usize,
n_byte_groups: usize,
n_vectors: usize,
n_blocks: usize,
mask: Option<&[u64]>,
k: usize,
heap_s: &mut [f32],
heap_i: &mut [u64],
heap_sz: &mut usize,
heap_min: &mut f32,
heap_mi: &mut usize,
) {
for b in 0..n_blocks {
let base_vec = b * BLOCK;
if !block_has_allowed(mask, base_vec) {
continue;
}
for lane in 0..BLOCK {
let vi = base_vec + lane;
if vi >= n_vectors {
break;
}
if let Some(m) = mask {
if !mask_allows(m, vi) {
continue;
}
}
let mut score = qlut_bias;
for g in 0..n_byte_groups {
let byte_val =
crate::pack::read_code(blocked_codes, bits, n_byte_groups, b, g, lane) as usize;
let hi = byte_val >> 4;
let lo = byte_val & 0x0F;
score += qlut_scale * qlut_uint8[g * 32 + hi] as f32;
score += qlut_scale * qlut_uint8[g * 32 + 16 + lo] as f32;
}
score *= vec_scales[vi];
if *heap_sz < k {
heap_s[*heap_sz] = score;
heap_i[*heap_sz] = vi as u64;
*heap_sz += 1;
if *heap_sz == k {
let (m, mi) = rescan_min(heap_s, heap_i, k);
*heap_min = m;
*heap_mi = mi;
}
} else if score > *heap_min {
heap_s[*heap_mi] = score;
heap_i[*heap_mi] = vi as u64;
let (m, mi) = rescan_min(heap_s, heap_i, k);
*heap_min = m;
*heap_mi = mi;
}
}
}
}
fn calibrate_queries(
q_rot: &[f32],
tqplus_shift: &[f32],
tqplus_scale: &[f32],
nq: usize,
dim: usize,
) -> (Vec<f32>, Vec<f32>) {
if tqplus_shift.is_empty() {
debug_assert!(tqplus_scale.is_empty());
return (q_rot.to_vec(), vec![0.0f32; nq]);
}
debug_assert_eq!(tqplus_shift.len(), dim);
debug_assert_eq!(tqplus_scale.len(), dim);
let mut q_calib = vec![0.0f32; nq * dim];
let mut bias_corrs = vec![0.0f32; nq];
q_calib
.par_chunks_mut(dim)
.zip(bias_corrs.par_iter_mut())
.enumerate()
.for_each(|(qi, (calib_row, bias))| {
let q_row = &q_rot[qi * dim..(qi + 1) * dim];
let mut bc = 0.0f64;
for d in 0..dim {
calib_row[d] = q_row[d] / tqplus_scale[d];
bc -= (q_row[d] as f64) * (tqplus_shift[d] as f64);
}
*bias = bc as f32;
});
(q_calib, bias_corrs)
}
pub(crate) fn search(
queries: &[f32], nq: usize,
rotation: &Rotation,
blocked_codes: &[u8],
centroids: &[f32],
vec_scales: &[f32],
tqplus_shift: &[f32], tqplus_scale: &[f32], bits: usize,
dim: usize,
n_vectors: usize,
n_blocks: usize,
k: usize,
mask: Option<&[u64]>,
) -> (Vec<f32>, Vec<i64>) {
let n_allowed = match mask {
Some(m) => m.iter().map(|w| w.count_ones() as usize).sum::<usize>(),
None => n_vectors,
};
let k = k.min(n_allowed);
if k == 0 {
return (Vec::new(), Vec::new());
}
let n_byte_groups = dim / (8 / bits);
let mut q_rot = queries.to_vec();
q_rot
.par_chunks_mut(dim)
.for_each_init(|| vec![0.0f32; dim], |scratch, row| {
rotation.apply_with_scratch(row, scratch)
});
let (q_for_lut, bias_corrs) =
calibrate_queries(&q_rot, tqplus_shift, tqplus_scale, nq, dim);
let query_luts: Vec<QueryNeonLut> = (0..nq)
.into_par_iter()
.map(|qi| {
let row = &q_for_lut[qi * dim..(qi + 1) * dim];
let mut lut = build_query_neon_lut_from_slice(row, centroids, bits, dim);
lut.bias += bias_corrs[qi];
if let Some(pd) = lut.pd.as_mut() {
pd.bias += bias_corrs[qi];
}
lut
})
.collect();
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments)]
fn scan_range_neon<const MASKED: bool>(
codes: &[u8],
lut: &QueryNeonLut,
n_byte_groups: usize,
scales_slice: &[f32],
block_bytes: usize,
range_blocks: usize,
range_vecs: usize,
k: usize,
mask: Option<&[u64]>,
) -> Vec<(f32, u64)> {
let mut heap: Vec<(f32, u64)> = Vec::with_capacity(k);
let mut heap_min = f32::NEG_INFINITY;
let mut heap_mi = 0usize;
let mut out = [[0.0f32; BLOCK]; 1];
let vm8_single: Option<Vec<i8>> = lut.pd.as_ref().and_then(|pd| {
crate::pack::vm8_for(4, n_byte_groups)
.then(|| build_smmla_a_vm8::<2>(&[pd, pd], n_byte_groups / 8))
});
for b in 0..range_blocks {
let base = b * BLOCK;
let end = (base + BLOCK).min(range_vecs);
if MASKED && !block_has_allowed(mask, base) {
continue;
}
unsafe {
if let Some(pd) = lut.pd.as_ref() {
if let Some(a) = vm8_single.as_deref() {
score_block_vm8_single(
codes, pd, a, b * block_bytes, n_byte_groups,
scales_slice, base, range_vecs, &mut out,
);
} else {
score_block_permute_dot_neon::<1>(
codes, &[pd], b * block_bytes, n_byte_groups,
scales_slice, base, range_vecs, &mut out,
);
}
} else {
score_4bit_block_neon(
codes, &lut.uint8_luts, b * block_bytes, n_byte_groups,
lut.scale, lut.bias, scales_slice, base, range_vecs, &mut out[0],
);
}
}
if heap.len() >= k {
let block_max = unsafe {
use std::arch::aarch64::*;
let p = out[0].as_ptr();
let m0 = vmaxq_f32(vld1q_f32(p), vld1q_f32(p.add(4)));
let m1 = vmaxq_f32(vld1q_f32(p.add(8)), vld1q_f32(p.add(12)));
let m2 = vmaxq_f32(vld1q_f32(p.add(16)), vld1q_f32(p.add(20)));
let m3 = vmaxq_f32(vld1q_f32(p.add(24)), vld1q_f32(p.add(28)));
vmaxvq_f32(vmaxq_f32(vmaxq_f32(m0, m1), vmaxq_f32(m2, m3)))
};
if block_max <= heap_min {
continue;
}
}
for (lane, &s) in out[0][..end - base].iter().enumerate() {
if MASKED && !mask_allows(mask.expect("MASKED implies a mask"), base + lane) {
continue;
}
if heap.len() < k {
heap.push((s, (base + lane) as u64));
if heap.len() == k {
heap_mi = 0;
for (h, &(hs, hix)) in heap.iter().enumerate().skip(1) {
if hs < heap[heap_mi].0
|| (hs == heap[heap_mi].0 && hix > heap[heap_mi].1)
{
heap_mi = h;
}
}
heap_min = heap[heap_mi].0;
}
} else if s > heap_min {
heap[heap_mi] = (s, (base + lane) as u64);
heap_mi = 0;
for (h, &(hs, hix)) in heap.iter().enumerate().skip(1) {
if hs < heap[heap_mi].0 || (hs == heap[heap_mi].0 && hix > heap[heap_mi].1)
{
heap_mi = h;
}
}
heap_min = heap[heap_mi].0;
}
}
}
heap
}
#[cfg(target_arch = "aarch64")]
#[allow(clippy::too_many_arguments)]
fn search_single_query_block_parallel_neon(
blocked_codes: &[u8],
lut: &QueryNeonLut,
n_byte_groups: usize,
vec_scales: &[f32],
n_vectors: usize,
n_blocks: usize,
k: usize,
mask: Option<&[u64]>,
) -> (Vec<f32>, Vec<i64>) {
let n_threads = rayon::current_num_threads().max(1);
let blocks_per_range = block_range_stride(n_blocks, n_threads);
let ranges: Vec<usize> = (0..n_blocks).step_by(blocks_per_range).collect();
let block_bytes = n_byte_groups * BLOCK;
let mut candidates: Vec<(f32, u64)> = ranges
.into_par_iter()
.flat_map(|block_start| {
let range_blocks = blocks_per_range.min(n_blocks - block_start);
let vec_start = block_start * BLOCK;
let range_vecs = (range_blocks * BLOCK).min(n_vectors - vec_start);
let codes = &blocked_codes
[block_start * block_bytes..(block_start + range_blocks) * block_bytes];
let scales_slice = &vec_scales[vec_start..vec_start + range_vecs];
let mask_slice = mask.map(|m| &m[vec_start / 64..]);
let heap = if mask_slice.is_some() {
scan_range_neon::<true>(
codes, lut, n_byte_groups, scales_slice, block_bytes,
range_blocks, range_vecs, k, mask_slice,
)
} else {
scan_range_neon::<false>(
codes, lut, n_byte_groups, scales_slice, block_bytes,
range_blocks, range_vecs, k, None,
)
};
heap.into_iter()
.map(|(s, i)| (s, i + vec_start as u64))
.collect::<Vec<_>>()
})
.collect();
candidates.sort_unstable_by(|a, b| {
b.0.partial_cmp(&a.0)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.1.cmp(&b.1))
});
candidates.truncate(k);
(
candidates.iter().map(|p| p.0).collect(),
candidates.iter().map(|p| p.1 as i64).collect(),
)
}
#[cfg(target_arch = "aarch64")]
let results = {
if nq == 1 && n_blocks >= SINGLE_QUERY_PARALLEL_MIN_BLOCKS {
vec![search_single_query_block_parallel_neon(
blocked_codes, &query_luts[0], n_byte_groups, vec_scales,
n_vectors, n_blocks, k, mask,
)]
} else {
let pd_batched = query_luts.first().is_some_and(|l| l.pd.is_some());
let qbs: usize = if !pd_batched {
4
} else if nq >= 12 {
12
} else if nq >= 8 {
8
} else {
4
};
const QBS_MAX: usize = 12;
const QBS_LUT: usize = 4;
let n_quads = nq.div_ceil(qbs).max(1);
let n_threads = rayon::current_num_threads().max(1);
let n_ranges = n_block_ranges(
nq, n_quads, n_blocks, n_vectors, k, n_threads,
TILES_PER_THREAD_NEON,
if bits == 2 { MIN_TILE_BLOCKS_NEON * 2 } else { MIN_TILE_BLOCKS_NEON },
false,
);
let n_ranges = smooth_tile_count(n_ranges, n_quads, n_threads);
let blocks_per_range = n_blocks.div_ceil(n_ranges).max(1);
let tiles: Vec<(usize, usize)> = (0..n_blocks.max(1))
.step_by(blocks_per_range)
.flat_map(|b| (0..nq).step_by(qbs).map(move |q| (q, b)))
.collect();
let tile_results: Vec<(usize, Vec<Vec<(f32, u64)>>)> = tiles
.into_par_iter()
.map(|(qi_start, block_start)| {
let block_end = (block_start + blocks_per_range).min(n_blocks);
let qi_end = (qi_start + qbs).min(nq);
let batch_size = qi_end - qi_start;
let mut heap_s = vec![vec![f32::NEG_INFINITY; k]; batch_size];
let mut heap_i = vec![vec![0u64; k]; batch_size];
let mut heap_sz = [0usize; QBS_MAX];
let mut heap_min = [f32::NEG_INFINITY; QBS_MAX];
let mut heap_mi = [0usize; QBS_MAX];
macro_rules! pd_scan {
($n:literal, $np:literal) => {{
let pds: [&QueryPermuteDot; $n] = std::array::from_fn(|i| {
query_luts[qi_start + i]
.pd
.as_ref()
.expect("pd built for every query")
});
let vm8 = crate::pack::vm8_for(bits, n_byte_groups);
let a_buf = if vm8 {
build_smmla_a_vm8::<$n>(&pds, n_byte_groups / 8)
} else if have_i8mm() {
build_smmla_a::<$n>(&pds, n_byte_groups / 4)
} else {
Vec::new()
};
let mut block_out = [[0.0f32; BLOCK]; $n];
for block_idx in block_start..block_end {
let base_vec = block_idx * BLOCK;
if !block_has_allowed(mask, base_vec) {
continue;
}
let block_offset = block_idx * n_byte_groups * BLOCK;
let end_lane = (base_vec + BLOCK).min(n_vectors) - base_vec;
unsafe {
if vm8 {
score_block_smmla_vm8::<$n, $np>(
blocked_codes, &pds, &a_buf, block_offset,
n_byte_groups, vec_scales, base_vec, n_vectors,
&mut block_out,
);
} else if a_buf.is_empty() {
score_block_permute_dot_neon::<$n>(
blocked_codes, &pds, block_offset, n_byte_groups,
vec_scales, base_vec, n_vectors, &mut block_out,
);
} else {
score_block_permute_smmla_neon::<$n, $np>(
blocked_codes, &pds, &a_buf, block_offset,
n_byte_groups, vec_scales, base_vec, n_vectors,
&mut block_out,
);
}
for q in 0..$n {
neon_block_topk_update(
&block_out[q], base_vec, end_lane, mask, k,
&mut heap_s[q], &mut heap_i[q], &mut heap_sz[q],
&mut heap_min[q], &mut heap_mi[q],
);
}
}
}
}};
}
if pd_batched && batch_size == 12 {
pd_scan!(12, 6)
} else if pd_batched && batch_size == 8 {
pd_scan!(8, 4)
} else if pd_batched && batch_size == 4 {
pd_scan!(4, 2)
} else if !pd_batched && batch_size == QBS_LUT {
let lut_refs: [&[u8]; QBS_LUT] = [
&query_luts[qi_start].uint8_luts,
&query_luts[qi_start + 1].uint8_luts,
&query_luts[qi_start + 2].uint8_luts,
&query_luts[qi_start + 3].uint8_luts,
];
let scales: [f32; QBS_LUT] = [
query_luts[qi_start].scale,
query_luts[qi_start + 1].scale,
query_luts[qi_start + 2].scale,
query_luts[qi_start + 3].scale,
];
let biases: [f32; QBS_LUT] = [
query_luts[qi_start].bias,
query_luts[qi_start + 1].bias,
query_luts[qi_start + 2].bias,
query_luts[qi_start + 3].bias,
];
let mut block_out = [[0.0f32; BLOCK]; QBS_LUT];
for block_idx in block_start..block_end {
let base_vec = block_idx * BLOCK;
if !block_has_allowed(mask, base_vec) {
continue;
}
let block_offset = block_idx * n_byte_groups * BLOCK;
let end_lane = (base_vec + BLOCK).min(n_vectors) - base_vec;
unsafe {
score_4query_block_neon(
blocked_codes, lut_refs, block_offset, n_byte_groups,
scales, biases, vec_scales, base_vec, n_vectors,
&mut block_out,
);
for q in 0..QBS_LUT {
neon_block_topk_update(
&block_out[q], base_vec, end_lane, mask, k,
&mut heap_s[q], &mut heap_i[q], &mut heap_sz[q],
&mut heap_min[q], &mut heap_mi[q],
);
}
}
}
} else {
for qi_off in 0..batch_size {
let qi = qi_start + qi_off;
let qlut = &query_luts[qi];
let vm8_single: Option<Vec<i8>> = qlut.pd.as_ref().and_then(|pd| {
crate::pack::vm8_for(bits, n_byte_groups)
.then(|| build_smmla_a_vm8::<2>(&[pd, pd], n_byte_groups / 8))
});
for block_idx in block_start..block_end {
let base_vec = block_idx * BLOCK;
if !block_has_allowed(mask, base_vec) {
continue;
}
let block_offset = block_idx * n_byte_groups * BLOCK;
let end_lane = (base_vec + BLOCK).min(n_vectors) - base_vec;
let mut block_out = [[0.0f32; BLOCK]; 1];
unsafe {
if let Some(pd) = qlut.pd.as_ref() {
if let Some(a) = vm8_single.as_deref() {
score_block_vm8_single(
blocked_codes, pd, a, block_offset, n_byte_groups,
vec_scales, base_vec, n_vectors, &mut block_out,
);
} else {
score_block_permute_dot_neon::<1>(
blocked_codes, &[pd], block_offset, n_byte_groups,
vec_scales, base_vec, n_vectors, &mut block_out,
);
}
} else {
score_4bit_block_neon(
blocked_codes, &qlut.uint8_luts, block_offset, n_byte_groups,
qlut.scale, qlut.bias, vec_scales, base_vec, n_vectors,
&mut block_out[0],
);
}
neon_block_topk_update(
&block_out[0], base_vec, end_lane, mask, k,
&mut heap_s[qi_off], &mut heap_i[qi_off],
&mut heap_sz[qi_off], &mut heap_min[qi_off],
&mut heap_mi[qi_off],
);
}
}
}
}
let cands: Vec<Vec<(f32, u64)>> = (0..batch_size)
.map(|qi_off| {
let sz = heap_sz[qi_off];
heap_s[qi_off][..sz]
.iter()
.zip(heap_i[qi_off][..sz].iter())
.map(|(&s, &i)| (s, i))
.collect()
})
.collect();
(qi_start, cands)
})
.collect();
let mut merged: Vec<Vec<(f32, u64)>> = vec![Vec::new(); nq];
for (qi_start, cands) in tile_results {
for (off, c) in cands.into_iter().enumerate() {
merged[qi_start + off].extend(c);
}
}
merged
.into_iter()
.map(|mut pairs| {
pairs.sort_unstable_by(|a, b| {
b.0.partial_cmp(&a.0)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.1.cmp(&b.1))
});
pairs.truncate(k);
let s: Vec<f32> = pairs.iter().map(|p| p.0).collect();
let i: Vec<i64> = pairs.iter().map(|p| p.1 as i64).collect();
(s, i)
})
.collect::<Vec<_>>()
}
};
#[cfg(target_arch = "x86_64")]
#[allow(clippy::too_many_arguments)]
fn search_single_query_block_parallel(
blocked_codes: &[u8],
lut: &QueryNeonLut,
n_byte_groups: usize,
vec_scales: &[f32],
n_vectors: usize,
n_blocks: usize,
k: usize,
use_avx512: bool,
mask: Option<&[u64]>,
) -> (Vec<f32>, Vec<i64>) {
let n_threads = rayon::current_num_threads().max(1);
let blocks_per_range = block_range_stride(n_blocks, n_threads);
let ranges: Vec<usize> = (0..n_blocks).step_by(blocks_per_range).collect();
let block_bytes = n_byte_groups * BLOCK;
let mut candidates: Vec<(f32, u64)> = ranges
.into_par_iter()
.flat_map(|block_start| {
let range_blocks = blocks_per_range.min(n_blocks - block_start);
let vec_start = block_start * BLOCK;
let range_vecs = (range_blocks * BLOCK).min(n_vectors - vec_start);
let codes =
&blocked_codes[block_start * block_bytes..(block_start + range_blocks) * block_bytes];
let scales_slice = &vec_scales[vec_start..vec_start + range_vecs];
let mask_slice = mask.map(|m| &m[vec_start / 64..]);
let lut_refs = [lut.uint8_luts.as_slice(); 4];
let scale_vals = [lut.scale; 4];
let bias_vals = [lut.bias; 4];
let mut heap_scores = vec![vec![f32::NEG_INFINITY; k]];
let mut heap_indices = vec![vec![0u64; k]];
let mut heap_sizes = vec![0usize];
let mut heap_mins = vec![f32::NEG_INFINITY];
let mut heap_min_idxs = vec![0usize];
unsafe {
if let Some(pd) = lut.pd.as_ref() {
let pd_refs = [pd; 1];
let args = (
codes, &pd_refs,
n_byte_groups, scales_slice, range_vecs,
1, k, mask_slice,
);
if have_gfni() {
search_multi_query_permute_dot_gfni::<1, 8>(
args.0, args.1, args.2, args.3, args.4, args.5, args.6, args.7,
&mut heap_scores, &mut heap_indices,
&mut heap_sizes, &mut heap_mins, &mut heap_min_idxs,
);
} else {
search_multi_query_permute_dot::<1, 8>(
args.0, args.1, args.2, args.3, args.4, args.5, args.6, args.7,
&mut heap_scores, &mut heap_indices,
&mut heap_sizes, &mut heap_mins, &mut heap_min_idxs,
);
}
} else if !lut.split.is_empty() {
let split_refs = [lut.split.as_slice(); 4];
search_multi_query_vnni_dispatch(
codes, &split_refs, &scale_vals, &bias_vals,
n_byte_groups, scales_slice, range_vecs,
1, k, mask_slice,
&mut heap_scores, &mut heap_indices,
&mut heap_sizes, &mut heap_mins, &mut heap_min_idxs,
);
} else if use_avx512 {
search_multi_query_avx512bw(
codes, &lut_refs, &scale_vals, &bias_vals,
n_byte_groups, scales_slice, range_vecs,
1, k, mask_slice,
&mut heap_scores, &mut heap_indices,
&mut heap_sizes, &mut heap_mins, &mut heap_min_idxs,
);
} else {
search_multi_query_avx2(
codes, &lut_refs, &scale_vals, &bias_vals,
n_byte_groups, scales_slice, range_vecs,
1, k, mask_slice,
&mut heap_scores, &mut heap_indices,
&mut heap_sizes, &mut heap_mins, &mut heap_min_idxs,
);
}
}
let sz = heap_sizes[0];
heap_scores[0][..sz]
.iter()
.zip(heap_indices[0][..sz].iter())
.map(|(&s, &i)| (s, i + vec_start as u64))
.collect::<Vec<_>>()
})
.collect();
candidates.sort_unstable_by(|a, b| {
b.0.partial_cmp(&a.0)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.1.cmp(&b.1))
});
candidates.truncate(k);
(
candidates.iter().map(|p| p.0).collect(),
candidates.iter().map(|p| p.1 as i64).collect(),
)
}
#[cfg(target_arch = "x86_64")]
let results = {
#[cfg(test)]
let force_scalar_single =
FORCE_SCALAR_FALLBACK.load(std::sync::atomic::Ordering::Relaxed);
#[cfg(not(test))]
let force_scalar_single = false;
let avx2_fma_ok =
is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma");
let use_avx512 = is_x86_feature_detected!("avx512bw")
&& is_x86_feature_detected!("avx512f")
&& avx2_fma_ok;
let simd_ok = use_avx512 || avx2_fma_ok;
if nq == 1
&& n_blocks >= SINGLE_QUERY_PARALLEL_MIN_BLOCKS
&& simd_ok
&& !force_scalar_single
{
vec![search_single_query_block_parallel(
blocked_codes, &query_luts[0], n_byte_groups, vec_scales,
n_vectors, n_blocks, k, use_avx512, mask,
)]
} else {
let wide_batch_kernel = query_luts.first().is_some_and(|q| q.pd.is_some());
let nq_batch: usize = if wide_batch_kernel
&& rayon::current_num_threads().max(1) == 1
&& nq.div_ceil(10) < nq.div_ceil(8)
{
10
} else {
8
};
#[cfg(test)]
let force_scalar_any = FORCE_SCALAR_FALLBACK.load(std::sync::atomic::Ordering::Relaxed);
#[cfg(not(test))]
let force_scalar_any = false;
let n_quads = nq.div_ceil(nq_batch).max(1);
let n_threads = rayon::current_num_threads().max(1);
let n_ranges = n_block_ranges(
nq,
n_quads,
n_blocks,
n_vectors,
k,
n_threads,
TILES_PER_THREAD,
MIN_TILE_BLOCKS_X86,
serial_required(mask.is_some(), simd_ok, force_scalar_any),
);
let n_ranges = smooth_tile_count(n_ranges, n_quads, n_threads);
let blocks_per_range = n_blocks.div_ceil(n_ranges).max(1);
let block_bytes = n_byte_groups * BLOCK;
let tiles: Vec<(usize, usize)> = (0..n_blocks.max(1))
.step_by(blocks_per_range)
.flat_map(move |b| (0..nq).step_by(nq_batch).map(move |q| (q, b)))
.collect();
let tile_results: Vec<(usize, Vec<Vec<(f32, u64)>>)> = tiles
.into_par_iter()
.map(|(qi_start, block_start)| {
let range_blocks = blocks_per_range.min(n_blocks - block_start);
let vec_start = block_start * BLOCK;
let range_vecs = (range_blocks * BLOCK).min(n_vectors - vec_start);
let codes = &blocked_codes
[block_start * block_bytes..(block_start + range_blocks) * block_bytes];
let scales_slice = &vec_scales[vec_start..vec_start + range_vecs];
let qi_end = (qi_start + nq_batch).min(nq);
let batch_nq = qi_end - qi_start;
let pad_qi = qi_end - 1;
let lut_refs: Vec<&[u8]> = (0..nq_batch)
.map(|i| {
let qi = if qi_start + i < qi_end { qi_start + i } else { pad_qi };
query_luts[qi].uint8_luts.as_slice()
}).collect();
let scale_vals: Vec<f32> = (0..nq_batch)
.map(|i| {
let qi = if qi_start + i < qi_end { qi_start + i } else { pad_qi };
query_luts[qi].scale
}).collect();
let bias_vals: Vec<f32> = (0..nq_batch)
.map(|i| {
let qi = if qi_start + i < qi_end { qi_start + i } else { pad_qi };
query_luts[qi].bias
}).collect();
let mut heap_scores: Vec<Vec<f32>> = (0..batch_nq)
.map(|_| vec![f32::NEG_INFINITY; k]).collect();
let mut heap_indices: Vec<Vec<u64>> = (0..batch_nq)
.map(|_| vec![0u64; k]).collect();
let mut heap_sizes = vec![0usize; batch_nq];
let mut heap_mins = vec![f32::NEG_INFINITY; batch_nq];
let mut heap_min_idxs = vec![0usize; batch_nq];
#[cfg(test)]
let force_scalar =
FORCE_SCALAR_FALLBACK.load(std::sync::atomic::Ordering::Relaxed);
#[cfg(not(test))]
let force_scalar = false;
unsafe {
if !force_scalar && query_luts[pad_qi].pd.is_some() {
macro_rules! pd_dispatch {
($n:literal) => {{
let pd_refs: [&QueryPermuteDot; $n] =
std::array::from_fn(|i| {
let qi = if qi_start + i < qi_end {
qi_start + i
} else {
pad_qi
};
query_luts[qi]
.pd
.as_ref()
.expect("pd built for every query")
});
if have_gfni() {
search_multi_query_permute_dot_gfni::<$n, 1>(
codes, &pd_refs,
n_byte_groups, scales_slice, range_vecs,
batch_nq, k, mask,
&mut heap_scores, &mut heap_indices,
&mut heap_sizes, &mut heap_mins,
&mut heap_min_idxs,
);
} else {
search_multi_query_permute_dot::<$n, 1>(
codes, &pd_refs,
n_byte_groups, scales_slice, range_vecs,
batch_nq, k, mask,
&mut heap_scores, &mut heap_indices,
&mut heap_sizes, &mut heap_mins,
&mut heap_min_idxs,
);
}
}};
}
if nq_batch == 10 {
pd_dispatch!(10)
} else {
pd_dispatch!(8)
}
} else if !force_scalar && !query_luts[pad_qi].split.is_empty() {
let split_refs: Vec<&[u8]> = (0..nq_batch)
.map(|i| {
let qi = if qi_start + i < qi_end { qi_start + i } else { pad_qi };
query_luts[qi].split.as_slice()
})
.collect();
search_multi_query_vnni_dispatch(
codes, &split_refs, &scale_vals, &bias_vals,
n_byte_groups, scales_slice, range_vecs,
batch_nq, k, mask,
&mut heap_scores, &mut heap_indices,
&mut heap_sizes, &mut heap_mins, &mut heap_min_idxs,
);
} else if !force_scalar
&& is_x86_feature_detected!("avx512bw")
&& is_x86_feature_detected!("avx512f")
&& is_x86_feature_detected!("avx2")
&& is_x86_feature_detected!("fma")
{
let mut cs = 0;
while cs < batch_nq {
let ce = (cs + 4).min(batch_nq);
let mut ch_luts = [lut_refs[cs]; 4];
let mut ch_scales = [scale_vals[cs]; 4];
let mut ch_biases = [bias_vals[cs]; 4];
let len = ce - cs;
ch_luts[..len].copy_from_slice(&lut_refs[cs..ce]);
ch_scales[..len].copy_from_slice(&scale_vals[cs..ce]);
ch_biases[..len].copy_from_slice(&bias_vals[cs..ce]);
search_multi_query_avx512bw(
codes, &ch_luts, &ch_scales, &ch_biases,
n_byte_groups, scales_slice, range_vecs,
ce - cs, k, mask,
&mut heap_scores[cs..ce], &mut heap_indices[cs..ce],
&mut heap_sizes[cs..ce], &mut heap_mins[cs..ce],
&mut heap_min_idxs[cs..ce],
);
cs = ce;
}
} else if !force_scalar
&& is_x86_feature_detected!("avx2")
&& is_x86_feature_detected!("fma")
{
let mut cs = 0;
while cs < batch_nq {
let ce = (cs + 4).min(batch_nq);
let mut ch_luts = [lut_refs[cs]; 4];
let mut ch_scales = [scale_vals[cs]; 4];
let mut ch_biases = [bias_vals[cs]; 4];
let len = ce - cs;
ch_luts[..len].copy_from_slice(&lut_refs[cs..ce]);
ch_scales[..len].copy_from_slice(&scale_vals[cs..ce]);
ch_biases[..len].copy_from_slice(&bias_vals[cs..ce]);
search_multi_query_avx2(
codes, &ch_luts, &ch_scales, &ch_biases,
n_byte_groups, scales_slice, range_vecs,
ce - cs, k, mask,
&mut heap_scores[cs..ce], &mut heap_indices[cs..ce],
&mut heap_sizes[cs..ce], &mut heap_mins[cs..ce],
&mut heap_min_idxs[cs..ce],
);
cs = ce;
}
} else {
for qo in 0..batch_nq {
score_query_into_heap(
lut_refs[qo],
scale_vals[qo],
bias_vals[qo],
blocked_codes,
vec_scales,
bits,
n_byte_groups,
n_vectors,
n_blocks,
mask,
k,
&mut heap_scores[qo],
&mut heap_indices[qo],
&mut heap_sizes[qo],
&mut heap_mins[qo],
&mut heap_min_idxs[qo],
);
}
}
}
let cands: Vec<Vec<(f32, u64)>> = (0..batch_nq)
.map(|qo| {
let sz = heap_sizes[qo];
heap_scores[qo][..sz]
.iter()
.zip(heap_indices[qo][..sz].iter())
.map(|(&s, &i)| (s, i + vec_start as u64))
.collect()
})
.collect();
(qi_start, cands)
})
.collect();
let mut merged: Vec<Vec<(f32, u64)>> = vec![Vec::new(); nq];
for (qi_start, cands) in tile_results {
for (off, c) in cands.into_iter().enumerate() {
merged[qi_start + off].extend(c);
}
}
merged
.into_iter()
.map(|mut pairs| {
pairs.sort_unstable_by(|a, b| {
b.0.partial_cmp(&a.0)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.1.cmp(&b.1))
});
pairs.truncate(k);
let s: Vec<f32> = pairs.iter().map(|p| p.0).collect();
let i: Vec<i64> = pairs.iter().map(|p| p.1 as i64).collect();
(s, i)
})
.collect::<Vec<_>>()
}
};
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
let results = {
let results: Vec<(Vec<f32>, Vec<i64>)> = (0..nq)
.into_par_iter()
.map(|qi| {
let qlut = &query_luts[qi];
let mut heap_s = vec![f32::NEG_INFINITY; k];
let mut heap_i = vec![0u64; k];
let mut heap_sz = 0usize;
let mut heap_min = f32::NEG_INFINITY;
let mut heap_mi = 0usize;
score_query_into_heap(
&qlut.uint8_luts,
qlut.scale,
qlut.bias,
blocked_codes,
vec_scales,
bits,
n_byte_groups,
n_vectors,
n_blocks,
mask,
k,
&mut heap_s,
&mut heap_i,
&mut heap_sz,
&mut heap_min,
&mut heap_mi,
);
let mut pairs: Vec<(f32, u64)> = heap_s[..heap_sz].iter()
.zip(heap_i[..heap_sz].iter()).map(|(&s, &i)| (s, i)).collect();
pairs.sort_unstable_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal).then_with(|| a.1.cmp(&b.1)));
(pairs.iter().map(|p| p.0).collect(), pairs.iter().map(|p| p.1 as i64).collect())
})
.collect();
results
};
let mut all_scores = Vec::with_capacity(nq * k);
let mut all_indices = Vec::with_capacity(nq * k);
for (s, i) in &results {
let pad = k.saturating_sub(s.len());
all_scores.extend_from_slice(s);
all_scores.extend(std::iter::repeat(f32::NEG_INFINITY).take(pad));
all_indices.extend_from_slice(i);
all_indices.extend(std::iter::repeat(0i64).take(pad));
}
(all_scores, all_indices)
}
#[cfg(test)]
mod gate_tests {
use super::*;
#[test]
fn single_query_gate_is_at_least_one_tile_wide() {
assert!(
SINGLE_QUERY_PARALLEL_MIN_BLOCKS >= MIN_TILE_BLOCKS,
"single-query pool gate ({SINGLE_QUERY_PARALLEL_MIN_BLOCKS} blocks) fires below \
the batch dispatch's own tile granularity ({MIN_TILE_BLOCKS} blocks): an nq=1 \
search would enter the shared pool at a size where the work is not worth \
splitting (#336)",
);
}
#[test]
fn sub_gate_single_query_never_splits_the_block_axis() {
let n_vectors = (SINGLE_QUERY_PARALLEL_MIN_BLOCKS - 1) * BLOCK;
let n_blocks = n_vectors.div_ceil(BLOCK);
assert!(!single_query_parallelizes(n_vectors));
for &min_tile in &[1usize, 8, 64, MIN_TILE_BLOCKS] {
assert_eq!(
n_block_ranges(1, 1, n_blocks, n_vectors, 10, 16, TILES_PER_THREAD, min_tile, false),
1,
"nq=1 below the pool gate split the block axis at min_tile={min_tile}",
);
}
assert!(n_block_ranges(64, 16, n_blocks, n_vectors, 10, 16, TILES_PER_THREAD, 1, false) > 1);
}
#[test]
fn above_gate_single_query_does_split() {
let n_vectors = SINGLE_QUERY_PARALLEL_MIN_BLOCKS * BLOCK * 4;
assert!(single_query_parallelizes(n_vectors));
assert!(
n_block_ranges(
1,
1,
n_vectors.div_ceil(BLOCK),
n_vectors,
10,
16,
TILES_PER_THREAD,
MIN_TILE_BLOCKS,
false
) > 1
);
}
#[test]
fn the_tile_target_binds_when_the_caps_do_not() {
let n_blocks = 100 * MIN_TILE_BLOCKS; let n_vectors = n_blocks * BLOCK;
assert_eq!(
n_block_ranges(64, 16, n_blocks, n_vectors, 1, 16, TILES_PER_THREAD, MIN_TILE_BLOCKS, false),
TILES_PER_THREAD,
"with both caps clear, the range count IS the per-worker tile target",
);
#[cfg(target_arch = "aarch64")]
{
assert_eq!(TILES_PER_THREAD_NEON, TILES_PER_THREAD * 2);
assert_eq!(MIN_TILE_BLOCKS_NEON, MIN_TILE_BLOCKS / 2);
}
}
#[test]
fn each_serial_condition_forces_one_range_on_its_own() {
let n_vectors = SINGLE_QUERY_PARALLEL_MIN_BLOCKS * BLOCK * 4;
let n_blocks = n_vectors.div_ceil(BLOCK);
assert!(
single_query_parallelizes(n_vectors),
"fixture must sit above the pool gate or the table is vacuous",
);
assert_eq!(
n_block_ranges(64, 16, n_blocks, n_vectors, 10, 16, TILES_PER_THREAD, MIN_TILE_BLOCKS, false),
4,
"baseline range count changed; the rows below prove only that \
the guard fires, so this is the one place the arithmetic \
beneath it is pinned",
);
assert_eq!(
n_block_ranges(64, 1, n_blocks, n_vectors, 10, 1, TILES_PER_THREAD, MIN_TILE_BLOCKS, false),
1,
"a single-threaded pool must not split the block axis",
);
assert_eq!(
n_block_ranges(64, 16, n_blocks, n_vectors, 10, 16, TILES_PER_THREAD, MIN_TILE_BLOCKS, true),
1,
"an explicitly serial call must not split the block axis",
);
let small = (SINGLE_QUERY_PARALLEL_MIN_BLOCKS - 1) * BLOCK;
assert_eq!(
n_block_ranges(1, 1, small.div_ceil(BLOCK), small, 10, 16, TILES_PER_THREAD, 1, false),
1,
"nq=1 below the pool gate must not split the block axis (#147)",
);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn each_term_of_the_serial_predicate_forces_serial_alone() {
assert!(!serial_required(false, true, false));
assert!(serial_required(true, true, false), "a mask alone must force serial");
assert!(serial_required(false, false, false), "absent SIMD alone must force serial");
assert!(serial_required(false, true, true), "forced scalar alone must force serial");
assert!(serial_required(true, false, true));
}
#[cfg(target_arch = "x86_64")]
#[test]
fn split_lut_for_vnni_is_the_documented_byte_map() {
let n_byte_groups = 8usize;
let src: Vec<u8> = (0..n_byte_groups * 32).map(|i| (i % 253) as u8).collect();
let out = split_lut_for_vnni(&src, n_byte_groups);
for g in 0..n_byte_groups {
let c = (g / 4) * 128;
let j = g % 4;
let s = g * 32;
assert_eq!(
&out[c + j * 16..c + j * 16 + 16],
&src[s + 16..s + 32],
"hi half of group {g}",
);
assert_eq!(
&out[c + 64 + j * 16..c + 64 + j * 16 + 16],
&src[s..s + 16],
"lo half of group {g}",
);
}
}
#[test]
fn permute_dot_weights_land_at_their_slots_and_seed_the_bias() {
let dim = 16usize; let n_byte_groups = dim / 2;
let q_rot_row: Vec<f32> = (0..dim).map(|d| (d as f32 - 5.5) * 0.11).collect();
let centroids: Vec<f32> = (0..16).map(|i| (i as f32 - 7.5) / 8.0).collect();
let pd = build_permute_dot(&q_rot_row, ¢roids, dim);
let qmax = q_rot_row.iter().fold(0.0f32, |m, &v| m.max(v.abs()));
let inv_qs = 1.0 / (qmax / 127.0); for g in 0..n_byte_groups {
let (q4, j) = (g / 4, g % 4);
let lo = (q_rot_row[2 * g + 1] * inv_qs).round().clamp(-127.0, 127.0) as i8;
let hi = (q_rot_row[2 * g] * inv_qs).round().clamp(-127.0, 127.0) as i8;
assert_eq!(pd.weights[q4 * 8 + j], lo, "low-nibble weight of group {g}");
assert_eq!(pd.weights[q4 * 8 + 4 + j], hi, "high-nibble weight of group {g}");
}
let wsum: i32 = pd.weights.iter().map(|&w| w as i32).sum();
assert_eq!(pd.zero, -128 * wsum, "accumulator seed must cancel the +128 bias");
}
#[cfg(target_arch = "aarch64")]
#[test]
fn smmla_a_operands_match_the_weight_maps() {
let dim = 32usize; let q0: Vec<f32> = (0..dim).map(|d| (d as f32 - 9.5) * 0.07).collect();
let q1: Vec<f32> = (0..dim).map(|d| (14.5 - d as f32) * 0.05).collect();
let centroids: Vec<f32> = (0..16).map(|i| (i as f32 - 7.5) / 8.0).collect();
let pd0 = build_permute_dot(&q0, ¢roids, dim);
let pd1 = build_permute_dot(&q1, ¢roids, dim);
let pds: [&QueryPermuteDot; 2] = [&pd0, &pd1];
let quads = dim / 8;
let a = build_smmla_a::<2>(&pds, quads);
for q4 in 0..quads {
for (r, pd) in pds.iter().enumerate() {
let dst = q4 * 16; let w = &pd.weights[q4 * 8..q4 * 8 + 8];
for j in 0..4 {
assert_eq!(a[dst + r * 8 + 2 * j], w[4 + j], "even dim, quad {q4} row {r} j {j}");
assert_eq!(a[dst + r * 8 + 2 * j + 1], w[j], "odd dim, quad {q4} row {r} j {j}");
}
}
}
let octs = dim / 16;
let a8 = build_smmla_a_vm8::<2>(&pds, octs);
for q8 in 0..octs {
for (r, pd) in pds.iter().enumerate() {
let dst = q8 * 32; for j in 0..8 {
let g = 8 * q8 + j;
let (q4, slot) = (g / 4, g % 4);
assert_eq!(a8[dst + r * 8 + j], pd.weights[q4 * 8 + 4 + slot], "even dim, oct {q8} row {r} j {j}");
assert_eq!(a8[dst + 16 + r * 8 + j], pd.weights[q4 * 8 + slot], "odd dim, oct {q8} row {r} j {j}");
}
}
}
}
}