#![allow(dead_code)]
#[inline]
pub fn adc_distance(codes: &[u8], lut: &[Vec<f32>]) -> f32 {
debug_assert_eq!(codes.len(), lut.len());
codes
.iter()
.zip(lut.iter())
.map(|(&code, table)| table[code as usize])
.sum()
}
pub fn adc_batch_distances(codes_batch: &[u8], num_codebooks: usize, lut: &[Vec<f32>]) -> Vec<f32> {
assert_adc_batch_shape(codes_batch, num_codebooks);
assert_eq!(
lut.len(),
num_codebooks,
"lut codebook count {} does not match num_codebooks {}",
lut.len(),
num_codebooks
);
let n_candidates = codes_batch.len() / num_codebooks;
let mut distances = Vec::with_capacity(n_candidates);
for i in 0..n_candidates {
let codes = &codes_batch[i * num_codebooks..(i + 1) * num_codebooks];
distances.push(adc_distance(codes, lut));
}
distances
}
pub fn adc_batch_distances_into<L: PackedLUTData>(
codes_batch: &[u8],
num_codebooks: usize,
lut: &L,
distances: &mut Vec<f32>,
) {
assert_packed_adc_batch_shape(codes_batch, num_codebooks, lut);
let n_candidates = codes_batch.len() / num_codebooks;
distances.clear();
distances.resize(n_candidates, 0.0);
for i in 0..n_candidates {
let codes = &codes_batch[i * num_codebooks..(i + 1) * num_codebooks];
distances[i] = lut.adc_distance(codes);
}
}
#[inline]
fn assert_adc_batch_shape(codes_batch: &[u8], num_codebooks: usize) {
assert!(num_codebooks > 0, "num_codebooks must be non-zero");
assert_eq!(
codes_batch.len() % num_codebooks,
0,
"codes_batch length {} is not a multiple of num_codebooks {}",
codes_batch.len(),
num_codebooks
);
}
#[inline]
fn assert_packed_adc_batch_shape<L: PackedLUTData>(
codes_batch: &[u8],
num_codebooks: usize,
lut: &L,
) {
assert_adc_batch_shape(codes_batch, num_codebooks);
assert_eq!(
lut.num_codebooks(),
num_codebooks,
"packed LUT codebook count {} does not match num_codebooks {}",
lut.num_codebooks(),
num_codebooks
);
let expected_len = match lut.num_codebooks().checked_mul(lut.codebook_size()) {
Some(value) => value,
None => panic!("packed LUT shape overflows usize"),
};
assert!(
lut.data().len() >= expected_len,
"packed LUT length {} is smaller than required length {}",
lut.data().len(),
expected_len
);
}
#[derive(Debug, Clone)]
pub struct PackedLUT {
data: Vec<f32>,
num_codebooks: usize,
codebook_size: usize,
}
impl PackedLUT {
pub fn from_nested(lut: &[Vec<f32>]) -> Self {
let num_codebooks = lut.len();
let codebook_size = if lut.is_empty() { 0 } else { lut[0].len() };
let mut data = Vec::with_capacity(num_codebooks * codebook_size);
for codebook in lut {
data.extend_from_slice(codebook);
}
Self {
data,
num_codebooks,
codebook_size,
}
}
pub fn from_flat(table: &[f32], num_codebooks: usize, codebook_size: usize) -> Self {
debug_assert_eq!(table.len(), num_codebooks * codebook_size);
Self {
data: table.to_vec(),
num_codebooks,
codebook_size,
}
}
#[inline]
pub fn lookup(&self, codebook: usize, code: u8) -> f32 {
self.data[codebook * self.codebook_size + code as usize]
}
#[inline]
pub fn adc_distance(&self, codes: &[u8]) -> f32 {
debug_assert_eq!(codes.len(), self.num_codebooks);
let mut sum = 0.0f32;
for (m, &code) in codes.iter().enumerate() {
sum += self.data[m * self.codebook_size + code as usize];
}
sum
}
#[inline]
pub fn codebook_ptr(&self, codebook_idx: usize) -> *const f32 {
assert!(
codebook_idx < self.num_codebooks,
"codebook_idx {} out of bounds (num_codebooks={})",
codebook_idx,
self.num_codebooks
);
debug_assert!(self.data.len() >= self.num_codebooks * self.codebook_size);
self.data
.as_ptr()
.wrapping_add(codebook_idx * self.codebook_size)
}
pub fn num_codebooks(&self) -> usize {
self.num_codebooks
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct PackedLUTRef<'a> {
data: &'a [f32],
num_codebooks: usize,
codebook_size: usize,
}
impl<'a> PackedLUTRef<'a> {
pub(crate) fn from_flat(table: &'a [f32], num_codebooks: usize, codebook_size: usize) -> Self {
debug_assert_eq!(table.len(), num_codebooks * codebook_size);
Self {
data: table,
num_codebooks,
codebook_size,
}
}
}
pub(crate) trait PackedLUTData {
fn data(&self) -> &[f32];
fn num_codebooks(&self) -> usize;
fn codebook_size(&self) -> usize;
#[inline]
fn codebook_ptr(&self, codebook_idx: usize) -> *const f32 {
assert!(
codebook_idx < self.num_codebooks(),
"codebook_idx {} out of bounds (num_codebooks={})",
codebook_idx,
self.num_codebooks()
);
debug_assert!(self.data().len() >= self.num_codebooks() * self.codebook_size());
self.data()
.as_ptr()
.wrapping_add(codebook_idx * self.codebook_size())
}
#[inline]
fn adc_distance(&self, codes: &[u8]) -> f32 {
debug_assert_eq!(codes.len(), self.num_codebooks());
let mut sum = 0.0f32;
for (m, &code) in codes.iter().enumerate() {
sum += self.data()[m * self.codebook_size() + code as usize];
}
sum
}
}
impl PackedLUTData for PackedLUT {
#[inline]
fn data(&self) -> &[f32] {
&self.data
}
#[inline]
fn num_codebooks(&self) -> usize {
self.num_codebooks
}
#[inline]
fn codebook_size(&self) -> usize {
self.codebook_size
}
}
impl PackedLUTData for PackedLUTRef<'_> {
#[inline]
fn data(&self) -> &[f32] {
self.data
}
#[inline]
fn num_codebooks(&self) -> usize {
self.num_codebooks
}
#[inline]
fn codebook_size(&self) -> usize {
self.codebook_size
}
}
#[cfg(target_arch = "x86_64")]
#[allow(unsafe_code)]
mod x86_64 {
use super::*;
pub(super) fn dispatch_into<L: PackedLUTData>(
codes_batch: &[u8],
num_codebooks: usize,
lut: &L,
distances: &mut Vec<f32>,
) -> bool {
if lut.codebook_size() != 256 || lut.data().len() < num_codebooks.saturating_mul(256) {
return false;
}
let n_candidates = codes_batch.len() / num_codebooks;
if n_candidates >= 16 && std::is_x86_feature_detected!("avx512f") {
unsafe { adc_batch_avx512_into(codes_batch, num_codebooks, lut, distances) };
return true;
}
if n_candidates >= 8 && std::is_x86_feature_detected!("avx2") {
unsafe { adc_batch_avx2_into(codes_batch, num_codebooks, lut, distances) };
return true;
}
false
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
pub(super) unsafe fn adc_batch_avx2_into<L: PackedLUTData>(
codes_batch: &[u8],
num_codebooks: usize,
lut: &L,
distances: &mut Vec<f32>,
) {
use std::arch::x86_64::{
__m256, __m256i, _mm256_add_ps, _mm256_i32gather_ps, _mm256_setzero_ps,
_mm256_storeu_ps,
};
let n_candidates = codes_batch.len() / num_codebooks;
distances.clear();
distances.resize(n_candidates, 0.0);
let chunks_8 = n_candidates / 8;
for chunk in 0..chunks_8 {
let base_idx = chunk * 8;
unsafe {
let mut sum: __m256 = _mm256_setzero_ps();
for m in 0..num_codebooks {
let mut indices = [0i32; 8];
for i in 0..8 {
indices[i] = codes_batch[(base_idx + i) * num_codebooks + m] as i32;
}
let indices_ptr = indices.as_ptr() as *const __m256i;
let idx_vec = std::ptr::read_unaligned(indices_ptr);
let lut_base = lut.codebook_ptr(m);
let gathered = _mm256_i32gather_ps(lut_base, idx_vec, 4);
sum = _mm256_add_ps(sum, gathered);
}
_mm256_storeu_ps(distances.as_mut_ptr().add(base_idx), sum);
}
}
let tail_start = chunks_8 * 8;
for i in tail_start..n_candidates {
let codes = &codes_batch[i * num_codebooks..(i + 1) * num_codebooks];
distances[i] = lut.adc_distance(codes);
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
#[allow(clippy::incompatible_msrv)] pub(super) unsafe fn adc_batch_avx512_into<L: PackedLUTData>(
codes_batch: &[u8],
num_codebooks: usize,
lut: &L,
distances: &mut Vec<f32>,
) {
use std::arch::x86_64::{
__m512, __m512i, _mm512_add_ps, _mm512_i32gather_ps, _mm512_setzero_ps,
_mm512_storeu_ps,
};
let n_candidates = codes_batch.len() / num_codebooks;
distances.clear();
distances.resize(n_candidates, 0.0);
let chunks_16 = n_candidates / 16;
for chunk in 0..chunks_16 {
let base_idx = chunk * 16;
unsafe {
let mut sum: __m512 = _mm512_setzero_ps();
for m in 0..num_codebooks {
let mut indices = [0i32; 16];
for i in 0..16 {
indices[i] = codes_batch[(base_idx + i) * num_codebooks + m] as i32;
}
let indices_ptr = indices.as_ptr() as *const __m512i;
let idx_vec = std::ptr::read_unaligned(indices_ptr);
let lut_base = lut.codebook_ptr(m);
let gathered = _mm512_i32gather_ps(idx_vec, lut_base, 4);
sum = _mm512_add_ps(sum, gathered);
}
_mm512_storeu_ps(distances.as_mut_ptr().add(base_idx), sum);
}
}
let tail_start = chunks_16 * 16;
for i in tail_start..n_candidates {
let codes = &codes_batch[i * num_codebooks..(i + 1) * num_codebooks];
distances[i] = lut.adc_distance(codes);
}
}
}
#[cfg(target_arch = "aarch64")]
#[allow(unsafe_code)]
mod aarch64 {
use super::*;
pub(super) fn dispatch_into<L: PackedLUTData>(
codes_batch: &[u8],
num_codebooks: usize,
lut: &L,
distances: &mut Vec<f32>,
) -> bool {
let n_candidates = codes_batch.len() / num_codebooks;
if n_candidates < 4 {
return false;
}
if lut.codebook_size() == 256 && lut.data().len() >= num_codebooks.saturating_mul(256) {
unsafe {
adc_batch_neon_flat_256_into(codes_batch, num_codebooks, lut.data(), distances)
};
} else {
unsafe { adc_batch_neon_into(codes_batch, num_codebooks, lut, distances) };
}
true
}
pub(super) fn fastscan_block(
block_data: &[u8],
lut_quantized: &[u8],
num_codebooks: usize,
) -> [u16; 32] {
let required_len = fastscan_bytes_per_block(num_codebooks);
assert!(block_data.len() >= required_len);
assert!(lut_quantized.len() >= required_len);
unsafe { fastscan_block_neon(block_data, lut_quantized, num_codebooks) }
}
#[target_feature(enable = "neon")]
pub(super) unsafe fn adc_batch_neon_into<L: PackedLUTData>(
codes_batch: &[u8],
num_codebooks: usize,
lut: &L,
distances: &mut Vec<f32>,
) {
use std::arch::aarch64::{float32x4_t, vaddq_f32, vdupq_n_f32, vsetq_lane_f32, vst1q_f32};
let n_candidates = codes_batch.len() / num_codebooks;
distances.clear();
distances.resize(n_candidates, 0.0);
let chunks_4 = n_candidates / 4;
for chunk in 0..chunks_4 {
let base_idx = chunk * 4;
unsafe {
let mut sum: float32x4_t = vdupq_n_f32(0.0);
for m in 0..num_codebooks {
let c0 = codes_batch[base_idx * num_codebooks + m] as usize;
let c1 = codes_batch[(base_idx + 1) * num_codebooks + m] as usize;
let c2 = codes_batch[(base_idx + 2) * num_codebooks + m] as usize;
let c3 = codes_batch[(base_idx + 3) * num_codebooks + m] as usize;
let lut_base = m * lut.codebook_size();
let lut_data = lut.data();
let v0 = lut_data[lut_base + c0];
let v1 = lut_data[lut_base + c1];
let v2 = lut_data[lut_base + c2];
let v3 = lut_data[lut_base + c3];
let lane0 = vsetq_lane_f32(v0, vdupq_n_f32(0.0), 0);
let lane01 = vsetq_lane_f32(v1, lane0, 1);
let lane012 = vsetq_lane_f32(v2, lane01, 2);
let gathered = vsetq_lane_f32(v3, lane012, 3);
sum = vaddq_f32(sum, gathered);
}
vst1q_f32(distances.as_mut_ptr().add(base_idx), sum);
}
}
let tail_start = chunks_4 * 4;
for i in tail_start..n_candidates {
let codes = &codes_batch[i * num_codebooks..(i + 1) * num_codebooks];
distances[i] = lut.adc_distance(codes);
}
}
#[target_feature(enable = "neon")]
pub(super) unsafe fn adc_batch_neon_flat_256_into(
codes_batch: &[u8],
num_codebooks: usize,
lut_data: &[f32],
distances: &mut Vec<f32>,
) {
use std::arch::aarch64::{float32x4_t, vaddq_f32, vdupq_n_f32, vsetq_lane_f32, vst1q_f32};
debug_assert_eq!(codes_batch.len() % num_codebooks, 0);
debug_assert!(lut_data.len() >= num_codebooks * 256);
let n_candidates = codes_batch.len() / num_codebooks;
distances.clear();
distances.resize(n_candidates, 0.0);
let codes_ptr = codes_batch.as_ptr();
let lut_ptr = lut_data.as_ptr();
let chunks_4 = n_candidates / 4;
for chunk in 0..chunks_4 {
let base_idx = chunk * 4;
unsafe {
let mut sum: float32x4_t = vdupq_n_f32(0.0);
let c0_ptr = codes_ptr.add(base_idx * num_codebooks);
let c1_ptr = c0_ptr.add(num_codebooks);
let c2_ptr = c1_ptr.add(num_codebooks);
let c3_ptr = c2_ptr.add(num_codebooks);
for m in 0..num_codebooks {
let lut_base = lut_ptr.add(m * 256);
let v0 = *lut_base.add(*c0_ptr.add(m) as usize);
let v1 = *lut_base.add(*c1_ptr.add(m) as usize);
let v2 = *lut_base.add(*c2_ptr.add(m) as usize);
let v3 = *lut_base.add(*c3_ptr.add(m) as usize);
let lane0 = vsetq_lane_f32(v0, vdupq_n_f32(0.0), 0);
let lane01 = vsetq_lane_f32(v1, lane0, 1);
let lane012 = vsetq_lane_f32(v2, lane01, 2);
let gathered = vsetq_lane_f32(v3, lane012, 3);
sum = vaddq_f32(sum, gathered);
}
vst1q_f32(distances.as_mut_ptr().add(base_idx), sum);
}
}
let tail_start = chunks_4 * 4;
for (i, out) in distances.iter_mut().enumerate().skip(tail_start) {
let code_base = i * num_codebooks;
let mut sum = 0.0f32;
unsafe {
for m in 0..num_codebooks {
let code = *codes_ptr.add(code_base + m) as usize;
sum += *lut_ptr.add(m * 256 + code);
}
}
*out = sum;
}
}
#[target_feature(enable = "neon")]
pub(super) unsafe fn fastscan_block_neon(
block_data: &[u8],
lut_quantized: &[u8],
num_codebooks: usize,
) -> [u16; 32] {
use std::arch::aarch64::{
uint16x8_t, vaddw_u8, vandq_u8, vdupq_n_u16, vdupq_n_u8, vget_high_u8, vget_low_u8,
vld1q_u8, vqtbl1q_u8, vshrq_n_u8, vst1q_u16,
};
let block_ptr = block_data.as_ptr();
let lut_ptr = lut_quantized.as_ptr();
let low_mask = vdupq_n_u8(0x0f);
let mut low_0_7: uint16x8_t = vdupq_n_u16(0);
let mut low_8_15: uint16x8_t = vdupq_n_u16(0);
let mut high_0_7: uint16x8_t = vdupq_n_u16(0);
let mut high_8_15: uint16x8_t = vdupq_n_u16(0);
for m in 0..num_codebooks {
let (packed, table) = unsafe {
(
vld1q_u8(block_ptr.add(m * 16)),
vld1q_u8(lut_ptr.add(m * 16)),
)
};
let low_idx = vandq_u8(packed, low_mask);
let high_idx = vshrq_n_u8::<4>(packed);
let low_vals = vqtbl1q_u8(table, low_idx);
let high_vals = vqtbl1q_u8(table, high_idx);
low_0_7 = vaddw_u8(low_0_7, vget_low_u8(low_vals));
low_8_15 = vaddw_u8(low_8_15, vget_high_u8(low_vals));
high_0_7 = vaddw_u8(high_0_7, vget_low_u8(high_vals));
high_8_15 = vaddw_u8(high_8_15, vget_high_u8(high_vals));
}
let mut accum = [0u16; 32];
unsafe {
vst1q_u16(accum.as_mut_ptr(), low_0_7);
vst1q_u16(accum.as_mut_ptr().add(8), low_8_15);
vst1q_u16(accum.as_mut_ptr().add(16), high_0_7);
vst1q_u16(accum.as_mut_ptr().add(24), high_8_15);
}
accum
}
}
pub fn adc_batch_dispatch<L: PackedLUTData>(
codes_batch: &[u8],
num_codebooks: usize,
lut: &L,
) -> Vec<f32> {
let mut distances = Vec::new();
adc_batch_dispatch_into(codes_batch, num_codebooks, lut, &mut distances);
distances
}
pub fn adc_batch_dispatch_into<L: PackedLUTData>(
codes_batch: &[u8],
num_codebooks: usize,
lut: &L,
distances: &mut Vec<f32>,
) {
assert_packed_adc_batch_shape(codes_batch, num_codebooks, lut);
#[cfg(target_arch = "x86_64")]
{
if x86_64::dispatch_into(codes_batch, num_codebooks, lut, distances) {
return;
}
}
#[cfg(target_arch = "aarch64")]
{
if aarch64::dispatch_into(codes_batch, num_codebooks, lut, distances) {
return;
}
}
adc_batch_distances_into(codes_batch, num_codebooks, lut, distances);
}
#[derive(Debug, Clone)]
pub struct PackedCodes4bit {
pub data: Vec<u8>,
pub num_vectors: usize,
pub num_codebooks: usize,
pub block_size: usize,
}
#[inline]
fn fastscan_bytes_per_block(num_codebooks: usize) -> usize {
match num_codebooks.checked_mul(16) {
Some(value) => value,
None => panic!("fastscan codebook count overflows usize"),
}
}
impl PackedCodes4bit {
pub fn pack(codes: &[u8], num_vectors: usize, num_codebooks: usize) -> Self {
let expected_codes_len = match num_vectors.checked_mul(num_codebooks) {
Some(value) => value,
None => panic!("packed code shape overflows usize"),
};
assert_eq!(
codes.len(),
expected_codes_len,
"codes length {} does not match num_vectors {} * num_codebooks {}",
codes.len(),
num_vectors,
num_codebooks
);
let block_size = 32usize;
let num_blocks = num_vectors.div_ceil(block_size);
let bytes_per_block = fastscan_bytes_per_block(num_codebooks);
let total_bytes = match num_blocks.checked_mul(bytes_per_block) {
Some(value) => value,
None => panic!("packed FastScan storage length overflows usize"),
};
let mut data = vec![0u8; total_bytes];
for block in 0..num_blocks {
let block_base = block * block_size;
let block_data_offset = block * bytes_per_block;
for m in 0..num_codebooks {
for lane in 0..16usize {
let vi_lo = block_base + lane;
let vi_hi = block_base + 16 + lane;
let code_lo = if vi_lo < num_vectors {
codes[vi_lo * num_codebooks + m] & 0x0F
} else {
0
};
let code_hi = if vi_hi < num_vectors {
codes[vi_hi * num_codebooks + m] & 0x0F
} else {
0
};
data[block_data_offset + m * 16 + lane] = code_lo | (code_hi << 4);
}
}
}
Self {
data,
num_vectors,
num_codebooks,
block_size,
}
}
pub fn num_blocks(&self) -> usize {
self.num_vectors.div_ceil(self.block_size)
}
pub fn bytes_per_block(&self) -> usize {
fastscan_bytes_per_block(self.num_codebooks)
}
pub fn block_data(&self, block_idx: usize) -> &[u8] {
let bpb = self.bytes_per_block();
let start = block_idx * bpb;
&self.data[start..start + bpb]
}
}
pub fn quantize_lut(lut: &[Vec<f32>]) -> (Vec<u8>, f32, f32) {
if lut.is_empty() {
return (Vec::new(), 1.0, 0.0);
}
let mut global_min = f32::INFINITY;
let mut global_max = f32::NEG_INFINITY;
for table in lut {
for &v in table {
if v < global_min {
global_min = v;
}
if v > global_max {
global_max = v;
}
}
}
let range = global_max - global_min;
let (scale, offset) = if range < f32::EPSILON {
(1.0, global_min)
} else {
(range / 255.0, global_min)
};
let inv_scale = if range < f32::EPSILON {
0.0
} else {
255.0 / range
};
let mut quantized = Vec::with_capacity(lut.len() * 16);
for table in lut {
for &v in table {
let q = ((v - offset) * inv_scale).round().clamp(0.0, 255.0) as u8;
quantized.push(q);
}
}
(quantized, scale, offset)
}
pub fn quantize_lut_flat(lut: &[f32], num_codebooks: usize) -> (Vec<u8>, f32, f32) {
debug_assert_eq!(lut.len(), num_codebooks * 16);
if lut.is_empty() {
return (Vec::new(), 1.0, 0.0);
}
let mut global_min = f32::INFINITY;
let mut global_max = f32::NEG_INFINITY;
for &value in lut {
if value < global_min {
global_min = value;
}
if value > global_max {
global_max = value;
}
}
let range = global_max - global_min;
let (scale, offset) = if range < f32::EPSILON {
(1.0, global_min)
} else {
(range / 255.0, global_min)
};
let inv_scale = if range < f32::EPSILON {
0.0
} else {
255.0 / range
};
let quantized = lut
.iter()
.map(|&value| ((value - offset) * inv_scale).round().clamp(0.0, 255.0) as u8)
.collect();
(quantized, scale, offset)
}
pub fn fastscan_block_portable(
block_data: &[u8],
lut_quantized: &[u8],
num_codebooks: usize,
) -> [u16; 32] {
let mut accum = [0u16; 32];
let required_len = fastscan_bytes_per_block(num_codebooks);
assert!(block_data.len() >= required_len);
assert!(lut_quantized.len() >= required_len);
for m in 0..num_codebooks {
let lut_offset = m * 16;
let data_offset = m * 16;
for lane in 0..16usize {
let packed_byte = block_data[data_offset + lane];
let code_lo = (packed_byte & 0x0F) as usize; let code_hi = (packed_byte >> 4) as usize;
accum[lane] += lut_quantized[lut_offset + code_lo] as u16;
accum[lane + 16] += lut_quantized[lut_offset + code_hi] as u16;
}
}
accum
}
pub fn fastscan_batch(packed: &PackedCodes4bit, lut: &[Vec<f32>]) -> Vec<f32> {
assert_eq!(lut.len(), packed.num_codebooks);
let (lut_q, scale, offset) = quantize_lut(lut);
fastscan_batch_quantized(packed, &lut_q, scale, offset)
}
pub fn fastscan_batch_flat(packed: &PackedCodes4bit, lut: &[f32]) -> Vec<f32> {
let (lut_q, scale, offset) = quantize_lut_flat(lut, packed.num_codebooks);
fastscan_batch_quantized(packed, &lut_q, scale, offset)
}
fn fastscan_batch_quantized(
packed: &PackedCodes4bit,
lut_q: &[u8],
scale: f32,
offset: f32,
) -> Vec<f32> {
let num_codebooks = packed.num_codebooks;
let num_blocks = packed.num_blocks();
let bpb = packed.bytes_per_block();
let base_offset = offset * num_codebooks as f32;
let mut distances = Vec::with_capacity(packed.num_vectors);
for block_idx in 0..num_blocks {
let block_start = block_idx * bpb;
let block_end = (block_start + bpb).min(packed.data.len());
let block_data = &packed.data[block_start..block_end];
#[cfg(target_arch = "aarch64")]
let accum = aarch64::fastscan_block(block_data, lut_q, num_codebooks);
#[cfg(not(target_arch = "aarch64"))]
let accum = fastscan_block_portable(block_data, lut_q, num_codebooks);
let vecs_in_block = if block_idx == num_blocks - 1 {
let remaining = packed.num_vectors - block_idx * 32;
remaining.min(32)
} else {
32
};
for &a in accum.iter().take(vecs_in_block) {
distances.push(a as f32 * scale + base_offset);
}
}
distances
}
#[cfg(test)]
mod tests {
use super::*;
fn create_test_lut(num_codebooks: usize, codebook_size: usize) -> Vec<Vec<f32>> {
(0..num_codebooks)
.map(|m| {
(0..codebook_size)
.map(|c| (m * codebook_size + c) as f32 * 0.1)
.collect()
})
.collect()
}
#[test]
fn test_adc_distance_basic() {
let lut = create_test_lut(4, 256);
let codes = vec![0u8, 1, 2, 3];
let dist = adc_distance(&codes, &lut);
let expected = 0.0 + 25.7 + 51.4 + 77.1;
assert!(
(dist - expected).abs() < 0.01,
"got {}, expected {}",
dist,
expected
);
}
#[test]
fn test_packed_lut_equivalence() {
let nested_lut = create_test_lut(8, 256);
let packed_lut = PackedLUT::from_nested(&nested_lut);
let codes = vec![10u8, 20, 30, 40, 50, 60, 70, 80];
let nested_dist = adc_distance(&codes, &nested_lut);
let packed_dist = packed_lut.adc_distance(&codes);
assert!(
(nested_dist - packed_dist).abs() < 1e-6,
"nested={}, packed={}",
nested_dist,
packed_dist
);
}
#[test]
fn test_borrowed_lut_dispatch_matches_owned_lut_and_scalar_oracle() {
let nested_lut = create_test_lut(5, 256);
let flat_lut: Vec<f32> = nested_lut.iter().flatten().copied().collect();
let owned_lut = PackedLUT::from_flat(&flat_lut, 5, 256);
let borrowed_lut = PackedLUTRef::from_flat(&flat_lut, 5, 256);
let n_candidates = 257;
let num_codebooks = 5;
let codes_batch: Vec<u8> = (0..n_candidates * num_codebooks)
.map(|i| ((i * 37) % 256) as u8)
.collect();
let owned = adc_batch_dispatch(&codes_batch, num_codebooks, &owned_lut);
let borrowed = adc_batch_dispatch(&codes_batch, num_codebooks, &borrowed_lut);
assert_eq!(borrowed, owned);
for i in 0..n_candidates {
let codes = &codes_batch[i * num_codebooks..(i + 1) * num_codebooks];
let scalar = adc_distance(codes, &nested_lut);
assert!(
(borrowed[i] - scalar).abs() < 1e-4,
"candidate {}: borrowed={}, scalar={}",
i,
borrowed[i],
scalar,
);
}
}
#[test]
fn test_adc_batch_correctness() {
let lut = create_test_lut(4, 256);
let packed_lut = PackedLUT::from_nested(&lut);
let n_candidates = 100;
let num_codebooks = 4;
let codes_batch: Vec<u8> = (0..n_candidates * num_codebooks)
.map(|i| (i % 256) as u8)
.collect();
let batch_result = adc_batch_dispatch(&codes_batch, num_codebooks, &packed_lut);
for i in 0..n_candidates {
let codes = &codes_batch[i * num_codebooks..(i + 1) * num_codebooks];
let expected = packed_lut.adc_distance(codes);
let actual = batch_result[i];
assert!(
(expected - actual).abs() < 1e-5,
"candidate {}: expected {}, got {}",
i,
expected,
actual
);
}
}
#[test]
fn test_adc_batch_simd_consistency() {
let lut = create_test_lut(16, 256);
let packed_lut = PackedLUT::from_nested(&lut);
let n_candidates = 1000;
let num_codebooks = 16;
let codes_batch: Vec<u8> = (0..n_candidates * num_codebooks)
.map(|i| ((i * 7) % 256) as u8)
.collect();
let result = adc_batch_dispatch(&codes_batch, num_codebooks, &packed_lut);
for i in 0..n_candidates {
let codes = &codes_batch[i * num_codebooks..(i + 1) * num_codebooks];
let expected = packed_lut.adc_distance(codes);
let actual = result[i];
assert!(
(expected - actual).abs() < 1e-4,
"mismatch at {}: expected {}, got {}",
i,
expected,
actual
);
}
}
#[test]
fn test_empty_batch() {
let lut = create_test_lut(4, 256);
let packed_lut = PackedLUT::from_nested(&lut);
let result = adc_batch_dispatch(&[], 4, &packed_lut);
assert!(result.is_empty());
}
#[test]
#[should_panic(expected = "num_codebooks must be non-zero")]
fn test_adc_batch_rejects_zero_codebooks() {
let lut = create_test_lut(1, 256);
let packed_lut = PackedLUT::from_nested(&lut);
let _ = adc_batch_dispatch(&[], 0, &packed_lut);
}
#[test]
#[should_panic(expected = "is not a multiple of num_codebooks")]
fn test_adc_batch_rejects_partial_candidate_row() {
let lut = create_test_lut(2, 256);
let packed_lut = PackedLUT::from_nested(&lut);
let _ = adc_batch_dispatch(&[1, 2, 3], 2, &packed_lut);
}
#[test]
#[should_panic(expected = "does not match num_codebooks")]
fn test_adc_batch_rejects_lut_codebook_mismatch() {
let lut = create_test_lut(3, 256);
let packed_lut = PackedLUT::from_nested(&lut);
let _ = adc_batch_dispatch(&[1, 2, 3, 4], 4, &packed_lut);
}
#[test]
fn test_non_256_lut_dispatch_matches_scalar_oracle() {
let lut = create_test_lut(3, 16);
let packed_lut = PackedLUT::from_nested(&lut);
let n_candidates = 37;
let num_codebooks = 3;
let codes_batch: Vec<u8> = (0..n_candidates * num_codebooks)
.map(|i| ((i * 7 + 3) % 16) as u8)
.collect();
let batch_result = adc_batch_dispatch(&codes_batch, num_codebooks, &packed_lut);
for i in 0..n_candidates {
let codes = &codes_batch[i * num_codebooks..(i + 1) * num_codebooks];
let expected = packed_lut.adc_distance(codes);
let actual = batch_result[i];
assert!(
(expected - actual).abs() < 1e-5,
"candidate {}: expected {}, got {}",
i,
expected,
actual
);
}
}
#[test]
fn test_single_candidate() {
let lut = create_test_lut(4, 256);
let packed_lut = PackedLUT::from_nested(&lut);
let codes = vec![5u8, 10, 15, 20];
let result = adc_batch_dispatch(&codes, 4, &packed_lut);
assert_eq!(result.len(), 1);
let expected = packed_lut.adc_distance(&codes);
assert!((result[0] - expected).abs() < 1e-6);
}
fn create_4bit_lut(num_codebooks: usize) -> Vec<Vec<f32>> {
(0..num_codebooks)
.map(|m| (0..16).map(|c| (m * 16 + c) as f32 * 0.5).collect())
.collect()
}
fn scalar_adc_4bit(codes: &[u8], lut: &[Vec<f32>]) -> f32 {
codes
.iter()
.zip(lut.iter())
.map(|(&c, table)| table[(c & 0x0F) as usize])
.sum()
}
#[test]
fn test_packed_codes_4bit_roundtrip() {
let num_vectors = 5;
let num_codebooks = 4;
let codes: Vec<u8> = (0..num_vectors * num_codebooks)
.map(|idx| {
let i = idx / num_codebooks;
let m = idx % num_codebooks;
((i + m) % 16) as u8
})
.collect();
let packed = PackedCodes4bit::pack(&codes, num_vectors, num_codebooks);
assert_eq!(packed.num_vectors, num_vectors);
assert_eq!(packed.num_codebooks, num_codebooks);
assert_eq!(packed.block_size, 32);
assert_eq!(packed.num_blocks(), 1);
let bpb = packed.bytes_per_block();
let block = &packed.data[0..bpb];
for i in 0..num_vectors {
for m in 0..num_codebooks {
let byte = block[m * 16 + (i % 16)];
let code = if i < 16 { byte & 0x0F } else { byte >> 4 };
assert_eq!(
code,
codes[i * num_codebooks + m],
"mismatch at vec={}, cb={}",
i,
m
);
}
}
}
#[test]
#[should_panic(expected = "does not match num_vectors")]
fn test_packed_codes_4bit_rejects_partial_rows() {
let codes = [0u8, 1, 2];
let _ = PackedCodes4bit::pack(&codes, 2, 2);
}
#[test]
#[should_panic]
fn test_fastscan_block_rejects_short_lut() {
let block_data = [0u8; 16];
let lut_quantized = [0u8; 15];
let _ = fastscan_block_portable(&block_data, &lut_quantized, 1);
}
#[cfg(target_arch = "aarch64")]
#[test]
fn test_neon_fastscan_block_matches_portable() {
let num_vectors = 32;
let num_codebooks = 7;
let codes: Vec<u8> = (0..num_vectors * num_codebooks)
.map(|i| ((i * 11 + 3) % 16) as u8)
.collect();
let packed = PackedCodes4bit::pack(&codes, num_vectors, num_codebooks);
let lut: Vec<f32> = (0..num_codebooks * 16)
.map(|i| ((i * 17 + 5) % 251) as f32 * 0.01)
.collect();
let (lut_q, _, _) = quantize_lut_flat(&lut, num_codebooks);
let portable = fastscan_block_portable(packed.block_data(0), &lut_q, num_codebooks);
let neon = aarch64::fastscan_block(packed.block_data(0), &lut_q, num_codebooks);
assert_eq!(neon, portable);
}
#[test]
fn test_quantize_lut_range() {
let lut = create_4bit_lut(8);
let (quantized, scale, offset) = quantize_lut(&lut);
assert_eq!(quantized.len(), 8 * 16);
assert_eq!(quantized[0], 0); assert_eq!(quantized[quantized.len() - 1], 255);
for (m, table) in lut.iter().enumerate() {
for (c, &original) in table.iter().enumerate() {
let reconstructed = quantized[m * 16 + c] as f32 * scale + offset;
let err = (reconstructed - original).abs();
assert!(
err <= scale / 2.0 + 1e-5,
"m={}, c={}: original={}, reconstructed={}, err={}",
m,
c,
original,
reconstructed,
err,
);
}
}
}
#[test]
fn test_flat_fastscan_matches_nested() {
let num_vectors = 64;
let num_codebooks = 5;
let lut = create_4bit_lut(num_codebooks);
let flat_lut: Vec<f32> = lut.iter().flat_map(|table| table.iter().copied()).collect();
let codes: Vec<u8> = (0..num_vectors * num_codebooks)
.map(|i| (i % 16) as u8)
.collect();
let packed = PackedCodes4bit::pack(&codes, num_vectors, num_codebooks);
let nested = fastscan_batch(&packed, &lut);
let flat = fastscan_batch_flat(&packed, &flat_lut);
assert_eq!(nested, flat);
}
#[test]
fn test_fastscan_vs_scalar_adc() {
let num_codebooks = 8;
let num_vectors = 100;
let lut = create_4bit_lut(num_codebooks);
let codes: Vec<u8> = (0..num_vectors * num_codebooks)
.map(|i| ((i * 7 + 3) % 16) as u8)
.collect();
let scalar_dists: Vec<f32> = (0..num_vectors)
.map(|i| {
let c = &codes[i * num_codebooks..(i + 1) * num_codebooks];
scalar_adc_4bit(c, &lut)
})
.collect();
let packed = PackedCodes4bit::pack(&codes, num_vectors, num_codebooks);
let fastscan_dists = fastscan_batch(&packed, &lut);
assert_eq!(fastscan_dists.len(), num_vectors);
let (_, scale, _) = quantize_lut(&lut);
let tolerance = num_codebooks as f32 * scale * 0.5 + 1e-3;
for i in 0..num_vectors {
let err = (fastscan_dists[i] - scalar_dists[i]).abs();
assert!(
err <= tolerance,
"vec {}: fastscan={}, scalar={}, err={}, tol={}",
i,
fastscan_dists[i],
scalar_dists[i],
err,
tolerance,
);
}
}
#[test]
fn test_fastscan_multiple_blocks() {
let num_codebooks = 4;
let num_vectors = 100; let lut = create_4bit_lut(num_codebooks);
let codes: Vec<u8> = (0..num_vectors * num_codebooks)
.map(|i| ((i * 11 + 5) % 16) as u8)
.collect();
let packed = PackedCodes4bit::pack(&codes, num_vectors, num_codebooks);
assert_eq!(packed.num_blocks(), 4);
let fastscan_dists = fastscan_batch(&packed, &lut);
assert_eq!(fastscan_dists.len(), num_vectors);
let (_, scale, _) = quantize_lut(&lut);
let tolerance = num_codebooks as f32 * scale * 0.5 + 1e-3;
for i in 0..num_vectors {
let c = &codes[i * num_codebooks..(i + 1) * num_codebooks];
let scalar = scalar_adc_4bit(c, &lut);
let err = (fastscan_dists[i] - scalar).abs();
assert!(
err <= tolerance,
"vec {}: fastscan={}, scalar={}, err={}, tol={}",
i,
fastscan_dists[i],
scalar,
err,
tolerance,
);
}
}
#[test]
fn test_fastscan_odd_codebooks() {
let num_codebooks = 5;
let num_vectors = 40;
let lut = create_4bit_lut(num_codebooks);
let codes: Vec<u8> = (0..num_vectors * num_codebooks)
.map(|i| (i % 16) as u8)
.collect();
let packed = PackedCodes4bit::pack(&codes, num_vectors, num_codebooks);
let fastscan_dists = fastscan_batch(&packed, &lut);
assert_eq!(fastscan_dists.len(), num_vectors);
let (_, scale, _) = quantize_lut(&lut);
let tolerance = num_codebooks as f32 * scale * 0.5 + 1e-3;
for i in 0..num_vectors {
let c = &codes[i * num_codebooks..(i + 1) * num_codebooks];
let scalar = scalar_adc_4bit(c, &lut);
let err = (fastscan_dists[i] - scalar).abs();
assert!(
err <= tolerance,
"vec {}: fastscan={}, scalar={}, err={}",
i,
fastscan_dists[i],
scalar,
err,
);
}
}
}