use std::arch::aarch64::*;
use super::neon::{expand, fold4, fold_tail};
use super::Args;
use crate::lut::Lut;
use crate::MAX_DIM;
const RP: usize = 4;
const TP: usize = 2;
const PAIR_BYTES: usize = 2 * MAX_DIM;
#[inline(always)]
unsafe fn smmla(acc: int32x4_t, a: int8x16_t, b: int8x16_t) -> int32x4_t {
let out: int32x4_t;
unsafe {
std::arch::asm!(
"smmla {out:v}.4s, {a:v}.16b, {b:v}.16b",
out = inout(vreg) acc => out,
a = in(vreg) a,
b = in(vreg) b,
options(pure, nomem, nostack),
);
}
out
}
#[inline(always)]
unsafe fn interleave(w0: &[i8; MAX_DIM], w1: &[i8; MAX_DIM], dim: usize, out: &mut [i8; PAIR_BYTES]) {
unsafe {
let mut k = 0usize;
while k < dim {
let a = vld1q_s8(w0.as_ptr().add(k));
let b = vld1q_s8(w1.as_ptr().add(k));
let o = out.as_mut_ptr().add(2 * k);
vst1q_s8(o, vcombine_s8(vget_low_s8(a), vget_low_s8(b)));
vst1q_s8(o.add(16), vcombine_s8(vget_high_s8(a), vget_high_s8(b)));
k += 16;
}
}
}
#[inline(always)]
#[allow(clippy::too_many_arguments)]
#[allow(clippy::needless_range_loop)]
unsafe fn block<const R: usize, const N: usize>(
wp: &[[i8; PAIR_BYTES]; N],
n_valid: usize,
d8n: usize,
qbase: *const i8,
pair_stride: usize,
row0: usize,
nq: usize,
sqw: &[f32],
crows: &[*const f32],
invs: &[f32],
best: &mut [f32],
) {
unsafe {
let zero = vdupq_n_s32(0);
let mut acc = [[zero; N]; R];
for g in 0..d8n {
let mut q = [vdupq_n_s8(0); R];
for (rp, qv) in q.iter_mut().enumerate() {
*qv = vld1q_s8(qbase.add(rp * pair_stride + g * 16));
}
let mut w = [vdupq_n_s8(0); N];
for (tp, wv) in w.iter_mut().enumerate() {
*wv = vld1q_s8(wp[tp].as_ptr().add(g * 16));
}
for rp in 0..R {
for tp in 0..N {
acc[rp][tp] = smmla(acc[rp][tp], q[rp], w[tp]);
}
}
}
for tp in 0..N {
for half in 0..2 {
let t = 2 * tp + half;
if t >= n_valid {
break;
}
let (crow, inv) = (crows[t], invs[t]);
let mut rp = 0usize;
while rp < R {
let a = acc[rp][tp];
let paired = rp + 1 < R;
let b = if paired { acc[rp + 1][tp] } else { zero };
let v = if half == 0 {
vuzp1q_s32(a, b)
} else {
vuzp2q_s32(a, b)
};
let base = row0 + 2 * rp;
let take = (nq - base).min(if paired { 4 } else { 2 });
if take == 4 {
fold4(v, base, sqw, crow, inv, best);
} else {
let mut lanes = [0i32; 4];
vst1q_s32(lanes.as_mut_ptr(), v);
fold_tail(
&lanes[..take],
&sqw[base..],
crow.add(base),
inv,
&mut best[base..],
);
}
rp += 2;
}
}
}
}
}
#[inline(always)]
#[allow(clippy::too_many_arguments)]
unsafe fn tokens<const N: usize>(
wp: &[[i8; PAIR_BYTES]; N],
n_valid: usize,
d8n: usize,
pairs: *const i8,
pair_stride: usize,
nq: usize,
sqw: &[f32],
crows: &[*const f32],
invs: &[f32],
best: &mut [f32],
) {
unsafe {
let npairs = nq.div_ceil(2);
let mut p0 = 0usize;
while npairs - p0 >= RP {
let qb = pairs.add(p0 * pair_stride);
block::<RP, N>(
wp,
n_valid,
d8n,
qb,
pair_stride,
2 * p0,
nq,
sqw,
crows,
invs,
best,
);
p0 += RP;
}
let qb = pairs.add(p0 * pair_stride);
match npairs - p0 {
3 => block::<3, N>(
wp,
n_valid,
d8n,
qb,
pair_stride,
2 * p0,
nq,
sqw,
crows,
invs,
best,
),
2 => block::<2, N>(
wp,
n_valid,
d8n,
qb,
pair_stride,
2 * p0,
nq,
sqw,
crows,
invs,
best,
),
1 => block::<1, N>(
wp,
n_valid,
d8n,
qb,
pair_stride,
2 * p0,
nq,
sqw,
crows,
invs,
best,
),
_ => {}
}
}
}
#[target_feature(enable = "i8mm")]
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 d8n = dim / 8;
let pair_stride = 2 * q.stride();
let pairs = q.pairs().as_ptr();
debug_assert!(q.pairs().len() >= nq.div_ceil(2) * pair_stride);
let sqw = q.sqw();
best.clear();
best.resize(nq, f32::NEG_INFINITY);
unsafe {
let mut tabs = [vdupq_n_s8(0); 8];
for (tab, src) in tabs.iter_mut().zip(nib.tables.iter()).take(kpb) {
*tab = vld1q_s8(src.as_ptr());
}
const NT: usize = 2 * TP;
let zero = [0i8; MAX_DIM];
let mut ws = [[0i8; MAX_DIM]; NT];
let mut wp = [[0i8; PAIR_BYTES]; TP];
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);
}
for tp in 0..TP {
interleave(&ws[2 * tp], &ws[2 * tp + 1], dim, &mut wp[tp]);
}
tokens::<TP>(&wp, NT, d8n, pairs, pair_stride, nq, sqw, &crows, &invs, best);
t += NT;
}
while t < a.n_tokens {
let n = (a.n_tokens - t).min(2);
for j in 0..n {
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);
}
let partner = if n == 2 { &ws[1] } else { &zero };
interleave(&ws[0], partner, dim, &mut wp[0]);
let one = [wp[0]];
tokens::<1>(
&one,
n,
d8n,
pairs,
pair_stride,
nq,
sqw,
&crows[..n],
&invs[..n],
best,
);
t += n;
}
}
best.iter().sum()
}