use std::arch::x86_64::*;
use super::avx2::{expand, lane_mask, tile8};
use super::Args;
use crate::lut::Lut;
use crate::MAX_DIM;
const NT: usize = 4;
#[inline(always)]
unsafe fn sum_weights(w: &[i8; MAX_DIM], dim: usize) -> i32 {
unsafe {
let ones = _mm256_set1_epi8(1);
let mut s = _mm256_setzero_si256();
let mut k = 0usize;
while k < dim {
let v = _mm256_loadu_si256(w.as_ptr().add(k) as *const __m256i);
s = _mm256_dpbusd_avx_epi32(s, ones, v);
k += 32;
}
let lo = _mm256_castsi256_si128(s);
let hi = _mm256_extracti128_si256(s, 1);
let mut tmp = [0i32; 4];
_mm_storeu_si128(tmp.as_mut_ptr() as *mut __m128i, _mm_add_epi32(lo, hi));
tmp.iter().sum()
}
}
#[inline(always)]
#[allow(clippy::too_many_arguments)]
unsafe fn block<const RT: usize, const N: usize>(
ws: &[[i8; MAX_DIM]; N],
corr: &[__m256i; 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 mut acc = [[zero; RT]; N];
for g in 0..d4n {
let mut q = [zero; RT];
for (rt, qv) in q.iter_mut().enumerate() {
*qv = _mm256_loadu_si256(tile_ptrs[rt].add(g * 64) as *const __m256i);
}
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 {
acc[j][rt] = _mm256_dpbusd_avx_epi32(acc[j][rt], q[rt], wb);
}
}
}
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(_mm256_sub_epi32(*acc_rt, corr[j]));
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)]
#[allow(clippy::too_many_arguments)]
unsafe fn tokens<const N: usize>(
ws: &[[i8; MAX_DIM]; N],
corr: &[__m256i; 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, corr, 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, corr, crows, invs, &tp, d4n, t0 * 8, nq, sqw, best);
}
}
}
#[target_feature(enable = "avx2,avxvnni")]
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 corr = [_mm256_setzero_si256(); 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],
);
corr[j] = _mm256_set1_epi32(128 * sum_weights(&ws[j], dim));
crows[j] = a.crow(tt);
invs[j] = a.inv(tt);
}
tokens::<NT>(
&ws,
&corr,
&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_c = [_mm256_set1_epi32(128 * sum_weights(&ws[0], dim))];
let one_r = [a.crow(t)];
let one_i = [a.inv(t)];
tokens::<1>(
&one_w,
&one_c,
&one_r,
&one_i,
tiles,
tile_stride,
d4n,
n8,
nq,
sqw,
best,
);
t += 1;
}
}
best.iter().sum()
}