#[cfg(target_arch = "x86_64")]
use core::arch::x86_64::*;
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
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 = _mm512_setzero_ps();
let mut i = 0usize;
while i + 16 <= n {
let va = _mm512_loadu_ps(pa.add(i));
let vb = _mm512_loadu_ps(pb.add(i));
let d = _mm512_sub_ps(va, vb);
acc = _mm512_fmadd_ps(d, d, acc);
i += 16;
}
let rem = n - i;
if rem > 0 {
let mask: __mmask16 = (1u16 << rem) - 1;
let va = _mm512_maskz_loadu_ps(mask, pa.add(i));
let vb = _mm512_maskz_loadu_ps(mask, pb.add(i));
let d = _mm512_sub_ps(va, vb);
acc = _mm512_fmadd_ps(d, d, acc);
}
_mm512_reduce_add_ps(acc)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
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 = _mm512_setzero_ps();
let mut i = 0usize;
while i + 16 <= n {
let va = _mm512_loadu_ps(pa.add(i));
let vb = _mm512_loadu_ps(pb.add(i));
acc = _mm512_fmadd_ps(va, vb, acc);
i += 16;
}
let rem = n - i;
if rem > 0 {
let mask: __mmask16 = (1u16 << rem) - 1;
let va = _mm512_maskz_loadu_ps(mask, pa.add(i));
let vb = _mm512_maskz_loadu_ps(mask, pb.add(i));
acc = _mm512_fmadd_ps(va, vb, acc);
}
_mm512_reduce_add_ps(acc)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
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 = _mm512_setzero_ps();
let mut na_acc = _mm512_setzero_ps();
let mut nb_acc = _mm512_setzero_ps();
let mut i = 0usize;
while i + 16 <= n {
let va = _mm512_loadu_ps(pa.add(i));
let vb = _mm512_loadu_ps(pb.add(i));
dot_acc = _mm512_fmadd_ps(va, vb, dot_acc);
na_acc = _mm512_fmadd_ps(va, va, na_acc);
nb_acc = _mm512_fmadd_ps(vb, vb, nb_acc);
i += 16;
}
let rem = n - i;
if rem > 0 {
let mask: __mmask16 = (1u16 << rem) - 1;
let va = _mm512_maskz_loadu_ps(mask, pa.add(i));
let vb = _mm512_maskz_loadu_ps(mask, pb.add(i));
dot_acc = _mm512_fmadd_ps(va, vb, dot_acc);
na_acc = _mm512_fmadd_ps(va, va, na_acc);
nb_acc = _mm512_fmadd_ps(vb, vb, nb_acc);
}
(
_mm512_reduce_add_ps(dot_acc),
_mm512_reduce_add_ps(na_acc),
_mm512_reduce_add_ps(nb_acc),
)
}
}
#[cfg(target_arch = "x86_64")]
#[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 = "avx512f")]
#[inline]
unsafe fn widen_f16(p: *const u8) -> __m512 {
unsafe {
let packed = _mm256_loadu_si256(p.cast::<__m256i>());
_mm512_cvtph_ps(packed)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
#[inline]
unsafe fn widen_i8(p: *const i8, scale: __m512) -> __m512 {
unsafe {
let packed = _mm_loadu_si128(p.cast::<__m128i>());
let widened = _mm512_cvtepi8_epi32(packed);
let floats = _mm512_cvtepi32_ps(widened);
_mm512_mul_ps(floats, scale)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
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 = _mm512_setzero_ps();
let mut i = 0usize;
while i + 16 <= n {
let vq = _mm512_loadu_ps(pq.add(i));
let vs = widen_f16(ps.add(2 * i));
let d = _mm512_sub_ps(vq, vs);
acc = _mm512_fmadd_ps(d, d, acc);
i += 16;
}
let mut tail = _mm512_reduce_add_ps(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 = "avx512f")]
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 = _mm512_setzero_ps();
let mut i = 0usize;
while i + 16 <= n {
let vq = _mm512_loadu_ps(pq.add(i));
let vs = widen_f16(ps.add(2 * i));
acc = _mm512_fmadd_ps(vq, vs, acc);
i += 16;
}
let mut tail = _mm512_reduce_add_ps(acc);
while i < n {
tail += *pq.add(i) * f16_at(ps, i);
i += 1;
}
tail
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
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 = _mm512_setzero_ps();
let mut nq_acc = _mm512_setzero_ps();
let mut ns_acc = _mm512_setzero_ps();
let mut i = 0usize;
while i + 16 <= n {
let vq = _mm512_loadu_ps(pq.add(i));
let vs = widen_f16(ps.add(2 * i));
dot_acc = _mm512_fmadd_ps(vq, vs, dot_acc);
nq_acc = _mm512_fmadd_ps(vq, vq, nq_acc);
ns_acc = _mm512_fmadd_ps(vs, vs, ns_acc);
i += 16;
}
let mut dot = _mm512_reduce_add_ps(dot_acc);
let mut nq = _mm512_reduce_add_ps(nq_acc);
let mut ns = _mm512_reduce_add_ps(ns_acc);
while i < n {
let q = *pq.add(i);
let s = f16_at(ps, i);
dot += q * s;
nq += q * q;
ns += s * s;
i += 1;
}
(dot, nq, ns)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
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 = _mm512_set1_ps(scale);
let mut acc = _mm512_setzero_ps();
let mut i = 0usize;
while i + 16 <= n {
let vq = _mm512_loadu_ps(pq.add(i));
let vs = widen_i8(pc.add(i), vscale);
let d = _mm512_sub_ps(vq, vs);
acc = _mm512_fmadd_ps(d, d, acc);
i += 16;
}
let mut tail = _mm512_reduce_add_ps(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 = "avx512f")]
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 = _mm512_set1_ps(scale);
let mut acc = _mm512_setzero_ps();
let mut i = 0usize;
while i + 16 <= n {
let vq = _mm512_loadu_ps(pq.add(i));
let vs = widen_i8(pc.add(i), vscale);
acc = _mm512_fmadd_ps(vq, vs, acc);
i += 16;
}
let mut tail = _mm512_reduce_add_ps(acc);
while i < n {
tail += *pq.add(i) * (f32::from(*pc.add(i)) * scale);
i += 1;
}
tail
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
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 = _mm512_set1_ps(scale);
let mut dot_acc = _mm512_setzero_ps();
let mut nq_acc = _mm512_setzero_ps();
let mut ns_acc = _mm512_setzero_ps();
let mut i = 0usize;
while i + 16 <= n {
let vq = _mm512_loadu_ps(pq.add(i));
let vs = widen_i8(pc.add(i), vscale);
dot_acc = _mm512_fmadd_ps(vq, vs, dot_acc);
nq_acc = _mm512_fmadd_ps(vq, vq, nq_acc);
ns_acc = _mm512_fmadd_ps(vs, vs, ns_acc);
i += 16;
}
let mut dot = _mm512_reduce_add_ps(dot_acc);
let mut nq = _mm512_reduce_add_ps(nq_acc);
let mut ns = _mm512_reduce_add_ps(ns_acc);
while i < n {
let q = *pq.add(i);
let s = f32::from(*pc.add(i)) * scale;
dot += q * s;
nq += q * q;
ns += s * s;
i += 1;
}
(dot, nq, ns)
}
}