#[cfg(target_arch = "aarch64")]
use std::arch::aarch64::*;
use crate::vector::core::distance::DistanceMetric;
#[cfg(target_arch = "aarch64")]
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 = "aarch64")]
use crate::vector::index::pq_fastscan_storage::BYTES_PER_SUB_PER_BLOCK;
#[inline]
pub fn is_neon_supported() -> bool {
cfg!(target_arch = "aarch64")
}
#[cfg(target_arch = "aarch64")]
pub unsafe fn distance_pq_fastscan_block_neon(
metric: DistanceMetric,
query: &PqFastScanQuery,
packed_block: &[u8],
) -> [f32; BLOCK_SIZE] {
unsafe {
let m = query.params.m as usize;
let mut acc_0_7 = vdupq_n_u16(0);
let mut acc_8_15 = vdupq_n_u16(0);
let mut acc_16_23 = vdupq_n_u16(0);
let mut acc_24_31 = vdupq_n_u16(0);
let lut_ptr = query.lut4_global.as_ptr();
let block_ptr = packed_block.as_ptr();
for m_idx in 0..m {
let lut = vld1q_u8(lut_ptr.add(m_idx * 16));
let codes = vld1q_u8(block_ptr.add(m_idx * BYTES_PER_SUB_PER_BLOCK));
let low_nib = vandq_u8(codes, vdupq_n_u8(0x0F));
let high_nib = vshrq_n_u8(codes, 4);
let dist_lo = vqtbl1q_u8(lut, low_nib);
let dist_hi = vqtbl1q_u8(lut, high_nib);
acc_0_7 = vaddq_u16(acc_0_7, vmovl_u8(vget_low_u8(dist_lo)));
acc_8_15 = vaddq_u16(acc_8_15, vmovl_high_u8(dist_lo));
acc_16_23 = vaddq_u16(acc_16_23, vmovl_u8(vget_low_u8(dist_hi)));
acc_24_31 = vaddq_u16(acc_24_31, vmovl_high_u8(dist_hi));
}
let mut u16_dists = [0u16; BLOCK_SIZE];
vst1q_u16(u16_dists.as_mut_ptr(), acc_0_7);
vst1q_u16(u16_dists.as_mut_ptr().add(8), acc_8_15);
vst1q_u16(u16_dists.as_mut_ptr().add(16), acc_16_23);
vst1q_u16(u16_dists.as_mut_ptr().add(24), acc_24_31);
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
}
#[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_path_returns_dispatch_compatible_distances() {
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);
for (v, &got) in scalar.iter().enumerate() {
let one_off =
distance_pq_fastscan_u8_global_scalar(DistanceMetric::Euclidean, &query, &block, v);
assert_eq!(got, one_off, "vec {v} mismatch");
}
}
#[cfg(target_arch = "aarch64")]
#[test]
fn neon_block_matches_scalar_for_random_codebook() {
if !is_neon_supported() {
eprintln!("NEON not supported on this target; 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_neon(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 = "aarch64")]
#[test]
fn neon_matches_scalar_for_cosine_metric() {
if !is_neon_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_neon(DistanceMetric::Cosine, &query, &block) };
assert_eq!(simd, scalar);
}
#[cfg(target_arch = "aarch64")]
#[test]
fn neon_matches_scalar_for_partial_block() {
if !is_neon_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_neon(DistanceMetric::Euclidean, &query, &block) };
assert_eq!(simd, scalar);
}
#[cfg(target_arch = "aarch64")]
#[test]
fn neon_handles_max_m_64_without_overflow() {
if !is_neon_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_neon(DistanceMetric::Euclidean, &query, &block) };
assert_eq!(simd, scalar);
}
}