#[cfg(target_arch = "aarch64")]
use core::arch::aarch64::*;
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
pub unsafe fn l2_sq(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
let n = a.len();
let pa = a.as_ptr();
let pb = b.as_ptr();
unsafe {
let mut acc = vdupq_n_f32(0.0);
let mut i = 0usize;
while i + 4 <= n {
let va = vld1q_f32(pa.add(i));
let vb = vld1q_f32(pb.add(i));
let d = vsubq_f32(va, vb);
acc = vfmaq_f32(acc, d, d); i += 4;
}
let mut tail = vaddvq_f32(acc); while i < n {
let d = *pa.add(i) - *pb.add(i);
tail += d * d;
i += 1;
}
tail
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
pub unsafe fn dot(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
let n = a.len();
let pa = a.as_ptr();
let pb = b.as_ptr();
unsafe {
let mut acc = vdupq_n_f32(0.0);
let mut i = 0usize;
while i + 4 <= n {
let va = vld1q_f32(pa.add(i));
let vb = vld1q_f32(pb.add(i));
acc = vfmaq_f32(acc, va, vb);
i += 4;
}
let mut acc_s = vaddvq_f32(acc);
while i < n {
acc_s += *pa.add(i) * *pb.add(i);
i += 1;
}
acc_s
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
pub unsafe fn cosine_parts(a: &[f32], b: &[f32]) -> (f32, f32, f32) {
debug_assert_eq!(a.len(), b.len());
let n = a.len();
let pa = a.as_ptr();
let pb = b.as_ptr();
unsafe {
let mut dot_acc = vdupq_n_f32(0.0);
let mut na_acc = vdupq_n_f32(0.0);
let mut nb_acc = vdupq_n_f32(0.0);
let mut i = 0usize;
while i + 4 <= n {
let va = vld1q_f32(pa.add(i));
let vb = vld1q_f32(pb.add(i));
dot_acc = vfmaq_f32(dot_acc, va, vb);
na_acc = vfmaq_f32(na_acc, va, va);
nb_acc = vfmaq_f32(nb_acc, vb, vb);
i += 4;
}
let mut d = vaddvq_f32(dot_acc);
let mut na = vaddvq_f32(na_acc);
let mut nb = vaddvq_f32(nb_acc);
while i < n {
let x = *pa.add(i);
let y = *pb.add(i);
d += x * y;
na += x * x;
nb += y * y;
i += 1;
}
(d, na, nb)
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon,fp16")]
#[inline]
unsafe fn widen_f16(p: *const u8) -> float32x4_t {
unsafe {
let bits = vld1_u16(p.cast::<u16>());
vcvt_f32_f16(vreinterpret_f16_u16(bits))
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon,fp16")]
pub unsafe fn l2_sq_f16(query: &[f32], stored: &[u8]) -> f32 {
debug_assert_eq!(stored.len(), 2 * query.len());
let n = query.len();
let pq = query.as_ptr();
let ps = stored.as_ptr();
unsafe {
let mut acc = vdupq_n_f32(0.0);
let mut i = 0usize;
while i + 4 <= n {
let vq = vld1q_f32(pq.add(i));
let vs = widen_f16(ps.add(2 * i));
let d = vsubq_f32(vq, vs);
acc = vfmaq_f32(acc, d, d);
i += 4;
}
let mut tail = vaddvq_f32(acc);
while i < n {
let s = f16_at(ps, i);
let d = *pq.add(i) - s;
tail += d * d;
i += 1;
}
tail
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon,fp16")]
pub unsafe fn dot_f16(query: &[f32], stored: &[u8]) -> f32 {
debug_assert_eq!(stored.len(), 2 * query.len());
let n = query.len();
let pq = query.as_ptr();
let ps = stored.as_ptr();
unsafe {
let mut acc = vdupq_n_f32(0.0);
let mut i = 0usize;
while i + 4 <= n {
let vq = vld1q_f32(pq.add(i));
let vs = widen_f16(ps.add(2 * i));
acc = vfmaq_f32(acc, vq, vs);
i += 4;
}
let mut acc_s = vaddvq_f32(acc);
while i < n {
acc_s += *pq.add(i) * f16_at(ps, i);
i += 1;
}
acc_s
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon,fp16")]
pub unsafe fn cosine_parts_f16(query: &[f32], stored: &[u8]) -> (f32, f32, f32) {
debug_assert_eq!(stored.len(), 2 * query.len());
let n = query.len();
let pq = query.as_ptr();
let ps = stored.as_ptr();
unsafe {
let mut dot_acc = vdupq_n_f32(0.0);
let mut nq_acc = vdupq_n_f32(0.0);
let mut ns_acc = vdupq_n_f32(0.0);
let mut i = 0usize;
while i + 4 <= n {
let vq = vld1q_f32(pq.add(i));
let vs = widen_f16(ps.add(2 * i));
dot_acc = vfmaq_f32(dot_acc, vq, vs);
nq_acc = vfmaq_f32(nq_acc, vq, vq);
ns_acc = vfmaq_f32(ns_acc, vs, vs);
i += 4;
}
let mut d = vaddvq_f32(dot_acc);
let mut nq = vaddvq_f32(nq_acc);
let mut ns = vaddvq_f32(ns_acc);
while i < n {
let q = *pq.add(i);
let s = f16_at(ps, i);
d += q * s;
nq += q * q;
ns += s * s;
i += 1;
}
(d, nq, ns)
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
unsafe fn f16_at(p: *const u8, i: usize) -> f32 {
let bits = unsafe { u16::from(*p.add(2 * i)) | (u16::from(*p.add(2 * i + 1)) << 8) };
half::f16::from_bits(bits).to_f32()
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
#[inline]
unsafe fn widen_i8(p: *const i8, scale: f32) -> float32x4_t {
unsafe {
let bytes = vld1_s8(p);
let s16 = vmovl_s8(bytes); let s32 = vmovl_s16(vget_low_s16(s16)); let floats = vcvtq_f32_s32(s32);
vmulq_n_f32(floats, scale)
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
pub unsafe fn l2_sq_i8(query: &[f32], scale: f32, codes: &[i8]) -> f32 {
debug_assert_eq!(codes.len(), query.len());
let n = query.len();
let pq = query.as_ptr();
let pc = codes.as_ptr();
unsafe {
let mut acc = vdupq_n_f32(0.0);
let mut i = 0usize;
while i + 8 <= n {
let vq = vld1q_f32(pq.add(i));
let vs = widen_i8(pc.add(i), scale);
let d = vsubq_f32(vq, vs);
acc = vfmaq_f32(acc, d, d);
i += 4;
}
let mut tail = vaddvq_f32(acc);
while i < n {
let s = f32::from(*pc.add(i)) * scale;
let d = *pq.add(i) - s;
tail += d * d;
i += 1;
}
tail
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
pub unsafe fn dot_i8(query: &[f32], scale: f32, codes: &[i8]) -> f32 {
debug_assert_eq!(codes.len(), query.len());
let n = query.len();
let pq = query.as_ptr();
let pc = codes.as_ptr();
unsafe {
let mut acc = vdupq_n_f32(0.0);
let mut i = 0usize;
while i + 8 <= n {
let vq = vld1q_f32(pq.add(i));
let vs = widen_i8(pc.add(i), scale);
acc = vfmaq_f32(acc, vq, vs);
i += 4;
}
let mut acc_s = vaddvq_f32(acc);
while i < n {
acc_s += *pq.add(i) * (f32::from(*pc.add(i)) * scale);
i += 1;
}
acc_s
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
pub unsafe fn cosine_parts_i8(query: &[f32], scale: f32, codes: &[i8]) -> (f32, f32, f32) {
debug_assert_eq!(codes.len(), query.len());
let n = query.len();
let pq = query.as_ptr();
let pc = codes.as_ptr();
unsafe {
let mut dot_acc = vdupq_n_f32(0.0);
let mut nq_acc = vdupq_n_f32(0.0);
let mut ns_acc = vdupq_n_f32(0.0);
let mut i = 0usize;
while i + 8 <= n {
let vq = vld1q_f32(pq.add(i));
let vs = widen_i8(pc.add(i), scale);
dot_acc = vfmaq_f32(dot_acc, vq, vs);
nq_acc = vfmaq_f32(nq_acc, vq, vq);
ns_acc = vfmaq_f32(ns_acc, vs, vs);
i += 4;
}
let mut d = vaddvq_f32(dot_acc);
let mut nq = vaddvq_f32(nq_acc);
let mut ns = vaddvq_f32(ns_acc);
while i < n {
let q = *pq.add(i);
let s = f32::from(*pc.add(i)) * scale;
d += q * s;
nq += q * q;
ns += s * s;
i += 1;
}
(d, nq, ns)
}
}