pub type Vector = Vec<f32>;
pub fn normalize(vector: &[f32]) -> Vector {
let norm = vector.iter().map(|x| x * x).sum::<f32>().sqrt();
if !norm.is_finite() {
let norm = vector
.iter()
.map(|&x| (x as f64).powi(2))
.sum::<f64>()
.sqrt();
vector.iter().map(|&x| (x as f64 / norm) as f32).collect()
} else if norm == 0.0 {
vec![0.0; vector.len()]
} else {
vector.iter().map(|x| x / norm).collect()
}
}
#[inline]
fn dot_scalar(left: &[f32], right: &[f32]) -> f32 {
debug_assert_eq!(left.len(), right.len());
let n = left.len().min(right.len());
let mut s = 0.0f32;
for i in 0..n {
s += left[i] * right[i];
}
s
}
#[inline]
pub fn dot(left: &[f32], right: &[f32]) -> f32 {
debug_assert_eq!(left.len(), right.len());
let n = left.len().min(right.len());
#[cfg(target_arch = "aarch64")]
{
let prefix = n & !15;
if prefix >= 16 {
let head = unsafe { dot_neon_multiple_of_16(left.as_ptr(), right.as_ptr(), prefix) };
if prefix == n {
return head;
}
return head + dot_scalar(&left[prefix..n], &right[prefix..n]);
}
}
dot_scalar(&left[..n], &right[..n])
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
unsafe fn dot_neon_multiple_of_16(a: *const f32, b: *const f32, len: usize) -> f32 {
unsafe {
use std::arch::aarch64::*;
debug_assert!(len % 16 == 0);
let mut acc0 = vdupq_n_f32(0.0);
let mut acc1 = vdupq_n_f32(0.0);
let mut acc2 = vdupq_n_f32(0.0);
let mut acc3 = vdupq_n_f32(0.0);
let mut i = 0usize;
while i < len {
let a0 = vld1q_f32(a.add(i));
let a1 = vld1q_f32(a.add(i + 4));
let a2 = vld1q_f32(a.add(i + 8));
let a3 = vld1q_f32(a.add(i + 12));
let b0 = vld1q_f32(b.add(i));
let b1 = vld1q_f32(b.add(i + 4));
let b2 = vld1q_f32(b.add(i + 8));
let b3 = vld1q_f32(b.add(i + 12));
acc0 = vfmaq_f32(acc0, a0, b0);
acc1 = vfmaq_f32(acc1, a1, b1);
acc2 = vfmaq_f32(acc2, a2, b2);
acc3 = vfmaq_f32(acc3, a3, b3);
i += 16;
}
let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
vaddvq_f32(acc)
}
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
unsafe fn dot_neon_128(a: *const f32, b: *const f32) -> f32 {
unsafe {
use std::arch::aarch64::*;
let mut acc0 = vdupq_n_f32(0.0);
let mut acc1 = vdupq_n_f32(0.0);
let mut acc2 = vdupq_n_f32(0.0);
let mut acc3 = vdupq_n_f32(0.0);
let mut i = 0usize;
while i < 128 {
let a0 = vld1q_f32(a.add(i));
let a1 = vld1q_f32(a.add(i + 4));
let a2 = vld1q_f32(a.add(i + 8));
let a3 = vld1q_f32(a.add(i + 12));
let b0 = vld1q_f32(b.add(i));
let b1 = vld1q_f32(b.add(i + 4));
let b2 = vld1q_f32(b.add(i + 8));
let b3 = vld1q_f32(b.add(i + 12));
acc0 = vfmaq_f32(acc0, a0, b0);
acc1 = vfmaq_f32(acc1, a1, b1);
acc2 = vfmaq_f32(acc2, a2, b2);
acc3 = vfmaq_f32(acc3, a3, b3);
i += 16;
}
let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
vaddvq_f32(acc)
}
}
pub fn maxsim(query: &[Vector], document: &[Vector]) -> f32 {
if query.is_empty() || document.is_empty() {
return 0.0;
}
let document: Vec<_> = document.iter().map(|v| normalize(v)).collect();
query
.iter()
.map(|query| {
let query = normalize(query);
document
.iter()
.map(|doc| dot(&query, doc))
.fold(f32::NEG_INFINITY, f32::max)
})
.sum()
}
pub fn maxsim_flat(query: &[Vector], document: &[f32], dimension: usize) -> f32 {
#[cfg(target_arch = "aarch64")]
{
if dimension > 0 && dimension % 16 == 0 && dimension <= 4096 {
return maxsim_flat_neon(query, document, dimension);
}
}
maxsim_flat_scalar(query, document, dimension)
}
#[inline]
fn maxsim_flat_scalar(query: &[Vector], document: &[f32], dimension: usize) -> f32 {
query
.iter()
.map(|q| {
document
.chunks_exact(dimension)
.map(|d| dot(q, d))
.fold(f32::NEG_INFINITY, f32::max)
})
.sum()
}
#[cfg(target_arch = "aarch64")]
fn maxsim_flat_neon(query: &[Vector], document: &[f32], dimension: usize) -> f32 {
if dimension == 128 {
return query
.iter()
.map(|q| {
debug_assert_eq!(q.len(), 128);
let qp = q.as_ptr();
let mut best = f32::NEG_INFINITY;
for doc in document.chunks_exact(128) {
let s = unsafe { dot_neon_128(qp, doc.as_ptr()) };
if s > best {
best = s;
}
}
best
})
.sum();
}
query
.iter()
.map(|q| {
debug_assert_eq!(q.len(), dimension);
let qp = q.as_ptr();
let mut best = f32::NEG_INFINITY;
for doc in document.chunks_exact(dimension) {
let s = unsafe { dot_neon_multiple_of_16(qp, doc.as_ptr(), dimension) };
if s > best {
best = s;
}
}
best
})
.sum()
}
#[cfg(test)]
mod tests {
use super::*;
fn deterministic(seed: u64, dim: usize) -> Vector {
let mut s = seed.wrapping_mul(0x9E37_79B9_7F4A_7C15);
(0..dim)
.map(|_| {
s = s
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(((s >> 33) as u32) as f32 / u32::MAX as f32) * 2.0 - 1.0
})
.collect()
}
fn flat(doc: &[Vector]) -> (Vec<f32>, usize) {
let dim = doc[0].len();
(doc.iter().flat_map(|v| v.iter().copied()).collect(), dim)
}
#[test]
fn maxsim_flat_scalar_and_dispatch_agree_on_dim128() {
let query: Vec<_> = (0..8).map(|i| deterministic(0x11 + i, 128)).collect();
let doc_tokens: Vec<_> = (0..50).map(|i| deterministic(0x2000 + i, 128)).collect();
let (flat_doc, dim) = flat(&doc_tokens);
let s = maxsim_flat_scalar(&query, &flat_doc, dim);
let d = maxsim_flat(&query, &flat_doc, dim);
assert!((s - d).abs() < 1e-3, "scalar={s} dispatch={d}");
}
#[test]
fn maxsim_flat_scalar_and_dispatch_agree_on_dim384() {
let query: Vec<_> = (0..12).map(|i| deterministic(0x33 + i, 384)).collect();
let doc_tokens: Vec<_> = (0..30).map(|i| deterministic(0x4000 + i, 384)).collect();
let (flat_doc, dim) = flat(&doc_tokens);
let s = maxsim_flat_scalar(&query, &flat_doc, dim);
let d = maxsim_flat(&query, &flat_doc, dim);
assert!((s - d).abs() < 1e-3, "scalar={s} dispatch={d}");
}
#[test]
fn maxsim_flat_dispatch_falls_back_when_dim_not_multiple_of_16() {
let query: Vec<_> = (0..4).map(|i| deterministic(0x55 + i, 100)).collect();
let doc_tokens: Vec<_> = (0..10).map(|i| deterministic(0x6000 + i, 100)).collect();
let (flat_doc, dim) = flat(&doc_tokens);
let s = maxsim_flat_scalar(&query, &flat_doc, dim);
let d = maxsim_flat(&query, &flat_doc, dim);
assert!((s - d).abs() < 1e-6, "scalar={s} dispatch={d}");
}
}