use super::{CODEBOOK_SCALE, CODEBOOK_U8, QUERY_HIGH_COEF, Query4bitSimd};
impl Query4bitSimd {
#[target_feature(enable = "sse4.1,ssse3")]
pub unsafe fn dotprod_raw_sse(&self, vector: &[u8]) -> i64 {
use core::arch::x86_64::*;
assert_eq!(
vector.len(),
self.expected_vector_bytes(),
"Query4bitSimd::dotprod_raw_sse: vector length mismatch ({} vs expected {})",
vector.len(),
self.expected_vector_bytes(),
);
unsafe {
let codebook = _mm_loadu_si128(CODEBOOK_U8.as_ptr().cast::<__m128i>());
let ones = _mm_set1_epi16(1);
let nibble_mask = _mm_set1_epi8(0x0F);
let mut acc_low = _mm_setzero_si128();
let mut acc_high = _mm_setzero_si128();
for (chunk_idx, [low, high]) in self.query_data.iter().enumerate() {
let v_packed =
_mm_loadl_epi64(vector.as_ptr().add(chunk_idx * 8).cast::<__m128i>());
let v_lo = _mm_and_si128(v_packed, nibble_mask);
let v_hi = _mm_and_si128(_mm_srli_epi16(v_packed, 4), nibble_mask);
let v = _mm_unpacklo_epi8(v_lo, v_hi);
let c_u = _mm_shuffle_epi8(codebook, v);
let q_low = _mm_loadu_si128(low.as_ptr().cast::<__m128i>());
let q_high = _mm_loadu_si128(high.as_ptr().cast::<__m128i>());
let prod_low = _mm_maddubs_epi16(c_u, q_low);
let prod_high = _mm_maddubs_epi16(c_u, q_high);
acc_low = _mm_add_epi32(acc_low, _mm_madd_epi16(prod_low, ones));
acc_high = _mm_add_epi32(acc_high, _mm_madd_epi16(prod_high, ones));
}
if let Some(buf) = self.tail_chunk_scratch(vector) {
let v_packed = _mm_loadl_epi64(buf.as_ptr().cast::<__m128i>());
let v_lo = _mm_and_si128(v_packed, nibble_mask);
let v_hi = _mm_and_si128(_mm_srli_epi16(v_packed, 4), nibble_mask);
let v = _mm_unpacklo_epi8(v_lo, v_hi);
let c_u = _mm_shuffle_epi8(codebook, v);
let q_low = _mm_loadu_si128(self.tail_low.as_ptr().cast::<__m128i>());
let q_high = _mm_loadu_si128(self.tail_high.as_ptr().cast::<__m128i>());
let prod_low = _mm_maddubs_epi16(c_u, q_low);
let prod_high = _mm_maddubs_epi16(c_u, q_high);
acc_low = _mm_add_epi32(acc_low, _mm_madd_epi16(prod_low, ones));
acc_high = _mm_add_epi32(acc_high, _mm_madd_epi16(prod_high, ones));
}
i64::from(hsum_i32_sse(acc_low)) + QUERY_HIGH_COEF * i64::from(hsum_i32_sse(acc_high))
}
}
#[target_feature(enable = "avx2")]
pub unsafe fn dotprod_raw_avx2(&self, vector: &[u8]) -> i64 {
use core::arch::x86_64::*;
assert_eq!(
vector.len(),
self.expected_vector_bytes(),
"Query4bitSimd::dotprod_raw_avx2: vector length mismatch ({} vs expected {})",
vector.len(),
self.expected_vector_bytes(),
);
unsafe {
let codebook_128 = _mm_loadu_si128(CODEBOOK_U8.as_ptr().cast::<__m128i>());
let codebook = _mm256_broadcastsi128_si256(codebook_128);
let ones = _mm256_set1_epi16(1);
let ones_128 = _mm_set1_epi16(1);
let nibble_mask = _mm_set1_epi8(0x0F);
let mut acc = _mm256_setzero_si256();
for (chunk_idx, chunk) in self.query_data.iter().enumerate() {
let low_high = _mm256_loadu_si256(chunk.as_ptr().cast::<__m256i>());
let v_packed =
_mm_loadl_epi64(vector.as_ptr().add(chunk_idx * 8).cast::<__m128i>());
let v_lo = _mm_and_si128(v_packed, nibble_mask);
let v_hi = _mm_and_si128(_mm_srli_epi16(v_packed, 4), nibble_mask);
let v128 = _mm_unpacklo_epi8(v_lo, v_hi);
let v = _mm256_broadcastsi128_si256(v128);
let c = _mm256_shuffle_epi8(codebook, v);
let prods = _mm256_maddubs_epi16(c, low_high);
acc = _mm256_add_epi32(acc, _mm256_madd_epi16(prods, ones));
}
let mut acc_low = _mm256_castsi256_si128(acc);
let mut acc_high = _mm256_extracti128_si256(acc, 1);
if let Some(buf) = self.tail_chunk_scratch(vector) {
let v_packed = _mm_loadl_epi64(buf.as_ptr().cast::<__m128i>());
let v_lo = _mm_and_si128(v_packed, nibble_mask);
let v_hi = _mm_and_si128(_mm_srli_epi16(v_packed, 4), nibble_mask);
let v = _mm_unpacklo_epi8(v_lo, v_hi);
let c_u = _mm_shuffle_epi8(codebook_128, v);
let q_low = _mm_loadu_si128(self.tail_low.as_ptr().cast::<__m128i>());
let q_high = _mm_loadu_si128(self.tail_high.as_ptr().cast::<__m128i>());
let prod_low = _mm_maddubs_epi16(c_u, q_low);
let prod_high = _mm_maddubs_epi16(c_u, q_high);
acc_low = _mm_add_epi32(acc_low, _mm_madd_epi16(prod_low, ones_128));
acc_high = _mm_add_epi32(acc_high, _mm_madd_epi16(prod_high, ones_128));
}
i64::from(hsum_i32_sse(acc_low)) + QUERY_HIGH_COEF * i64::from(hsum_i32_sse(acc_high))
}
}
#[target_feature(enable = "avx512f,avx512bw,avx512vnni,sse4.1,ssse3")]
pub unsafe fn dotprod_raw_avx512_vnni(&self, vector: &[u8]) -> i64 {
use core::arch::x86_64::*;
assert_eq!(
vector.len(),
self.expected_vector_bytes(),
"Query4bitSimd::dotprod_raw_avx512_vnni: vector length mismatch ({} vs expected {})",
vector.len(),
self.expected_vector_bytes(),
);
unsafe {
let codebook_128 = _mm_loadu_si128(CODEBOOK_U8.as_ptr().cast::<__m128i>());
let codebook_512 = _mm512_broadcast_i32x4(codebook_128);
let nibble_mask_128 = _mm_set1_epi8(0x0F);
let ones_128 = _mm_set1_epi16(1);
let mut acc = _mm512_setzero_si512();
let chunks = self.query_data.as_slice();
let n_pairs = chunks.len() / 2;
for i in 0..n_pairs {
let pair_ptr = chunks.as_ptr().add(2 * i).cast::<__m512i>();
let low_high_pair = _mm512_loadu_si512(pair_ptr);
let v_packed_16 = _mm_loadu_si128(vector.as_ptr().add(16 * i).cast::<__m128i>());
let v_lo = _mm_and_si128(v_packed_16, nibble_mask_128);
let v_hi = _mm_and_si128(_mm_srli_epi16(v_packed_16, 4), nibble_mask_128);
let v_chunk_a = _mm_unpacklo_epi8(v_lo, v_hi);
let v_chunk_b = _mm_unpackhi_epi8(v_lo, v_hi);
let v_dup_a = _mm256_broadcastsi128_si256(v_chunk_a);
let v_dup_b = _mm256_broadcastsi128_si256(v_chunk_b);
let v_512 = _mm512_inserti64x4(_mm512_castsi256_si512(v_dup_a), v_dup_b, 1);
let c_512 = _mm512_shuffle_epi8(codebook_512, v_512);
acc = _mm512_dpbusd_epi32(acc, c_512, low_high_pair);
}
let acc_256_lo = _mm512_castsi512_si256(acc);
let acc_256_hi = _mm512_extracti64x4_epi64(acc, 1);
let lane_a_low = _mm256_castsi256_si128(acc_256_lo);
let lane_a_high = _mm256_extracti128_si256(acc_256_lo, 1);
let lane_b_low = _mm256_castsi256_si128(acc_256_hi);
let lane_b_high = _mm256_extracti128_si256(acc_256_hi, 1);
let mut sum_low_xmm = _mm_add_epi32(lane_a_low, lane_b_low);
let mut sum_high_xmm = _mm_add_epi32(lane_a_high, lane_b_high);
if chunks.len() % 2 == 1 {
let tail_chunk = 2 * n_pairs;
let [low, high] = chunks[tail_chunk];
let v_packed =
_mm_loadl_epi64(vector.as_ptr().add(tail_chunk * 8).cast::<__m128i>());
let v_lo = _mm_and_si128(v_packed, nibble_mask_128);
let v_hi = _mm_and_si128(_mm_srli_epi16(v_packed, 4), nibble_mask_128);
let v = _mm_unpacklo_epi8(v_lo, v_hi);
let c_u = _mm_shuffle_epi8(codebook_128, v);
let q_low = _mm_loadu_si128(low.as_ptr().cast::<__m128i>());
let q_high = _mm_loadu_si128(high.as_ptr().cast::<__m128i>());
let prod_low = _mm_maddubs_epi16(c_u, q_low);
let prod_high = _mm_maddubs_epi16(c_u, q_high);
sum_low_xmm = _mm_add_epi32(sum_low_xmm, _mm_madd_epi16(prod_low, ones_128));
sum_high_xmm = _mm_add_epi32(sum_high_xmm, _mm_madd_epi16(prod_high, ones_128));
}
if let Some(buf) = self.tail_chunk_scratch(vector) {
let v_packed = _mm_loadl_epi64(buf.as_ptr().cast::<__m128i>());
let v_lo = _mm_and_si128(v_packed, nibble_mask_128);
let v_hi = _mm_and_si128(_mm_srli_epi16(v_packed, 4), nibble_mask_128);
let v = _mm_unpacklo_epi8(v_lo, v_hi);
let c_u = _mm_shuffle_epi8(codebook_128, v);
let q_low = _mm_loadu_si128(self.tail_low.as_ptr().cast::<__m128i>());
let q_high = _mm_loadu_si128(self.tail_high.as_ptr().cast::<__m128i>());
let prod_low = _mm_maddubs_epi16(c_u, q_low);
let prod_high = _mm_maddubs_epi16(c_u, q_high);
sum_low_xmm = _mm_add_epi32(sum_low_xmm, _mm_madd_epi16(prod_low, ones_128));
sum_high_xmm = _mm_add_epi32(sum_high_xmm, _mm_madd_epi16(prod_high, ones_128));
}
i64::from(hsum_i32_sse(sum_low_xmm))
+ QUERY_HIGH_COEF * i64::from(hsum_i32_sse(sum_high_xmm))
}
}
}
#[target_feature(enable = "sse2")]
unsafe fn hsum_i32_sse(v: core::arch::x86_64::__m128i) -> i32 {
use core::arch::x86_64::*;
let v = _mm_add_epi32(v, _mm_shuffle_epi32(v, 0x4E));
let v = _mm_add_epi32(v, _mm_shuffle_epi32(v, 0xB1));
_mm_cvtsi128_si32(v)
}
#[target_feature(enable = "sse4.1,ssse3")]
pub unsafe fn score_4bit_internal_sse(a: &[u8], b: &[u8]) -> f32 {
use core::arch::x86_64::*;
assert_eq!(
a.len(),
b.len(),
"score_4bit_internal_sse: vector length mismatch ({} vs {})",
a.len(),
b.len(),
);
unsafe {
let codebook_i8 = _mm_xor_si128(
_mm_loadu_si128(CODEBOOK_U8.as_ptr().cast::<__m128i>()),
_mm_set1_epi8(-128i8),
);
let nibble_mask = _mm_set1_epi8(0x0F);
let mut acc = _mm_setzero_si128();
let n_full = a.len() / 8;
for i in 0..n_full {
let va = _mm_loadl_epi64(a.as_ptr().add(i * 8).cast::<__m128i>());
let va_lo = _mm_and_si128(va, nibble_mask);
let va_hi = _mm_and_si128(_mm_srli_epi16(va, 4), nibble_mask);
let a_idx = _mm_unpacklo_epi8(va_lo, va_hi);
let c_a_i8 = _mm_shuffle_epi8(codebook_i8, a_idx);
let vb = _mm_loadl_epi64(b.as_ptr().add(i * 8).cast::<__m128i>());
let vb_lo = _mm_and_si128(vb, nibble_mask);
let vb_hi = _mm_and_si128(_mm_srli_epi16(vb, 4), nibble_mask);
let b_idx = _mm_unpacklo_epi8(vb_lo, vb_hi);
let c_b_i8 = _mm_shuffle_epi8(codebook_i8, b_idx);
let c_a_lo = _mm_cvtepi8_epi16(c_a_i8);
let c_a_hi = _mm_cvtepi8_epi16(_mm_srli_si128(c_a_i8, 8));
let c_b_lo = _mm_cvtepi8_epi16(c_b_i8);
let c_b_hi = _mm_cvtepi8_epi16(_mm_srli_si128(c_b_i8, 8));
let prod_lo = _mm_madd_epi16(c_a_lo, c_b_lo);
let prod_hi = _mm_madd_epi16(c_a_hi, c_b_hi);
acc = _mm_add_epi32(acc, _mm_add_epi32(prod_lo, prod_hi));
}
let simd_bytes = n_full * 8;
let sum = i64::from(hsum_i32_sse(acc))
+ super::score_4bit_internal_integer(&a[simd_bytes..], &b[simd_bytes..]);
sum as f32 / (CODEBOOK_SCALE * CODEBOOK_SCALE)
}
}
#[target_feature(enable = "avx2")]
pub unsafe fn score_4bit_internal_avx2(a: &[u8], b: &[u8]) -> f32 {
use core::arch::x86_64::*;
assert_eq!(
a.len(),
b.len(),
"score_4bit_internal_avx2: vector length mismatch ({} vs {})",
a.len(),
b.len(),
);
unsafe {
let codebook_i8_128 = _mm_xor_si128(
_mm_loadu_si128(CODEBOOK_U8.as_ptr().cast::<__m128i>()),
_mm_set1_epi8(-128i8),
);
let codebook_i8 = _mm256_broadcastsi128_si256(codebook_i8_128);
let nibble_mask = _mm_set1_epi8(0x0F);
let mut acc = _mm256_setzero_si256();
let n_iters = a.len() / 16;
for i in 0..n_iters {
let va = _mm_loadu_si128(a.as_ptr().add(16 * i).cast::<__m128i>());
let va_lo = _mm_and_si128(va, nibble_mask);
let va_hi = _mm_and_si128(_mm_srli_epi16(va, 4), nibble_mask);
let a_idx_0 = _mm_unpacklo_epi8(va_lo, va_hi);
let a_idx_1 = _mm_unpackhi_epi8(va_lo, va_hi);
let a_idx_256 = _mm256_inserti128_si256(_mm256_castsi128_si256(a_idx_0), a_idx_1, 1);
let c_a_i8 = _mm256_shuffle_epi8(codebook_i8, a_idx_256);
let vb = _mm_loadu_si128(b.as_ptr().add(16 * i).cast::<__m128i>());
let vb_lo = _mm_and_si128(vb, nibble_mask);
let vb_hi = _mm_and_si128(_mm_srli_epi16(vb, 4), nibble_mask);
let b_idx_0 = _mm_unpacklo_epi8(vb_lo, vb_hi);
let b_idx_1 = _mm_unpackhi_epi8(vb_lo, vb_hi);
let b_idx_256 = _mm256_inserti128_si256(_mm256_castsi128_si256(b_idx_0), b_idx_1, 1);
let c_b_i8 = _mm256_shuffle_epi8(codebook_i8, b_idx_256);
let c_a_lo = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(c_a_i8));
let c_a_hi = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(c_a_i8, 1));
let c_b_lo = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(c_b_i8));
let c_b_hi = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(c_b_i8, 1));
let prod_lo = _mm256_madd_epi16(c_a_lo, c_b_lo);
let prod_hi = _mm256_madd_epi16(c_a_hi, c_b_hi);
acc = _mm256_add_epi32(acc, _mm256_add_epi32(prod_lo, prod_hi));
}
let acc_lo = _mm256_castsi256_si128(acc);
let acc_hi = _mm256_extracti128_si256(acc, 1);
let simd_bytes = n_iters * 16;
let sum = i64::from(hsum_i32_sse(_mm_add_epi32(acc_lo, acc_hi)))
+ super::score_4bit_internal_integer(&a[simd_bytes..], &b[simd_bytes..]);
sum as f32 / (CODEBOOK_SCALE * CODEBOOK_SCALE)
}
}
#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
pub unsafe fn score_4bit_internal_avx512_vnni(a: &[u8], b: &[u8]) -> f32 {
use core::arch::x86_64::*;
assert_eq!(
a.len(),
b.len(),
"score_4bit_internal_avx512_vnni: vector length mismatch ({} vs {})",
a.len(),
b.len(),
);
unsafe {
let codebook_i8_128 = _mm_xor_si128(
_mm_loadu_si128(CODEBOOK_U8.as_ptr().cast::<__m128i>()),
_mm_set1_epi8(-128i8),
);
let codebook_i8_256 = _mm256_broadcastsi128_si256(codebook_i8_128);
let nibble_mask = _mm_set1_epi8(0x0F);
let mut acc = _mm512_setzero_si512();
let n_iters = a.len() / 16;
for i in 0..n_iters {
let va = _mm_loadu_si128(a.as_ptr().add(16 * i).cast::<__m128i>());
let va_lo = _mm_and_si128(va, nibble_mask);
let va_hi = _mm_and_si128(_mm_srli_epi16(va, 4), nibble_mask);
let a_idx_0 = _mm_unpacklo_epi8(va_lo, va_hi);
let a_idx_1 = _mm_unpackhi_epi8(va_lo, va_hi);
let a_idx_256 = _mm256_inserti128_si256(_mm256_castsi128_si256(a_idx_0), a_idx_1, 1);
let c_a_i8_256 = _mm256_shuffle_epi8(codebook_i8_256, a_idx_256);
let vb = _mm_loadu_si128(b.as_ptr().add(16 * i).cast::<__m128i>());
let vb_lo = _mm_and_si128(vb, nibble_mask);
let vb_hi = _mm_and_si128(_mm_srli_epi16(vb, 4), nibble_mask);
let b_idx_0 = _mm_unpacklo_epi8(vb_lo, vb_hi);
let b_idx_1 = _mm_unpackhi_epi8(vb_lo, vb_hi);
let b_idx_256 = _mm256_inserti128_si256(_mm256_castsi128_si256(b_idx_0), b_idx_1, 1);
let c_b_i8_256 = _mm256_shuffle_epi8(codebook_i8_256, b_idx_256);
let c_a_i16 = _mm512_cvtepi8_epi16(c_a_i8_256);
let c_b_i16 = _mm512_cvtepi8_epi16(c_b_i8_256);
acc = _mm512_dpwssd_epi32(acc, c_a_i16, c_b_i16);
}
let acc_256_lo = _mm512_castsi512_si256(acc);
let acc_256_hi = _mm512_extracti64x4_epi64(acc, 1);
let acc_256 = _mm256_add_epi32(acc_256_lo, acc_256_hi);
let acc_128 = _mm_add_epi32(
_mm256_castsi256_si128(acc_256),
_mm256_extracti128_si256(acc_256, 1),
);
let simd_bytes = n_iters * 16;
let sum = i64::from(hsum_i32_sse(acc_128))
+ super::score_4bit_internal_integer(&a[simd_bytes..], &b[simd_bytes..]);
sum as f32 / (CODEBOOK_SCALE * CODEBOOK_SCALE)
}
}
#[target_feature(enable = "sse4.1,ssse3")]
pub unsafe fn score_4bit_internal_weighted_sse(a: &[u8], b: &[u8], weights: &[i16]) -> i64 {
use core::arch::x86_64::*;
assert_eq!(
a.len(),
b.len(),
"score_4bit_internal_weighted_sse: vector length mismatch ({} vs {})",
a.len(),
b.len(),
);
assert_eq!(
weights.len(),
2 * a.len(),
"score_4bit_internal_weighted_sse: weights length {} != 2 · a.len() {}",
weights.len(),
2 * a.len(),
);
unsafe {
let codebook_i8 = _mm_xor_si128(
_mm_loadu_si128(CODEBOOK_U8.as_ptr().cast::<__m128i>()),
_mm_set1_epi8(-128i8),
);
let nibble_mask = _mm_set1_epi8(0x0F);
let mut acc = _mm_setzero_si128();
let n_full = a.len() / 8;
for i in 0..n_full {
let va = _mm_loadl_epi64(a.as_ptr().add(i * 8).cast::<__m128i>());
let va_lo = _mm_and_si128(va, nibble_mask);
let va_hi = _mm_and_si128(_mm_srli_epi16(va, 4), nibble_mask);
let a_idx = _mm_unpacklo_epi8(va_lo, va_hi);
let c_a_i8 = _mm_shuffle_epi8(codebook_i8, a_idx);
let vb = _mm_loadl_epi64(b.as_ptr().add(i * 8).cast::<__m128i>());
let vb_lo = _mm_and_si128(vb, nibble_mask);
let vb_hi = _mm_and_si128(_mm_srli_epi16(vb, 4), nibble_mask);
let b_idx = _mm_unpacklo_epi8(vb_lo, vb_hi);
let c_b_i8 = _mm_shuffle_epi8(codebook_i8, b_idx);
let c_a_lo = _mm_cvtepi8_epi16(c_a_i8);
let c_a_hi = _mm_cvtepi8_epi16(_mm_srli_si128(c_a_i8, 8));
let c_b_lo = _mm_cvtepi8_epi16(c_b_i8);
let c_b_hi = _mm_cvtepi8_epi16(_mm_srli_si128(c_b_i8, 8));
let prod_lo = _mm_mullo_epi16(c_a_lo, c_b_lo);
let prod_hi = _mm_mullo_epi16(c_a_hi, c_b_hi);
let w_lo = _mm_loadu_si128(weights.as_ptr().add(16 * i).cast::<__m128i>());
let w_hi = _mm_loadu_si128(weights.as_ptr().add(16 * i + 8).cast::<__m128i>());
let pw_lo = _mm_madd_epi16(prod_lo, w_lo);
let pw_hi = _mm_madd_epi16(prod_hi, w_hi);
let pw = _mm_add_epi32(pw_lo, pw_hi);
let pw_lo_i64 = _mm_cvtepi32_epi64(pw);
let pw_hi_i64 = _mm_cvtepi32_epi64(_mm_srli_si128(pw, 8));
acc = _mm_add_epi64(acc, pw_lo_i64);
acc = _mm_add_epi64(acc, pw_hi_i64);
}
let mut tmp = [0i64; 2];
_mm_storeu_si128(tmp.as_mut_ptr().cast::<__m128i>(), acc);
let simd_sum = tmp[0] + tmp[1];
let simd_bytes = n_full * 8;
let tail = super::score_4bit_internal_weighted_scalar(
&a[simd_bytes..],
&b[simd_bytes..],
&weights[2 * simd_bytes..],
);
simd_sum + tail
}
}
#[target_feature(enable = "avx2")]
pub unsafe fn score_4bit_internal_weighted_avx2(a: &[u8], b: &[u8], weights: &[i16]) -> i64 {
use core::arch::x86_64::*;
assert_eq!(
a.len(),
b.len(),
"score_4bit_internal_weighted_avx2: vector length mismatch ({} vs {})",
a.len(),
b.len(),
);
assert_eq!(
weights.len(),
2 * a.len(),
"score_4bit_internal_weighted_avx2: weights length {} != 2 · a.len() {}",
weights.len(),
2 * a.len(),
);
unsafe {
let codebook_i8_128 = _mm_xor_si128(
_mm_loadu_si128(CODEBOOK_U8.as_ptr().cast::<__m128i>()),
_mm_set1_epi8(-128i8),
);
let codebook_i8 = _mm256_broadcastsi128_si256(codebook_i8_128);
let nibble_mask = _mm_set1_epi8(0x0F);
let mut acc = _mm256_setzero_si256();
let n_iters = a.len() / 16;
for i in 0..n_iters {
let va = _mm_loadu_si128(a.as_ptr().add(16 * i).cast::<__m128i>());
let va_lo = _mm_and_si128(va, nibble_mask);
let va_hi = _mm_and_si128(_mm_srli_epi16(va, 4), nibble_mask);
let a_idx_0 = _mm_unpacklo_epi8(va_lo, va_hi);
let a_idx_1 = _mm_unpackhi_epi8(va_lo, va_hi);
let a_idx_256 = _mm256_inserti128_si256(_mm256_castsi128_si256(a_idx_0), a_idx_1, 1);
let c_a_i8 = _mm256_shuffle_epi8(codebook_i8, a_idx_256);
let vb = _mm_loadu_si128(b.as_ptr().add(16 * i).cast::<__m128i>());
let vb_lo = _mm_and_si128(vb, nibble_mask);
let vb_hi = _mm_and_si128(_mm_srli_epi16(vb, 4), nibble_mask);
let b_idx_0 = _mm_unpacklo_epi8(vb_lo, vb_hi);
let b_idx_1 = _mm_unpackhi_epi8(vb_lo, vb_hi);
let b_idx_256 = _mm256_inserti128_si256(_mm256_castsi128_si256(b_idx_0), b_idx_1, 1);
let c_b_i8 = _mm256_shuffle_epi8(codebook_i8, b_idx_256);
let c_a_lo = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(c_a_i8));
let c_a_hi = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(c_a_i8, 1));
let c_b_lo = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(c_b_i8));
let c_b_hi = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(c_b_i8, 1));
let prod_lo = _mm256_mullo_epi16(c_a_lo, c_b_lo);
let prod_hi = _mm256_mullo_epi16(c_a_hi, c_b_hi);
let w_lo = _mm256_loadu_si256(weights.as_ptr().add(32 * i).cast::<__m256i>());
let w_hi = _mm256_loadu_si256(weights.as_ptr().add(32 * i + 16).cast::<__m256i>());
let pw_lo = _mm256_madd_epi16(prod_lo, w_lo);
let pw_hi = _mm256_madd_epi16(prod_hi, w_hi);
let pw = _mm256_add_epi32(pw_lo, pw_hi);
let pw_lo_i64 = _mm256_cvtepi32_epi64(_mm256_castsi256_si128(pw));
let pw_hi_i64 = _mm256_cvtepi32_epi64(_mm256_extracti128_si256(pw, 1));
acc = _mm256_add_epi64(acc, pw_lo_i64);
acc = _mm256_add_epi64(acc, pw_hi_i64);
}
let acc_lo = _mm256_castsi256_si128(acc);
let acc_hi = _mm256_extracti128_si256(acc, 1);
let summed = _mm_add_epi64(acc_lo, acc_hi);
let mut tmp = [0i64; 2];
_mm_storeu_si128(tmp.as_mut_ptr().cast::<__m128i>(), summed);
let simd_sum = tmp[0] + tmp[1];
let simd_bytes = n_iters * 16;
let tail = super::score_4bit_internal_weighted_scalar(
&a[simd_bytes..],
&b[simd_bytes..],
&weights[2 * simd_bytes..],
);
simd_sum + tail
}
}
#[cfg(test)]
mod tests {
use rand::SeedableRng as _;
use rand::prelude::StdRng;
use super::super::super::shared::pack_codes;
use super::super::shared::{PARITY_DIMS, random_inputs};
use super::super::{
Query4bitSimd, score_4bit_internal_scalar, score_4bit_internal_weighted_scalar,
};
use super::{
score_4bit_internal_avx2, score_4bit_internal_avx512_vnni, score_4bit_internal_sse,
score_4bit_internal_weighted_avx2, score_4bit_internal_weighted_sse,
};
fn random_weights(rng: &mut StdRng, vec_bytes: usize) -> Vec<i16> {
use rand::RngExt;
(0..2 * vec_bytes)
.map(|_| rng.random_range(0..=i16::MAX))
.collect()
}
#[test]
fn test_sse_matches_scalar() {
if !std::is_x86_feature_detected!("ssse3") || !std::is_x86_feature_detected!("sse4.1") {
return;
}
let mut rng = StdRng::seed_from_u64(7);
for &dim in PARITY_DIMS {
let (simd_query, vector) = random_inputs(&mut rng, dim);
let scalar = simd_query.dotprod_raw(&vector);
let sse = unsafe { simd_query.dotprod_raw_sse(&vector) };
assert_eq!(scalar, sse, "scalar {scalar} != sse {sse} at dim {dim}");
}
}
#[test]
fn test_avx2_matches_scalar() {
if !std::is_x86_feature_detected!("avx2") {
return;
}
let mut rng = StdRng::seed_from_u64(7);
for &dim in PARITY_DIMS {
let (simd_query, vector) = random_inputs(&mut rng, dim);
let scalar = simd_query.dotprod_raw(&vector);
let avx2 = unsafe { simd_query.dotprod_raw_avx2(&vector) };
assert_eq!(scalar, avx2, "scalar {scalar} != avx2 {avx2} at dim {dim}");
}
}
#[test]
fn test_avx512_vnni_matches_scalar() {
if !(std::is_x86_feature_detected!("avx512f")
&& std::is_x86_feature_detected!("avx512bw")
&& std::is_x86_feature_detected!("avx512vnni"))
{
return;
}
let mut rng = StdRng::seed_from_u64(7);
for &dim in PARITY_DIMS {
let (simd_query, vector) = random_inputs(&mut rng, dim);
let scalar = simd_query.dotprod_raw(&vector);
let vnni512 = unsafe { simd_query.dotprod_raw_avx512_vnni(&vector) };
assert_eq!(
scalar, vnni512,
"scalar {scalar} != avx512_vnni {vnni512} at dim {dim}"
);
}
}
#[test]
fn test_saturation_safety_64k() {
let dim = 65_536;
let query = vec![1.0_f32; dim];
let indices: Vec<u8> = vec![15; dim]; let vector = pack_codes(&indices, 4);
let q = Query4bitSimd::new(&query);
let scalar = q.dotprod_raw(&vector);
unsafe {
if std::is_x86_feature_detected!("ssse3") && std::is_x86_feature_detected!("sse4.1") {
let sse = q.dotprod_raw_sse(&vector);
assert_eq!(scalar, sse, "sse disagrees at dim={dim}");
}
if std::is_x86_feature_detected!("avx2") {
let avx2 = q.dotprod_raw_avx2(&vector);
assert_eq!(scalar, avx2, "avx2 disagrees at dim={dim}");
}
if std::is_x86_feature_detected!("avx512f")
&& std::is_x86_feature_detected!("avx512bw")
&& std::is_x86_feature_detected!("avx512vnni")
{
let v512 = q.dotprod_raw_avx512_vnni(&vector);
assert_eq!(scalar, v512, "avx512_vnni disagrees at dim={dim}");
}
}
}
#[test]
fn test_score_sse_matches_scalar() {
if !std::is_x86_feature_detected!("ssse3") || !std::is_x86_feature_detected!("sse4.1") {
return;
}
let mut rng = StdRng::seed_from_u64(7);
for &dim in PARITY_DIMS {
let (_, vec_a) = random_inputs(&mut rng, dim);
let (_, vec_b) = random_inputs(&mut rng, dim);
let scalar = score_4bit_internal_scalar(&vec_a, &vec_b);
let sse = unsafe { score_4bit_internal_sse(&vec_a, &vec_b) };
assert_eq!(
scalar, sse,
"score: scalar {scalar} != sse {sse} at dim {dim}"
);
}
}
#[test]
fn test_score_avx2_matches_scalar() {
if !std::is_x86_feature_detected!("avx2") {
return;
}
let mut rng = StdRng::seed_from_u64(7);
for &dim in PARITY_DIMS {
let (_, vec_a) = random_inputs(&mut rng, dim);
let (_, vec_b) = random_inputs(&mut rng, dim);
let scalar = score_4bit_internal_scalar(&vec_a, &vec_b);
let avx2 = unsafe { score_4bit_internal_avx2(&vec_a, &vec_b) };
assert_eq!(
scalar, avx2,
"score: scalar {scalar} != avx2 {avx2} at dim {dim}"
);
}
}
#[test]
fn test_score_avx512_vnni_matches_scalar() {
if !(std::is_x86_feature_detected!("avx512f")
&& std::is_x86_feature_detected!("avx512bw")
&& std::is_x86_feature_detected!("avx512vnni"))
{
return;
}
let mut rng = StdRng::seed_from_u64(7);
for &dim in PARITY_DIMS {
let (_, vec_a) = random_inputs(&mut rng, dim);
let (_, vec_b) = random_inputs(&mut rng, dim);
let scalar = score_4bit_internal_scalar(&vec_a, &vec_b);
let vnni512 = unsafe { score_4bit_internal_avx512_vnni(&vec_a, &vec_b) };
assert_eq!(
scalar, vnni512,
"score: scalar {scalar} != avx512_vnni {vnni512} at dim {dim}"
);
}
}
#[test]
fn test_score_weighted_sse_matches_scalar() {
if !std::is_x86_feature_detected!("ssse3") || !std::is_x86_feature_detected!("sse4.1") {
return;
}
let mut rng = StdRng::seed_from_u64(0xBEEF);
for &dim in PARITY_DIMS {
let (_, vec_a) = random_inputs(&mut rng, dim);
let (_, vec_b) = random_inputs(&mut rng, dim);
let weights = random_weights(&mut rng, vec_a.len());
let scalar = score_4bit_internal_weighted_scalar(&vec_a, &vec_b, &weights);
let sse = unsafe { score_4bit_internal_weighted_sse(&vec_a, &vec_b, &weights) };
assert_eq!(
scalar, sse,
"weighted: scalar {scalar} != sse {sse} at dim {dim}"
);
}
}
#[test]
fn test_score_weighted_avx2_matches_scalar() {
if !std::is_x86_feature_detected!("avx2") {
return;
}
let mut rng = StdRng::seed_from_u64(0xBEEF);
for &dim in PARITY_DIMS {
let (_, vec_a) = random_inputs(&mut rng, dim);
let (_, vec_b) = random_inputs(&mut rng, dim);
let weights = random_weights(&mut rng, vec_a.len());
let scalar = score_4bit_internal_weighted_scalar(&vec_a, &vec_b, &weights);
let avx2 = unsafe { score_4bit_internal_weighted_avx2(&vec_a, &vec_b, &weights) };
assert_eq!(
scalar, avx2,
"weighted: scalar {scalar} != avx2 {avx2} at dim {dim}"
);
}
}
#[test]
fn test_score_weighted_saturation_safety_64k() {
let dim = 65_536;
let indices: Vec<u8> = vec![15; dim];
let vec_a = pack_codes(&indices, 4);
let vec_b = pack_codes(&indices, 4);
let max_weight: i16 = i16::MAX;
let weights: Vec<i16> = vec![max_weight; dim];
let scalar = score_4bit_internal_weighted_scalar(&vec_a, &vec_b, &weights);
unsafe {
if std::is_x86_feature_detected!("ssse3") && std::is_x86_feature_detected!("sse4.1") {
let sse = score_4bit_internal_weighted_sse(&vec_a, &vec_b, &weights);
assert_eq!(scalar, sse, "weighted score sse disagrees at dim={dim}");
}
if std::is_x86_feature_detected!("avx2") {
let avx2 = score_4bit_internal_weighted_avx2(&vec_a, &vec_b, &weights);
assert_eq!(scalar, avx2, "weighted score avx2 disagrees at dim={dim}");
}
}
}
#[test]
fn test_score_saturation_safety_64k() {
let dim = 65_536;
let indices: Vec<u8> = vec![15; dim]; let vec_a = pack_codes(&indices, 4);
let vec_b = pack_codes(&indices, 4);
let scalar = score_4bit_internal_scalar(&vec_a, &vec_b);
unsafe {
if std::is_x86_feature_detected!("ssse3") && std::is_x86_feature_detected!("sse4.1") {
let sse = score_4bit_internal_sse(&vec_a, &vec_b);
assert_eq!(scalar, sse, "score sse disagrees at dim={dim}");
}
if std::is_x86_feature_detected!("avx2") {
let avx2 = score_4bit_internal_avx2(&vec_a, &vec_b);
assert_eq!(scalar, avx2, "score avx2 disagrees at dim={dim}");
}
if std::is_x86_feature_detected!("avx512f")
&& std::is_x86_feature_detected!("avx512bw")
&& std::is_x86_feature_detected!("avx512vnni")
{
let vnni = score_4bit_internal_avx512_vnni(&vec_a, &vec_b);
assert_eq!(scalar, vnni, "score avx512_vnni disagrees at dim={dim}");
}
}
}
}