#[cfg(target_arch = "x86_64")]
use core::arch::x86_64::*;
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
#[inline]
unsafe fn hsum256(v: __m256) -> f32 {
let lo = _mm256_castps256_ps128(v);
let hi = _mm256_extractf128_ps(v, 1);
let sum128 = _mm_add_ps(lo, hi);
let shuf = _mm_movehdup_ps(sum128); let sums = _mm_add_ps(sum128, shuf);
let shuf2 = _mm_movehl_ps(shuf, sums);
let sums2 = _mm_add_ss(sums, shuf2);
_mm_cvtss_f32(sums2)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
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 = _mm256_setzero_ps();
let mut i = 0usize;
while i + 8 <= n {
let va = _mm256_loadu_ps(pa.add(i));
let vb = _mm256_loadu_ps(pb.add(i));
let d = _mm256_sub_ps(va, vb);
acc = _mm256_fmadd_ps(d, d, acc); i += 8;
}
let mut tail = hsum256(acc);
while i < n {
let d = *pa.add(i) - *pb.add(i);
tail += d * d;
i += 1;
}
tail
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
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 = _mm256_setzero_ps();
let mut i = 0usize;
while i + 8 <= n {
let va = _mm256_loadu_ps(pa.add(i));
let vb = _mm256_loadu_ps(pb.add(i));
acc = _mm256_fmadd_ps(va, vb, acc);
i += 8;
}
let mut acc_s = hsum256(acc);
while i < n {
acc_s += *pa.add(i) * *pb.add(i);
i += 1;
}
acc_s
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
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 = _mm256_setzero_ps();
let mut na_acc = _mm256_setzero_ps();
let mut nb_acc = _mm256_setzero_ps();
let mut i = 0usize;
while i + 8 <= n {
let va = _mm256_loadu_ps(pa.add(i));
let vb = _mm256_loadu_ps(pb.add(i));
dot_acc = _mm256_fmadd_ps(va, vb, dot_acc);
na_acc = _mm256_fmadd_ps(va, va, na_acc);
nb_acc = _mm256_fmadd_ps(vb, vb, nb_acc);
i += 8;
}
let mut d = hsum256(dot_acc);
let mut na = hsum256(na_acc);
let mut nb = hsum256(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 = "x86_64")]
#[target_feature(enable = "avx2,f16c")]
#[inline]
unsafe fn widen_f16(p: *const u8) -> __m256 {
unsafe {
let packed = _mm_loadu_si128(p.cast::<__m128i>());
_mm256_cvtph_ps(packed)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma,f16c")]
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 = _mm256_setzero_ps();
let mut i = 0usize;
while i + 8 <= n {
let vq = _mm256_loadu_ps(pq.add(i));
let vs = widen_f16(ps.add(2 * i));
let d = _mm256_sub_ps(vq, vs);
acc = _mm256_fmadd_ps(d, d, acc);
i += 8;
}
let mut tail = hsum256(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 = "x86_64")]
#[target_feature(enable = "avx2,fma,f16c")]
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 = _mm256_setzero_ps();
let mut i = 0usize;
while i + 8 <= n {
let vq = _mm256_loadu_ps(pq.add(i));
let vs = widen_f16(ps.add(2 * i));
acc = _mm256_fmadd_ps(vq, vs, acc);
i += 8;
}
let mut acc_s = hsum256(acc);
while i < n {
acc_s += *pq.add(i) * f16_at(ps, i);
i += 1;
}
acc_s
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma,f16c")]
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 = _mm256_setzero_ps();
let mut nq_acc = _mm256_setzero_ps();
let mut ns_acc = _mm256_setzero_ps();
let mut i = 0usize;
while i + 8 <= n {
let vq = _mm256_loadu_ps(pq.add(i));
let vs = widen_f16(ps.add(2 * i));
dot_acc = _mm256_fmadd_ps(vq, vs, dot_acc);
nq_acc = _mm256_fmadd_ps(vq, vq, nq_acc);
ns_acc = _mm256_fmadd_ps(vs, vs, ns_acc);
i += 8;
}
let mut d = hsum256(dot_acc);
let mut nq = hsum256(nq_acc);
let mut ns = hsum256(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 = "x86_64")]
#[target_feature(enable = "f16c")]
#[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 = "x86_64")]
#[target_feature(enable = "avx2")]
#[inline]
unsafe fn widen_i8(p: *const i8, scale: __m256) -> __m256 {
unsafe {
let packed = _mm_loadl_epi64(p.cast::<__m128i>());
let widened = _mm256_cvtepi8_epi32(packed);
let floats = _mm256_cvtepi32_ps(widened);
_mm256_mul_ps(floats, scale)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
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 vscale = _mm256_set1_ps(scale);
let mut acc = _mm256_setzero_ps();
let mut i = 0usize;
while i + 8 <= n {
let vq = _mm256_loadu_ps(pq.add(i));
let vs = widen_i8(pc.add(i), vscale);
let d = _mm256_sub_ps(vq, vs);
acc = _mm256_fmadd_ps(d, d, acc);
i += 8;
}
let mut tail = hsum256(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 = "x86_64")]
#[target_feature(enable = "avx2,fma")]
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 vscale = _mm256_set1_ps(scale);
let mut acc = _mm256_setzero_ps();
let mut i = 0usize;
while i + 8 <= n {
let vq = _mm256_loadu_ps(pq.add(i));
let vs = widen_i8(pc.add(i), vscale);
acc = _mm256_fmadd_ps(vq, vs, acc);
i += 8;
}
let mut acc_s = hsum256(acc);
while i < n {
acc_s += *pq.add(i) * (f32::from(*pc.add(i)) * scale);
i += 1;
}
acc_s
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
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 vscale = _mm256_set1_ps(scale);
let mut dot_acc = _mm256_setzero_ps();
let mut nq_acc = _mm256_setzero_ps();
let mut ns_acc = _mm256_setzero_ps();
let mut i = 0usize;
while i + 8 <= n {
let vq = _mm256_loadu_ps(pq.add(i));
let vs = widen_i8(pc.add(i), vscale);
dot_acc = _mm256_fmadd_ps(vq, vs, dot_acc);
nq_acc = _mm256_fmadd_ps(vq, vq, nq_acc);
ns_acc = _mm256_fmadd_ps(vs, vs, ns_acc);
i += 8;
}
let mut d = hsum256(dot_acc);
let mut nq = hsum256(nq_acc);
let mut ns = hsum256(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)
}
}