#![allow(unsafe_code, unsafe_op_in_unsafe_fn)]
use super::scalar;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ArmTier {
Smmla,
Sdot,
None,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub(crate) enum DenseI8Route {
Autovec,
Scalar,
Sdot,
Smmla,
}
#[must_use]
fn autovec_preferred() -> bool {
use std::sync::OnceLock;
static PREF: OnceLock<bool> = OnceLock::new();
*PREF.get_or_init(|| {
cfg!(target_vendor = "apple")
&& !matches!(
std::env::var("FOCR_INT8_AUTOVEC").ok().as_deref(),
Some("0") | Some("off") | Some("false") | Some("no")
)
&& !force_arch_honored()
})
}
fn force_arch_honored() -> bool {
let Ok(force) = std::env::var("FOCR_FORCE_ARCH") else {
return false;
};
matches!(
(force.trim().to_ascii_lowercase().as_str(), detect_tier()),
("smmla", ArmTier::Smmla) | ("sdot", ArmTier::Sdot) | ("scalar", ArmTier::None)
)
}
#[must_use]
pub(crate) fn effective_dense_route() -> DenseI8Route {
match detect_tier() {
ArmTier::Smmla => DenseI8Route::Smmla,
ArmTier::Sdot if autovec_preferred() => DenseI8Route::Autovec,
ArmTier::Sdot => DenseI8Route::Sdot,
ArmTier::None => DenseI8Route::Scalar,
}
}
pub fn detect_tier() -> ArmTier {
use std::sync::OnceLock;
static TIER: OnceLock<ArmTier> = OnceLock::new();
*TIER.get_or_init(detect_tier_uncached)
}
fn detect_tier_uncached() -> ArmTier {
#[cfg(target_arch = "aarch64")]
{
let has_i8mm = std::arch::is_aarch64_feature_detected!("i8mm");
let has_dotprod = std::arch::is_aarch64_feature_detected!("dotprod");
if let Ok(force) = std::env::var("FOCR_FORCE_ARCH") {
match force.trim().to_ascii_lowercase().as_str() {
"smmla" if has_i8mm => return ArmTier::Smmla,
"sdot" if has_dotprod => return ArmTier::Sdot,
"scalar" => return ArmTier::None,
_ => {}
}
}
#[cfg(target_vendor = "apple")]
{
if has_dotprod {
return ArmTier::Sdot;
}
if has_i8mm {
return ArmTier::Smmla;
}
}
#[cfg(not(target_vendor = "apple"))]
{
if has_i8mm {
return ArmTier::Smmla;
}
if has_dotprod {
return ArmTier::Sdot;
}
}
ArmTier::None
}
#[cfg(not(target_arch = "aarch64"))]
{
ArmTier::None
}
}
pub fn igemm_s8s8(a: &[i8], b: &[i8], m: usize, k: usize, n: usize, out: &mut [i32]) {
let _ = igemm_s8s8_with_route(a, b, m, k, n, out);
}
pub(crate) fn igemm_s8s8_with_route(
a: &[i8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) -> DenseI8Route {
let a_len = scalar::checked_len("igemm_s8s8", m, k, "m*k");
let b_len = scalar::checked_len("igemm_s8s8", n, k, "n*k");
let out_len = scalar::checked_len("igemm_s8s8", m, n, "m*n");
assert_eq!(
a.len(),
a_len,
"igemm_s8s8: a.len {} != m*k {}",
a.len(),
a_len
);
assert_eq!(
b.len(),
b_len,
"igemm_s8s8: b.len {} != n*k {}",
b.len(),
b_len
);
assert_eq!(
out.len(),
out_len,
"igemm_s8s8: out.len {} != m*n {}",
out.len(),
out_len
);
#[cfg(target_arch = "aarch64")]
{
match detect_tier() {
ArmTier::Smmla => {
unsafe { aarch64_impl::igemm_s8s8_smmla(a, b, m, k, n, out) };
return DenseI8Route::Smmla;
}
ArmTier::Sdot => {
if !autovec_preferred() {
unsafe { aarch64_impl::igemm_s8s8_sdot(a, b, m, k, n, out) };
return DenseI8Route::Sdot;
}
}
ArmTier::None => {}
}
}
scalar::igemm_s8s8(a, b, m, k, n, out);
match detect_tier() {
ArmTier::Sdot => DenseI8Route::Autovec,
ArmTier::None => DenseI8Route::Scalar,
ArmTier::Smmla => unreachable!("SMMLA branch must return before scalar fallback"),
}
}
pub fn igemm_s8s8_packed_b(
a: &[i8],
b_panels: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
let a_len = scalar::checked_len("igemm_s8s8_packed_b", m, k, "m*k");
let panels_len = crate::simd::pack::smmla_packed_len(n, k);
let out_len = scalar::checked_len("igemm_s8s8_packed_b", m, n, "m*n");
assert_eq!(
a.len(),
a_len,
"igemm_s8s8_packed_b: a.len {} != m*k {}",
a.len(),
a_len
);
assert_eq!(
b_panels.len(),
panels_len,
"igemm_s8s8_packed_b: b_panels.len {} != ceil(n/2)*ceil(k/8)*16 {}",
b_panels.len(),
panels_len
);
assert_eq!(
out.len(),
out_len,
"igemm_s8s8_packed_b: out.len {} != m*n {}",
out.len(),
out_len
);
out.fill(0);
#[cfg(target_arch = "aarch64")]
if detect_tier() == ArmTier::Smmla {
unsafe { aarch64_impl::igemm_s8s8_smmla_packed_b(a, b_panels, m, k, n, out) };
return;
}
let b = crate::simd::pack::smmla_unpack_panels(b_panels, n, k).expect("length asserted above");
igemm_s8s8(a, &b, m, k, n, out);
}
pub fn igemm_u8s8(a: &[u8], b: &[i8], m: usize, k: usize, n: usize, out: &mut [i32]) {
let _ = igemm_u8s8_with_route(a, b, m, k, n, out);
}
pub(crate) fn igemm_u8s8_with_route(
a: &[u8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) -> DenseI8Route {
let a_len = scalar::checked_len("igemm_u8s8", m, k, "m*k");
let b_len = scalar::checked_len("igemm_u8s8", n, k, "n*k");
let out_len = scalar::checked_len("igemm_u8s8", m, n, "m*n");
assert_eq!(
a.len(),
a_len,
"igemm_u8s8: a.len {} != m*k {}",
a.len(),
a_len
);
assert_eq!(
b.len(),
b_len,
"igemm_u8s8: b.len {} != n*k {}",
b.len(),
b_len
);
assert_eq!(
out.len(),
out_len,
"igemm_u8s8: out.len {} != m*n {}",
out.len(),
out_len
);
#[cfg(target_arch = "aarch64")]
{
let tier = detect_tier();
if matches!(tier, ArmTier::Smmla) || (matches!(tier, ArmTier::Sdot) && !autovec_preferred())
{
let a_signed: Vec<i8> = a.iter().map(|&x| x.wrapping_sub(128) as i8).collect();
match tier {
ArmTier::Smmla => {
unsafe { aarch64_impl::igemm_s8s8_smmla(&a_signed, b, m, k, n, out) };
}
ArmTier::Sdot => {
unsafe { aarch64_impl::igemm_s8s8_sdot(&a_signed, b, m, k, n, out) };
}
ArmTier::None => unreachable!(),
}
let mut rowsum = vec![0i32; n];
for (oc, rs) in rowsum.iter_mut().enumerate() {
let row = &b[oc * k..(oc + 1) * k];
let mut s: i32 = 0;
for &w in row {
s += i32::from(w);
}
*rs = s;
}
for r in 0..m {
let orow = &mut out[r * n..(r + 1) * n];
for (c, cell) in orow.iter_mut().enumerate() {
*cell += 128 * rowsum[c];
}
}
return match tier {
ArmTier::Smmla => DenseI8Route::Smmla,
ArmTier::Sdot => DenseI8Route::Sdot,
ArmTier::None => unreachable!(),
};
}
}
scalar::igemm_u8s8(a, b, m, k, n, out);
match detect_tier() {
ArmTier::Sdot => DenseI8Route::Autovec,
ArmTier::None => DenseI8Route::Scalar,
ArmTier::Smmla => unreachable!("SMMLA branch must return before scalar fallback"),
}
}
#[allow(clippy::too_many_arguments)]
pub fn igemm_s4s8_packed(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
out: &mut [f32],
) {
match detect_tier() {
ArmTier::Sdot => {
unsafe {
aarch64_impl::igemm_s4s8_packed_sdot(a, b_packed, scales, group, m, k, n, out);
}
}
ArmTier::Smmla => {
unsafe {
aarch64_impl::igemm_s4s8_packed_smmla(a, b_packed, scales, group, m, k, n, out);
}
}
ArmTier::None => {
super::int4::s4s8_packed_kernel_scalar(a, b_packed, scales, group, m, k, n, out);
}
}
}
#[cfg(target_arch = "aarch64")]
#[allow(unsafe_code, unsafe_op_in_unsafe_fn)]
mod aarch64_impl {
use core::arch::aarch64::{
int8x8_t, int8x16_t, int32x4_t, vaddvq_s32, vand_u8, vcombine_s8, vdotq_s32, vdup_n_u8,
vdupq_n_s32, vld1_u8, vld1q_s8, vld2_s8, vmmlaq_s32, vreinterpret_s8_u8, vshl_n_s8,
vshr_n_s8, vshr_n_u8, vzip_s8,
};
#[inline]
#[target_feature(enable = "neon")]
unsafe fn load16(s: &[i8], off: usize) -> int8x16_t {
debug_assert!(off + 16 <= s.len(), "load16 out of bounds");
vld1q_s8(s.as_ptr().add(off))
}
#[inline]
#[target_feature(enable = "neon")]
unsafe fn load16_tail(s: &[i8], off: usize, valid: usize) -> int8x16_t {
let mut buf = [0i8; 16];
let n = valid.min(16);
buf[..n].copy_from_slice(&s[off..off + n]);
vld1q_s8(buf.as_ptr())
}
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn igemm_s8s8_sdot(
a: &[i8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
let k16 = k / 16 * 16;
let mut r = 0;
while r < m {
let mr = (m - r).min(4);
let mut c = 0;
while c < n {
let nr = (n - c).min(4);
let mut acc = [[vdupq_n_s32(0); 4]; 4];
let mut t = 0;
while t < k16 {
let mut av = [vdupq_n_s32(0).into_i8(); 4];
let mut bv = [vdupq_n_s32(0).into_i8(); 4];
#[allow(clippy::needless_range_loop)]
for i in 0..mr {
av[i] = load16(a, (r + i) * k + t);
}
#[allow(clippy::needless_range_loop)]
for j in 0..nr {
bv[j] = load16(b, (c + j) * k + t);
}
for i in 0..mr {
for j in 0..nr {
acc[i][j] = vdotq_s32(acc[i][j], av[i], bv[j]);
}
}
t += 16;
}
if k16 < k {
let valid = k - k16;
let mut av = [vdupq_n_s32(0).into_i8(); 4];
let mut bv = [vdupq_n_s32(0).into_i8(); 4];
#[allow(clippy::needless_range_loop)]
for i in 0..mr {
av[i] = load16_tail(a, (r + i) * k + k16, valid);
}
#[allow(clippy::needless_range_loop)]
for j in 0..nr {
bv[j] = load16_tail(b, (c + j) * k + k16, valid);
}
for i in 0..mr {
for j in 0..nr {
acc[i][j] = vdotq_s32(acc[i][j], av[i], bv[j]);
}
}
}
for i in 0..mr {
for j in 0..nr {
out[(r + i) * n + (c + j)] += vaddvq_s32(acc[i][j]);
}
}
c += nr;
}
r += mr;
}
}
fn pack_panels(
src: &[i8],
base_row: usize,
rows: usize,
k: usize,
src_k: usize,
) -> (Vec<i8>, usize, usize) {
crate::simd::pack::smmla_pack_panels(src, base_row, rows, k, src_k)
}
#[target_feature(enable = "neon,i8mm")]
pub unsafe fn igemm_s8s8_smmla(
a: &[i8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
let mut r = 0;
while r < m {
let mr = (m - r).min(8);
let (apack, a_pairs, kb) = pack_panels(a, r, mr, k, k);
let mut c = 0;
while c < n {
let nr = (n - c).min(8);
let (bpack, b_pairs, _kb2) = pack_panels(b, c, nr, k, k);
for p in 0..a_pairs {
for q in 0..b_pairs {
let mut acc = vdupq_n_s32(0);
for block in 0..kb {
let aoff = (p * kb + block) * 16;
let boff = (q * kb + block) * 16;
let av = load16(&apack, aoff);
let bv = load16(&bpack, boff);
acc = vmmlaq_s32(acc, av, bv);
}
let tile = [
vgetq_lane0(acc),
vgetq_lane1(acc),
vgetq_lane2(acc),
vgetq_lane3(acc),
];
let ar0 = 2 * p;
let ar1 = 2 * p + 1;
let bc0 = 2 * q;
let bc1 = 2 * q + 1;
if ar0 < mr && bc0 < nr {
out[(r + ar0) * n + (c + bc0)] += tile[0];
}
if ar0 < mr && bc1 < nr {
out[(r + ar0) * n + (c + bc1)] += tile[1];
}
if ar1 < mr && bc0 < nr {
out[(r + ar1) * n + (c + bc0)] += tile[2];
}
if ar1 < mr && bc1 < nr {
out[(r + ar1) * n + (c + bc1)] += tile[3];
}
}
}
c += nr;
}
r += mr;
}
}
#[target_feature(enable = "neon,i8mm")]
pub unsafe fn igemm_s8s8_smmla_packed_b(
a: &[i8],
b_panels: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
let kb = k.div_ceil(8);
let mut r = 0;
while r < m {
let mr = (m - r).min(8);
let (apack, a_pairs, _kb) = pack_panels(a, r, mr, k, k);
let mut c = 0;
while c < n {
let nr = (n - c).min(8);
let b_base = (c / 2) * kb;
let b_pairs = nr.div_ceil(2);
for p in 0..a_pairs {
for q in 0..b_pairs {
let mut acc = vdupq_n_s32(0);
for block in 0..kb {
let aoff = (p * kb + block) * 16;
let boff = ((b_base + q * kb) + block) * 16;
let av = load16(&apack, aoff);
let bv = load16(b_panels, boff);
acc = vmmlaq_s32(acc, av, bv);
}
let tile = [
vgetq_lane0(acc),
vgetq_lane1(acc),
vgetq_lane2(acc),
vgetq_lane3(acc),
];
let ar0 = 2 * p;
let ar1 = 2 * p + 1;
let bc0 = 2 * q;
let bc1 = 2 * q + 1;
if ar0 < mr && bc0 < nr {
out[(r + ar0) * n + (c + bc0)] += tile[0];
}
if ar0 < mr && bc1 < nr {
out[(r + ar0) * n + (c + bc1)] += tile[1];
}
if ar1 < mr && bc0 < nr {
out[(r + ar1) * n + (c + bc0)] += tile[2];
}
if ar1 < mr && bc1 < nr {
out[(r + ar1) * n + (c + bc1)] += tile[3];
}
}
}
c += nr;
}
r += mr;
}
}
#[inline]
#[target_feature(enable = "neon")]
fn vgetq_lane0(v: int32x4_t) -> i32 {
core::arch::aarch64::vgetq_lane_s32::<0>(v)
}
#[inline]
#[target_feature(enable = "neon")]
fn vgetq_lane1(v: int32x4_t) -> i32 {
core::arch::aarch64::vgetq_lane_s32::<1>(v)
}
#[inline]
#[target_feature(enable = "neon")]
fn vgetq_lane2(v: int32x4_t) -> i32 {
core::arch::aarch64::vgetq_lane_s32::<2>(v)
}
#[inline]
#[target_feature(enable = "neon")]
fn vgetq_lane3(v: int32x4_t) -> i32 {
core::arch::aarch64::vgetq_lane_s32::<3>(v)
}
trait IntoI8 {
fn into_i8(self) -> int8x16_t;
}
impl IntoI8 for int32x4_t {
#[inline]
fn into_i8(self) -> int8x16_t {
unsafe { core::arch::aarch64::vreinterpretq_s8_s32(self) }
}
}
#[inline]
#[target_feature(enable = "neon")]
unsafe fn unpack8(wptr: *const u8) -> (int8x8_t, int8x8_t) {
let v = vld1_u8(wptr);
let lo_u = vand_u8(v, vdup_n_u8(0x0F));
let lo_s = vshr_n_s8(vshl_n_s8(vreinterpret_s8_u8(lo_u), 4), 4);
let hi_u = vshr_n_u8(v, 4);
let hi_s = vshr_n_s8(vshl_n_s8(vreinterpret_s8_u8(hi_u), 4), 4);
(lo_s, hi_s)
}
#[allow(clippy::too_many_arguments)]
#[target_feature(enable = "neon,dotprod")]
pub unsafe fn igemm_s4s8_packed_sdot(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
out: &mut [f32],
) {
let groups = k / group;
let sub_per_group = group / 16; let kbytes = k / 2;
let mut r = 0;
while r < m {
let mr = (m - r).min(4);
let mut c = 0;
while c < n {
let nr = (n - c).min(4);
let mut out_acc = [[0.0f32; 4]; 4];
for g in 0..groups {
let mut acc = [[vdupq_n_s32(0); 4]; 4];
for sub in 0..sub_per_group {
let sb = g * sub_per_group + sub;
let mut nat = [vdupq_n_s32(0).into_i8(); 4];
#[allow(clippy::needless_range_loop)]
for j in 0..nr {
let (lo, hi) =
unpack8(b_packed.as_ptr().add((c + j) * kbytes + sb * 8));
let zz = vzip_s8(lo, hi);
nat[j] = vcombine_s8(zz.0, zz.1);
}
let mut act = [vdupq_n_s32(0).into_i8(); 4];
#[allow(clippy::needless_range_loop)]
for i in 0..mr {
act[i] = vld1q_s8(a.as_ptr().add((r + i) * k + sb * 16));
}
for i in 0..mr {
for j in 0..nr {
acc[i][j] = vdotq_s32(acc[i][j], act[i], nat[j]);
}
}
}
for i in 0..mr {
for j in 0..nr {
let gi = vaddvq_s32(acc[i][j]);
out_acc[i][j] += scales[(c + j) * groups + g] * gi as f32;
}
}
}
for i in 0..mr {
for j in 0..nr {
out[(r + i) * n + (c + j)] = out_acc[i][j];
}
}
c += nr;
}
r += mr;
}
}
#[allow(clippy::too_many_arguments)]
#[target_feature(enable = "neon,i8mm")]
pub unsafe fn igemm_s4s8_packed_smmla(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
out: &mut [f32],
) {
let groups = k / group;
let sub_per_group = group / 16;
let kbytes = k / 2;
let zero8 = vreinterpret_s8_u8(vdup_n_u8(0));
let mut rp = 0;
while rp < m {
let mut cp = 0;
while cp < n {
let mut lane = [0.0f32; 4];
for g in 0..groups {
let mut acc = vdupq_n_s32(0);
for sub in 0..sub_per_group {
let sb = g * sub_per_group + sub;
let (lo0, hi0) = unpack8(b_packed.as_ptr().add(cp * kbytes + sb * 8));
let (lo1, hi1) = if cp + 1 < n {
unpack8(b_packed.as_ptr().add((cp + 1) * kbytes + sb * 8))
} else {
(zero8, zero8)
};
let t0 = vld2_s8(a.as_ptr().add(rp * k + sb * 16));
let (ev1, od1) = if rp + 1 < m {
let t1 = vld2_s8(a.as_ptr().add((rp + 1) * k + sb * 16));
(t1.0, t1.1)
} else {
(zero8, zero8)
};
acc = vmmlaq_s32(acc, vcombine_s8(t0.0, ev1), vcombine_s8(lo0, lo1));
acc = vmmlaq_s32(acc, vcombine_s8(t0.1, od1), vcombine_s8(hi0, hi1));
}
let s_c0 = scales[cp * groups + g];
let s_c1 = if cp + 1 < n {
scales[(cp + 1) * groups + g]
} else {
0.0
};
lane[0] += s_c0 * vgetq_lane0(acc) as f32;
lane[1] += s_c1 * vgetq_lane1(acc) as f32;
lane[2] += s_c0 * vgetq_lane2(acc) as f32;
lane[3] += s_c1 * vgetq_lane3(acc) as f32;
}
out[rp * n + cp] = lane[0];
if cp + 1 < n {
out[rp * n + cp + 1] = lane[1];
}
if rp + 1 < m {
out[(rp + 1) * n + cp] = lane[2];
}
if rp + 1 < m && cp + 1 < n {
out[(rp + 1) * n + cp + 1] = lane[3];
}
cp += 2;
}
rp += 2;
}
}
}
#[cfg(all(test, target_arch = "aarch64"))]
mod tests {
use super::*;
fn oracle_s8s8(a: &[i8], b: &[i8], m: usize, k: usize, n: usize) -> Vec<i32> {
let mut out = vec![0i32; m * n];
for r in 0..m {
for col in 0..n {
let mut acc: i32 = 0;
for t in 0..k {
acc += i32::from(a[r * k + t]) * i32::from(b[col * k + t]);
}
out[r * n + col] = acc;
}
}
out
}
fn oracle_u8s8(a: &[u8], b: &[i8], m: usize, k: usize, n: usize) -> Vec<i32> {
let mut out = vec![0i32; m * n];
for r in 0..m {
for col in 0..n {
let mut acc: i32 = 0;
for t in 0..k {
acc += i32::from(a[r * k + t]) * i32::from(b[col * k + t]);
}
out[r * n + col] = acc;
}
}
out
}
struct Rng(u64);
impl Rng {
fn next(&mut self) -> u64 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
x
}
fn i8(&mut self) -> i8 {
(self.next() & 0xff) as u8 as i8
}
fn u8(&mut self) -> u8 {
(self.next() & 0xff) as u8
}
}
fn rand_i8(rng: &mut Rng, len: usize) -> Vec<i8> {
(0..len).map(|_| rng.i8()).collect()
}
fn rand_u8(rng: &mut Rng, len: usize) -> Vec<u8> {
(0..len).map(|_| rng.u8()).collect()
}
fn run_sdot_s8s8(a: &[i8], b: &[i8], m: usize, k: usize, n: usize) -> Option<Vec<i32>> {
if !std::arch::is_aarch64_feature_detected!("dotprod") {
return None;
}
let mut out = vec![0i32; m * n];
unsafe { aarch64_impl::igemm_s8s8_sdot(a, b, m, k, n, &mut out) };
Some(out)
}
fn run_smmla_s8s8(a: &[i8], b: &[i8], m: usize, k: usize, n: usize) -> Option<Vec<i32>> {
if !std::arch::is_aarch64_feature_detected!("i8mm") {
return None;
}
let mut out = vec![0i32; m * n];
unsafe { aarch64_impl::igemm_s8s8_smmla(a, b, m, k, n, &mut out) };
Some(out)
}
#[test]
fn smmla_packed_b_matches_row_major_and_oracle() {
if !std::arch::is_aarch64_feature_detected!("i8mm") {
eprintln!(r#"{{"check":"smmla_packed_b_parity","result":"skip","reason":"no i8mm"}}"#);
return;
}
let mut rng = Rng(0x0dd0_beef_cafe_f00d);
let shapes = [
(1usize, 16usize, 8usize),
(1, 17, 5),
(3, 5, 7),
(4, 64, 8),
(8, 96, 96),
(1, 1280, 64),
(1, 6848, 4),
];
for &(m, k, n) in &shapes {
let mut a = vec![0i8; m * k];
let mut b = vec![0i8; n * k];
if (m, k, n) == (1, 6848, 4) {
a.fill(i8::MAX);
b.fill(i8::MIN);
} else {
for v in a.iter_mut() {
*v = rng.i8();
}
for v in b.iter_mut() {
*v = rng.i8();
}
}
let (panels, _, _) = crate::simd::pack::smmla_pack_panels(&b, 0, n, k, k);
let mut packed_out = vec![0i32; m * n];
unsafe {
aarch64_impl::igemm_s8s8_smmla_packed_b(&a, &panels, m, k, n, &mut packed_out);
};
let row_major = run_smmla_s8s8(&a, &b, m, k, n).expect("i8mm present");
let mut oracle = vec![0i32; m * n];
scalar::igemm_s8s8(&a, &b, m, k, n, &mut oracle);
assert_eq!(
packed_out, row_major,
"packed-B vs row-major SMMLA [{m},{k},{n}]"
);
assert_eq!(
packed_out, oracle,
"packed-B vs scalar oracle [{m},{k},{n}]"
);
eprintln!(
r#"{{"check":"smmla_packed_b_parity","m":{m},"k":{k},"n":{n},"result":"pass"}}"#
);
}
}
#[test]
fn packed_b_public_wrapper_matches_dispatch_and_zeroes_out() {
let mut rng = Rng(0x5eed_5eed_5eed_5eed);
for &(m, k, n) in &[(1usize, 48usize, 96usize), (4, 33, 7), (1, 6848, 4)] {
let mut a = vec![0i8; m * k];
let mut b = vec![0i8; n * k];
for v in a.iter_mut() {
*v = rng.i8();
}
for v in b.iter_mut() {
*v = rng.i8();
}
let (panels, _, _) = crate::simd::pack::smmla_pack_panels(&b, 0, n, k, k);
let mut want = vec![0i32; m * n];
igemm_s8s8(&a, &b, m, k, n, &mut want);
let mut got = vec![0x5a5a_5a5ai32; m * n]; igemm_s8s8_packed_b(&a, &panels, m, k, n, &mut got);
assert_eq!(got, want, "public packed-B wrapper [{m},{k},{n}]");
}
}
#[test]
fn sdot_and_smmla_match_oracle_randomized() {
let mut rng = Rng(0x1234_5678_9abc_def0);
let shapes = [
(1usize, 1usize, 1usize),
(1, 16, 1),
(1, 17, 1),
(3, 5, 7),
(4, 16, 4),
(5, 23, 6),
(8, 32, 8),
(7, 31, 9),
(10, 1280, 10),
(6, 896, 13),
];
for &(m, k, n) in &shapes {
let a = rand_i8(&mut rng, m * k);
let b = rand_i8(&mut rng, n * k);
let want = oracle_s8s8(&a, &b, m, k, n);
if let Some(got) = run_sdot_s8s8(&a, &b, m, k, n) {
assert_eq!(got, want, "SDOT mismatch at shape ({m},{k},{n})");
}
if let Some(got) = run_smmla_s8s8(&a, &b, m, k, n) {
assert_eq!(got, want, "SMMLA mismatch at shape ({m},{k},{n})");
}
let mut got = vec![0i32; m * n];
igemm_s8s8(&a, &b, m, k, n, &mut got);
assert_eq!(got, want, "dispatch S8S8 mismatch at shape ({m},{k},{n})");
}
}
#[test]
fn batched_m_equals_per_row_gemv_per_tier() {
let mut rng = Rng(0x0a1b_2c3d_4e5f_6071);
let m_sweep = [1usize, 2, 3, 4, 5, 7, 8, 9, 16, 17, 33, 64, 65, 128, 129];
let k = 320usize; let n = 13usize; let b = rand_i8(&mut rng, n * k);
for &m in &m_sweep {
let a = rand_i8(&mut rng, m * k);
let want = oracle_s8s8(&a, &b, m, k, n);
if let Some(batched) = run_sdot_s8s8(&a, &b, m, k, n) {
assert_eq!(batched, want, "SDOT M=B vs oracle (m={m})");
for r in 0..m {
let row = &a[r * k..(r + 1) * k];
let single = run_sdot_s8s8(row, &b, 1, k, n).expect("dotprod present");
assert_eq!(
&batched[r * n..(r + 1) * n],
&single[..],
"SDOT batched row {r} != standalone m=1 GEMV (m={m})"
);
}
}
if let Some(batched) = run_smmla_s8s8(&a, &b, m, k, n) {
assert_eq!(batched, want, "SMMLA M=B vs oracle (m={m})");
for r in 0..m {
let row = &a[r * k..(r + 1) * k];
let single = run_smmla_s8s8(row, &b, 1, k, n).expect("i8mm present");
assert_eq!(
&batched[r * n..(r + 1) * n],
&single[..],
"SMMLA batched row {r} != standalone m=1 GEMV (m={m})"
);
}
}
}
}
#[test]
fn accelerated_paths_add_into_seeded_out() {
let (m, k, n) = (2usize, 17usize, 3usize);
let a = vec![
1i8, -2, 3, -4, 5, -6, 7, -8, 9, -10, 11, -12, 13, -14, 15, -16, 17, -3, 4, -5, 6, -7,
8, -9, 10, -11, 12, -13, 14, -15, 16, -17, 18, -19,
];
let au = a
.iter()
.map(|&x| (x as i16 + 128) as u8)
.collect::<Vec<_>>();
let b = vec![
2i8, 1, -1, 3, -3, 4, -4, 5, -5, 6, -6, 7, -7, 8, -8, 9, -9, -2, 3, -4, 5, -6, 7, -8,
9, -10, 11, -12, 13, -14, 15, -16, 17, 1, -3, 5, -7, 9, -11, 13, -15, 17, -19, 21, -23,
25, -27, 29, -31, 33, -35,
];
let seed = (0..m * n)
.map(|idx| (idx as i32 * 17) - 41)
.collect::<Vec<_>>();
let mut want_s8 = seed.clone();
for (cell, dot) in want_s8.iter_mut().zip(oracle_s8s8(&a, &b, m, k, n)) {
*cell += dot;
}
if std::arch::is_aarch64_feature_detected!("dotprod") {
let mut got = seed.clone();
unsafe { aarch64_impl::igemm_s8s8_sdot(&a, &b, m, k, n, &mut got) };
assert_eq!(got, want_s8, "SDOT must add into seeded out");
}
if std::arch::is_aarch64_feature_detected!("i8mm") {
let mut got = seed.clone();
unsafe { aarch64_impl::igemm_s8s8_smmla(&a, &b, m, k, n, &mut got) };
assert_eq!(got, want_s8, "SMMLA must add into seeded out");
}
let mut got = seed.clone();
igemm_s8s8(&a, &b, m, k, n, &mut got);
assert_eq!(
got, want_s8,
"public S8S8 dispatch must add into seeded out"
);
let mut want_u8 = seed.clone();
for (cell, dot) in want_u8.iter_mut().zip(oracle_u8s8(&au, &b, m, k, n)) {
*cell += dot;
}
let mut got = seed;
igemm_u8s8(&au, &b, m, k, n, &mut got);
assert_eq!(
got, want_u8,
"public U8S8 dispatch must add into seeded out"
);
}
#[test]
fn s8s8_adversarial_all_max_at_worst_case_k() {
let k = 6848usize;
let m = 3usize;
let n = 5usize;
for &fill in &[127i8, -128i8] {
let a = vec![fill; m * k];
let b = vec![fill; n * k];
let want = oracle_s8s8(&a, &b, m, k, n);
let expect = if fill == 127 {
110_451_392
} else {
112_197_632
};
assert_eq!(want[0], expect, "oracle value at fill {fill}");
if let Some(got) = run_sdot_s8s8(&a, &b, m, k, n) {
assert_eq!(got, want, "SDOT all-{fill} @K=6848");
}
if let Some(got) = run_smmla_s8s8(&a, &b, m, k, n) {
assert_eq!(got, want, "SMMLA all-{fill} @K=6848");
}
}
}
#[test]
fn u8s8_matches_oracle_randomized_and_adversarial() {
let mut rng = Rng(0xdead_beef_0bad_f00d);
let shapes = [
(1usize, 1usize, 1usize),
(2, 15, 3),
(4, 16, 4),
(5, 17, 6),
(8, 64, 8),
(3, 896, 7),
(4, 6848, 5),
];
for &(m, k, n) in &shapes {
let a = rand_u8(&mut rng, m * k);
let b = rand_i8(&mut rng, n * k);
let want = oracle_u8s8(&a, &b, m, k, n);
let mut got = vec![0i32; m * n];
igemm_u8s8(&a, &b, m, k, n, &mut got);
assert_eq!(got, want, "U8S8 dispatch mismatch at shape ({m},{k},{n})");
}
let (m, k, n) = (2usize, 6848usize, 3usize);
let a = vec![255u8; m * k];
let b = vec![127i8; n * k];
let want = oracle_u8s8(&a, &b, m, k, n);
assert_eq!(want[0], 221_772_480, "U8S8 worst-case oracle value");
let mut got = vec![0i32; m * n];
igemm_u8s8(&a, &b, m, k, n, &mut got);
assert_eq!(got, want, "U8S8 adversarial @K=6848");
let b_neg = vec![-128i8; n * k];
let want_neg = oracle_u8s8(&a, &b_neg, m, k, n);
let mut got_neg = vec![0i32; m * n];
igemm_u8s8(&a, &b_neg, m, k, n, &mut got_neg);
assert_eq!(got_neg, want_neg, "U8S8 negative adversarial @K=6848");
}
#[test]
fn empty_and_unit_dims() {
let a = vec![3i8, -2, 5];
let b = vec![7i8, 11];
let (m, k, n) = (3usize, 1usize, 2usize);
let want = oracle_s8s8(&a, &b, m, k, n);
let mut got = vec![0i32; m * n];
igemm_s8s8(&a, &b, m, k, n, &mut got);
assert_eq!(got, want);
}
fn oracle_s4s8_packed(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
) -> Vec<f32> {
let groups = k / group;
let mut out = vec![0.0f32; m * n];
for mi in 0..m {
for ni in 0..n {
let mut acc_f = 0.0f32;
for g in 0..groups {
let mut gi: i32 = 0;
for kk in g * group..(g + 1) * group {
let byte = b_packed[ni * (k / 2) + kk / 2];
let w = if kk % 2 == 0 {
((byte << 4) as i8) >> 4
} else {
(byte as i8) >> 4
};
gi += i32::from(a[mi * k + kk]) * i32::from(w);
}
acc_f += scales[ni * groups + g] * gi as f32;
}
out[mi * n + ni] = acc_f;
}
}
out
}
fn run_sdot_s4s8(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
) -> Option<Vec<f32>> {
if !std::arch::is_aarch64_feature_detected!("dotprod") {
return None;
}
let mut out = vec![0.0f32; m * n];
unsafe {
aarch64_impl::igemm_s4s8_packed_sdot(a, b_packed, scales, group, m, k, n, &mut out);
}
Some(out)
}
fn run_smmla_s4s8(
a: &[i8],
b_packed: &[u8],
scales: &[f32],
group: usize,
m: usize,
k: usize,
n: usize,
) -> Option<Vec<f32>> {
if !std::arch::is_aarch64_feature_detected!("i8mm") {
return None;
}
let mut out = vec![0.0f32; m * n];
unsafe {
aarch64_impl::igemm_s4s8_packed_smmla(a, b_packed, scales, group, m, k, n, &mut out);
}
Some(out)
}
#[test]
fn s4s8_packed_tiers_match_scalar_oracle_randomized() {
let mut rng = Rng(0x5104_9ac4_ed00_0a1b); let cases = [
(1usize, 16usize, 1usize, 16usize),
(1, 32, 1, 32),
(2, 32, 3, 16),
(4, 64, 5, 32),
(3, 96, 2, 16),
(5, 128, 7, 32),
(8, 48, 9, 16),
(7, 160, 6, 32),
(9, 16, 4, 16),
(16, 64, 13, 16),
(6, 6848, 5, 16),
(4, 6848, 5, 32),
];
for (m, k, n, group) in cases {
let groups = k / group;
let a = rand_i8(&mut rng, m * k);
let b_packed = rand_u8(&mut rng, n * (k / 2));
let scales: Vec<f32> = (0..n * groups)
.map(|i| ((i as f32 * 0.013) - 0.4) + if i % 3 == 0 { 0.125 } else { -0.0625 })
.collect();
let want = oracle_s4s8_packed(&a, &b_packed, &scales, group, m, k, n);
if let Some(got) = run_sdot_s4s8(&a, &b_packed, &scales, group, m, k, n) {
assert_eq!(
got, want,
"SDOT-packed != oracle (m={m},k={k},n={n},g={group})"
);
}
if let Some(got) = run_smmla_s4s8(&a, &b_packed, &scales, group, m, k, n) {
assert_eq!(
got, want,
"SMMLA-packed != oracle (m={m},k={k},n={n},g={group})"
);
}
}
}
#[test]
fn s4s8_packed_adversarial_max_operands() {
for &(m, k, n, group) in &[
(3usize, 6848usize, 5usize, 16usize),
(4, 6848, 6, 32),
(5, 64, 7, 16),
] {
let groups = k / group;
let b_packed = vec![0x88u8; n * (k / 2)]; let a: Vec<i8> = (0..m * k)
.map(|i| if i % 2 == 0 { 127 } else { -127 })
.collect();
let scales = vec![1.0f32; n * groups];
let want = oracle_s4s8_packed(&a, &b_packed, &scales, group, m, k, n);
if let Some(got) = run_sdot_s4s8(&a, &b_packed, &scales, group, m, k, n) {
assert_eq!(got, want, "SDOT-packed adversarial (k={k},g={group})");
}
if let Some(got) = run_smmla_s4s8(&a, &b_packed, &scales, group, m, k, n) {
assert_eq!(got, want, "SMMLA-packed adversarial (k={k},g={group})");
}
}
}
}