use std::arch::x86_64::*;
use super::Args;
use crate::lut::{Lut, NibbleTables};
use crate::MAX_DIM;
const NT: usize = 4;
#[inline(always)]
pub(super) unsafe fn expand(
row: &[u8],
tabs: &[__m128i; 8],
nib: &NibbleTables,
kpb: usize,
pdim: usize,
w: &mut [i8; MAX_DIM],
) {
unsafe {
let low_mask = _mm_set1_epi8(0x0F);
let wp = w.as_mut_ptr();
let mut i = 0usize;
while i + 16 <= pdim {
let v = _mm_loadu_si128(row.as_ptr().add(i) as *const __m128i);
let hi = _mm_and_si128(_mm_srli_epi16(v, 4), low_mask);
let lo = _mm_and_si128(v, low_mask);
for (k, tab) in tabs.iter().enumerate().take(kpb) {
let idx = if nib.from_hi[k] { hi } else { lo };
_mm_storeu_si128(wp.add(k * pdim + i) as *mut __m128i, _mm_shuffle_epi8(*tab, idx));
}
i += 16;
}
if i < pdim {
let rem = pdim - i;
let mut src = [0u8; 16];
src[..rem].copy_from_slice(&row[i..pdim]);
let v = _mm_loadu_si128(src.as_ptr() as *const __m128i);
let hi = _mm_and_si128(_mm_srli_epi16(v, 4), low_mask);
let lo = _mm_and_si128(v, low_mask);
let mut dst = [0i8; 16];
for k in 0..kpb {
let idx = if nib.from_hi[k] { hi } else { lo };
_mm_storeu_si128(dst.as_mut_ptr() as *mut __m128i, _mm_shuffle_epi8(tabs[k], idx));
w[k * pdim + i..k * pdim + pdim].copy_from_slice(&dst[..rem]);
}
}
}
}
#[inline(always)]
pub(super) unsafe fn lane_mask(rem: usize) -> __m256i {
unsafe {
let idx = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7);
_mm256_cmpgt_epi32(_mm256_set1_epi32(rem as i32), idx)
}
}
#[inline(always)]
#[allow(clippy::too_many_arguments)]
unsafe fn block<const RT: usize, const N: usize>(
ws: &[[i8; MAX_DIM]; N],
crows: &[*const f32; N],
invs: &[f32; N],
tile_ptrs: &[*const u8; RT],
d4n: usize,
row0: usize,
nq: usize,
sqw: &[f32],
best: &mut [f32],
) {
unsafe {
let zero = _mm256_setzero_si256();
let ones = _mm256_set1_epi16(1);
let flip = _mm256_set1_epi8(-128);
let mut acc = [[zero; RT]; N];
for g in 0..d4n {
let mut qs = [zero; RT];
let mut qa = [zero; RT];
for rt in 0..RT {
let qu = _mm256_loadu_si256(tile_ptrs[rt].add(g * 64) as *const __m256i);
qs[rt] = _mm256_xor_si256(qu, flip);
qa[rt] = _mm256_abs_epi8(qs[rt]);
}
for j in 0..N {
let wb = _mm256_set1_epi32((ws[j].as_ptr().add(g * 4) as *const i32).read_unaligned());
for rt in 0..RT {
let prod = _mm256_maddubs_epi16(qa[rt], _mm256_sign_epi8(wb, qs[rt]));
acc[j][rt] = _mm256_add_epi32(acc[j][rt], _mm256_madd_epi16(prod, ones));
}
}
}
for j in 0..N {
let invv = _mm256_set1_ps(invs[j]);
for (rt, acc_rt) in acc[j].iter().enumerate() {
let r0 = row0 + rt * 8;
let rem = (nq - r0).min(8);
let m = lane_mask(rem);
let a = _mm256_cvtepi32_ps(*acc_rt);
let s = _mm256_mul_ps(
_mm256_add_ps(
_mm256_mul_ps(_mm256_maskload_ps(sqw.as_ptr().add(r0), m), a),
_mm256_maskload_ps(crows[j].add(r0), m),
),
invv,
);
let b = _mm256_maskload_ps(best.as_ptr().add(r0), m);
_mm256_maskstore_ps(best.as_mut_ptr().add(r0), m, _mm256_max_ps(b, s));
}
}
}
}
#[inline(always)]
pub(super) unsafe fn tile8(tiles: *const u8, tile_stride: usize, r8: usize) -> *const u8 {
unsafe { tiles.add((r8 / 2) * tile_stride + (r8 % 2) * 32) }
}
#[inline(always)]
#[allow(clippy::too_many_arguments)]
unsafe fn tokens<const N: usize>(
ws: &[[i8; MAX_DIM]; N],
crows: &[*const f32; N],
invs: &[f32; N],
tiles: *const u8,
tile_stride: usize,
d4n: usize,
n8: usize,
nq: usize,
sqw: &[f32],
best: &mut [f32],
) {
unsafe {
let mut t0 = 0usize;
while n8 - t0 >= 2 {
let tp = [tile8(tiles, tile_stride, t0), tile8(tiles, tile_stride, t0 + 1)];
block::<2, N>(ws, crows, invs, &tp, d4n, t0 * 8, nq, sqw, best);
t0 += 2;
}
if n8 - t0 == 1 {
let tp = [tile8(tiles, tile_stride, t0)];
block::<1, N>(ws, crows, invs, &tp, d4n, t0 * 8, nq, sqw, best);
}
}
}
#[target_feature(enable = "avx2")]
pub(super) unsafe fn maxsim(lut: &Lut, a: &Args<'_>, best: &mut Vec<f32>, _accs: &mut Vec<i32>) -> f32 {
let q = a.query;
let nq = q.n_tokens();
let dim = q.dim();
if nq == 0 || a.n_tokens == 0 {
return 0.0;
}
let nib = lut.nibble_tables().expect("select() guarantees nibble tables");
let kpb = lut.keys_per_byte();
let pdim = dim / kpb;
let d4n = dim / 4;
let n8 = nq.div_ceil(8);
let tile_stride = (q.stride() / 4) * 64;
let tiles = q.tiles_u8().as_ptr();
debug_assert!(q.tiles_u8().len() >= nq.div_ceil(16) * tile_stride);
let sqw = q.sqw();
best.clear();
best.resize(nq, f32::NEG_INFINITY);
unsafe {
let mut tabs = [_mm_setzero_si128(); 8];
for (tab, src) in tabs.iter_mut().zip(nib.tables.iter()).take(kpb) {
*tab = _mm_loadu_si128(src.as_ptr() as *const __m128i);
}
let mut ws = [[0i8; MAX_DIM]; NT];
let mut crows = [std::ptr::null::<f32>(); NT];
let mut invs = [0f32; NT];
let mut t = 0usize;
while t + NT <= a.n_tokens {
for j in 0..NT {
let tt = t + j;
expand(
&a.packed[tt * a.row_stride..tt * a.row_stride + pdim],
&tabs,
nib,
kpb,
pdim,
&mut ws[j],
);
crows[j] = a.crow(tt);
invs[j] = a.inv(tt);
}
tokens::<NT>(&ws, &crows, &invs, tiles, tile_stride, d4n, n8, nq, sqw, best);
t += NT;
}
while t < a.n_tokens {
expand(
&a.packed[t * a.row_stride..t * a.row_stride + pdim],
&tabs,
nib,
kpb,
pdim,
&mut ws[0],
);
let one_w = [ws[0]];
let one_r = [a.crow(t)];
let one_i = [a.inv(t)];
tokens::<1>(&one_w, &one_r, &one_i, tiles, tile_stride, d4n, n8, nq, sqw, best);
t += 1;
}
}
best.iter().sum()
}