#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
use crate::vector::core::distance::DistanceMetric;
#[cfg(target_arch = "x86_64")]
use crate::vector::core::distance_pq_fastscan::apply_metric;
use crate::vector::core::distance_pq_fastscan::{
PqFastScanQuery, distance_pq_fastscan_u8_global_scalar,
};
use crate::vector::index::pq_fastscan_storage::BLOCK_SIZE;
#[cfg(target_arch = "x86_64")]
use crate::vector::index::pq_fastscan_storage::BYTES_PER_SUB_PER_BLOCK;
#[inline]
pub fn is_avx2_supported() -> bool {
#[cfg(target_arch = "x86_64")]
{
is_x86_feature_detected!("avx2")
}
#[cfg(not(target_arch = "x86_64"))]
{
false
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
pub unsafe fn distance_pq_fastscan_block_avx2(
metric: DistanceMetric,
query: &PqFastScanQuery,
packed_block: &[u8],
) -> [f32; BLOCK_SIZE] {
unsafe {
let m = query.params.m as usize;
let zero = _mm256_setzero_si256();
let mask_0f = _mm_set1_epi8(0x0F);
let mut u16_acc_lo = zero;
let mut u16_acc_hi = zero;
let lut_ptr = query.lut4_global.as_ptr();
let block_ptr = packed_block.as_ptr();
for m_idx in 0..m {
let lut_xmm = _mm_loadu_si128(lut_ptr.add(m_idx * 16) as *const __m128i);
let codes_xmm =
_mm_loadu_si128(block_ptr.add(m_idx * BYTES_PER_SUB_PER_BLOCK) as *const __m128i);
let low_nibbles = _mm_and_si128(codes_xmm, mask_0f);
let high_nibbles = _mm_and_si128(_mm_srli_epi16(codes_xmm, 4), mask_0f);
let dist_lo_xmm = _mm_shuffle_epi8(lut_xmm, low_nibbles);
let dist_hi_xmm = _mm_shuffle_epi8(lut_xmm, high_nibbles);
let dist_ymm = _mm256_set_m128i(dist_hi_xmm, dist_lo_xmm);
let dist_u16_lo = _mm256_unpacklo_epi8(dist_ymm, zero);
let dist_u16_hi = _mm256_unpackhi_epi8(dist_ymm, zero);
u16_acc_lo = _mm256_add_epi16(u16_acc_lo, dist_u16_lo);
u16_acc_hi = _mm256_add_epi16(u16_acc_hi, dist_u16_hi);
}
let acc_0_to_16 = _mm256_permute2x128_si256(u16_acc_lo, u16_acc_hi, 0x20);
let acc_16_to_32 = _mm256_permute2x128_si256(u16_acc_lo, u16_acc_hi, 0x31);
let mut u16_dists = [0u16; BLOCK_SIZE];
_mm256_storeu_si256(u16_dists.as_mut_ptr() as *mut __m256i, acc_0_to_16);
_mm256_storeu_si256(u16_dists.as_mut_ptr().add(16) as *mut __m256i, acc_16_to_32);
let mut out = [0.0_f32; BLOCK_SIZE];
for (v, &sum) in u16_dists.iter().enumerate() {
let l2_sq = sum as f32 * query.lut_scale_global + query.lut_bias_sum;
out[v] = apply_metric(metric, l2_sq);
}
out
}
}
pub fn distance_pq_fastscan_block_scalar(
metric: DistanceMetric,
query: &PqFastScanQuery,
packed_block: &[u8],
) -> [f32; BLOCK_SIZE] {
let mut out = [0.0_f32; BLOCK_SIZE];
for (v, slot) in out.iter_mut().enumerate() {
*slot = distance_pq_fastscan_u8_global_scalar(metric, query, packed_block, v);
}
out
}
pub fn distance_pq_fastscan_block(
metric: DistanceMetric,
query: &PqFastScanQuery,
packed_block: &[u8],
) -> [f32; BLOCK_SIZE] {
#[cfg(target_arch = "x86_64")]
{
if is_avx2_supported() && (query.params.m as usize) <= 64 {
return unsafe { distance_pq_fastscan_block_avx2(metric, query, packed_block) };
}
}
#[cfg(target_arch = "aarch64")]
{
use crate::vector::index::pq_fastscan_neon::distance_pq_fastscan_block_neon;
if (query.params.m as usize) <= 64 {
return unsafe { distance_pq_fastscan_block_neon(metric, query, packed_block) };
}
}
distance_pq_fastscan_block_scalar(metric, query, packed_block)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::vector::core::quantization::{PqParams, pq_encode, pq_train_codebook};
use crate::vector::core::vector::Vector;
use crate::vector::index::pq_fastscan_storage::PqFastScanPool;
fn train_small_codebook(
m: usize,
sub_dim: usize,
n: usize,
) -> (PqParams, Vec<f32>, Vec<Vector>) {
let dim = m * sub_dim;
let params = PqParams::new(m as u16, 16, sub_dim as u16).unwrap();
let mut vectors = Vec::with_capacity(n);
for i in 0..n {
let mut v = Vec::with_capacity(dim);
for d in 0..dim {
let x = ((i * 31 + d * 17) % 257) as f32 - 128.0;
let y = ((i.wrapping_mul(d + 1)) % 91) as f32 * 0.1;
v.push(x + y);
}
vectors.push(Vector::new(v));
}
let codebook = pq_train_codebook(dim, params, &vectors).unwrap();
(params, codebook, vectors)
}
fn build_pool(vectors: &[Vector], params: PqParams, codebook: &[f32]) -> PqFastScanPool {
let codes: Vec<Vec<u8>> = vectors
.iter()
.map(|v| pq_encode(&v.data, params, codebook))
.collect();
let entries = codes
.iter()
.enumerate()
.map(|(i, c)| (i as u64, "f".to_string(), c.clone()));
PqFastScanPool::build(params, codebook.to_vec(), entries).unwrap()
}
fn block_slice(pool: &PqFastScanPool, block_idx: usize) -> Vec<u8> {
let stride = pool.block_stride();
let base = block_idx * stride;
pool.packed[base..base + stride].to_vec()
}
#[test]
fn scalar_and_dispatch_paths_agree_on_random_block() {
let (params, codebook, vectors) = train_small_codebook(4, 2, 200);
let pool = build_pool(&vectors, params, &codebook);
let query = PqFastScanQuery::prepare(&vectors[0].data, params, &codebook).unwrap();
let block = block_slice(&pool, 0);
let scalar = distance_pq_fastscan_block_scalar(DistanceMetric::Euclidean, &query, &block);
let dispatch = distance_pq_fastscan_block(DistanceMetric::Euclidean, &query, &block);
assert_eq!(scalar, dispatch);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avx2_block_matches_scalar_for_random_codebook() {
if !is_avx2_supported() {
eprintln!("AVX2 not supported on this CPU; skipping kernel test");
return;
}
let (params, codebook, vectors) = train_small_codebook(8, 2, 256);
let pool = build_pool(&vectors, params, &codebook);
let query = PqFastScanQuery::prepare(&vectors[3].data, params, &codebook).unwrap();
for block_idx in 0..pool.block_count() {
let block = block_slice(&pool, block_idx);
let scalar =
distance_pq_fastscan_block_scalar(DistanceMetric::Euclidean, &query, &block);
let simd = unsafe {
distance_pq_fastscan_block_avx2(DistanceMetric::Euclidean, &query, &block)
};
for v in 0..BLOCK_SIZE {
assert_eq!(
simd[v], scalar[v],
"block {block_idx} vec {v} mismatch (simd={} scalar={})",
simd[v], scalar[v]
);
}
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avx2_matches_scalar_for_cosine_metric() {
if !is_avx2_supported() {
return;
}
let (params, codebook, vectors) = train_small_codebook(4, 2, 200);
let pool = build_pool(&vectors, params, &codebook);
let query = PqFastScanQuery::prepare(&vectors[1].data, params, &codebook).unwrap();
let block = block_slice(&pool, 0);
let scalar = distance_pq_fastscan_block_scalar(DistanceMetric::Cosine, &query, &block);
let simd =
unsafe { distance_pq_fastscan_block_avx2(DistanceMetric::Cosine, &query, &block) };
assert_eq!(simd, scalar);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avx2_matches_scalar_for_partial_block() {
if !is_avx2_supported() {
return;
}
let m = 4;
let sub_dim = 2;
let params = PqParams::new(m as u16, 16, sub_dim as u16).unwrap();
let dim = m * sub_dim;
let codebook: Vec<f32> = (0..params.codebook_len())
.map(|i| (i as f32) * 0.01)
.collect();
let codes: Vec<Vec<u8>> = (0..5)
.map(|i| (0..m).map(|sub| ((i + sub) % 16) as u8).collect())
.collect();
let pool = PqFastScanPool::build(
params,
codebook.clone(),
codes
.iter()
.enumerate()
.map(|(i, c)| (i as u64, "f".to_string(), c.clone())),
)
.unwrap();
let query_vec: Vec<f32> = (0..dim).map(|d| 0.1 * d as f32).collect();
let query = PqFastScanQuery::prepare(&query_vec, params, &codebook).unwrap();
let block = block_slice(&pool, 0);
let scalar = distance_pq_fastscan_block_scalar(DistanceMetric::Euclidean, &query, &block);
let simd =
unsafe { distance_pq_fastscan_block_avx2(DistanceMetric::Euclidean, &query, &block) };
assert_eq!(simd, scalar);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avx2_handles_max_m_64_without_overflow() {
if !is_avx2_supported() {
return;
}
let m = 64usize;
let sub_dim = 1usize;
let params = PqParams::new(m as u16, 16, sub_dim as u16).unwrap();
let codebook: Vec<f32> = (0..params.codebook_len())
.map(|i| (i % 16) as f32)
.collect();
let codes: Vec<Vec<u8>> = (0..BLOCK_SIZE)
.map(|i| (0..m).map(|sub| ((i + sub) % 16) as u8).collect())
.collect();
let pool = PqFastScanPool::build(
params,
codebook.clone(),
codes
.iter()
.enumerate()
.map(|(i, c)| (i as u64, "f".to_string(), c.clone())),
)
.unwrap();
let query_vec: Vec<f32> = (0..m * sub_dim).map(|d| d as f32).collect();
let query = PqFastScanQuery::prepare(&query_vec, params, &codebook).unwrap();
let block = block_slice(&pool, 0);
let scalar = distance_pq_fastscan_block_scalar(DistanceMetric::Euclidean, &query, &block);
let simd =
unsafe { distance_pq_fastscan_block_avx2(DistanceMetric::Euclidean, &query, &block) };
assert_eq!(simd, scalar);
}
}