#![allow(unsafe_code, unsafe_op_in_unsafe_fn)]
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],
) -> super::dispatch::EffectiveI8Route {
let a_len = super::scalar::checked_len("igemm_s8s8", m, k, "m*k");
let b_len = super::scalar::checked_len("igemm_s8s8", n, k, "n*k");
let out_len = super::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
);
match super::dispatch::detected_tier() {
super::dispatch::IsaTier::Avx512Vnni => {
unsafe {
x86_avx512vnni::igemm_s8s8_avx512vnni(a, b, m, k, n, out);
}
super::dispatch::EffectiveI8Route::Avx512Vnni
}
super::dispatch::IsaTier::AvxVnni => {
unsafe {
x86_avxvnni::igemm_s8s8_avxvnni(a, b, m, k, n, out);
}
super::dispatch::EffectiveI8Route::AvxVnni
}
super::dispatch::IsaTier::Avx2 => {
unsafe {
x86_avx2::igemm_s8s8_avx2(a, b, m, k, n, out);
}
super::dispatch::EffectiveI8Route::Avx2
}
super::dispatch::IsaTier::Scalar => {
scalar_s8s8(a, b, m, k, n, out);
super::dispatch::EffectiveI8Route::Scalar
}
super::dispatch::IsaTier::Sdot | super::dispatch::IsaTier::Smmla => {
unreachable!("ARM ISA tier cannot be selected by an x86-64 build")
}
super::dispatch::IsaTier::WasmSimd128 => {
unreachable!("wasm ISA tier cannot be selected by an x86-64 build")
}
}
}
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],
) -> super::dispatch::EffectiveI8Route {
let a_len = super::scalar::checked_len("igemm_u8s8", m, k, "m*k");
let b_len = super::scalar::checked_len("igemm_u8s8", n, k, "n*k");
let out_len = super::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
);
match super::dispatch::detected_tier() {
super::dispatch::IsaTier::Avx512Vnni => {
unsafe {
x86_avx512vnni::igemm_u8s8_avx512vnni(a, b, m, k, n, out);
}
super::dispatch::EffectiveI8Route::Avx512Vnni
}
super::dispatch::IsaTier::AvxVnni => {
unsafe {
x86_avxvnni::igemm_u8s8_avxvnni(a, b, m, k, n, out);
}
super::dispatch::EffectiveI8Route::AvxVnni
}
super::dispatch::IsaTier::Avx2 => {
unsafe {
x86_avx2::igemm_u8s8_avx2(a, b, m, k, n, out);
}
super::dispatch::EffectiveI8Route::Avx2
}
super::dispatch::IsaTier::Scalar => {
scalar_u8s8(a, b, m, k, n, out);
super::dispatch::EffectiveI8Route::Scalar
}
super::dispatch::IsaTier::Sdot | super::dispatch::IsaTier::Smmla => {
unreachable!("ARM ISA tier cannot be selected by an x86-64 build")
}
super::dispatch::IsaTier::WasmSimd128 => {
unreachable!("wasm ISA tier cannot be selected by an x86-64 build")
}
}
}
pub fn scalar_s8s8(a: &[i8], b: &[i8], m: usize, k: usize, n: usize, out: &mut [i32]) {
for r in 0..m {
let arow = &a[r * k..r * k + k];
for c in 0..n {
let brow = &b[c * k..c * k + k];
let mut acc: i32 = 0;
for t in 0..k {
acc += i32::from(arow[t]) * i32::from(brow[t]);
}
out[r * n + c] += acc;
}
}
}
pub fn scalar_u8s8(a: &[u8], b: &[i8], m: usize, k: usize, n: usize, out: &mut [i32]) {
for r in 0..m {
let arow = &a[r * k..r * k + k];
for c in 0..n {
let brow = &b[c * k..c * k + k];
let mut acc: i32 = 0;
for t in 0..k {
acc += i32::from(arow[t]) * i32::from(brow[t]);
}
out[r * n + c] += acc;
}
}
}
#[cfg(target_arch = "x86_64")]
mod x86_avx2 {
#![allow(clippy::needless_range_loop)]
use core::arch::x86_64::*;
const MR: usize = 2;
const NR: usize = 2;
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn hsum_i32_avx2(v: __m256i) -> i32 {
let lo = _mm256_castsi256_si128(v);
let hi = _mm256_extracti128_si256::<1>(v);
let s = _mm_add_epi32(lo, hi); let s = _mm_hadd_epi32(s, s); let s = _mm_hadd_epi32(s, s); _mm_cvtsi128_si32(s)
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn load_i8x16_to_i16(p: *const i8) -> __m256i {
let lo = _mm_loadu_si128(p.cast::<__m128i>());
_mm256_cvtepi8_epi16(lo)
}
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn load_u8x16_to_i16(p: *const u8) -> __m256i {
let lo = _mm_loadu_si128(p.cast::<__m128i>());
_mm256_cvtepu8_epi16(lo)
}
#[target_feature(enable = "avx2")]
pub unsafe fn igemm_s8s8_avx2(
a: &[i8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
let k16 = k - (k % 16);
let mut r0 = 0;
while r0 < m {
let mr = MR.min(m - r0);
let mut c0 = 0;
while c0 < n {
let nr = NR.min(n - c0);
let mut acc = [[_mm256_setzero_si256(); NR]; MR];
let mut t = 0;
while t < k16 {
let mut av = [_mm256_setzero_si256(); MR];
for i in 0..mr {
av[i] = load_i8x16_to_i16(a.as_ptr().add((r0 + i) * k + t));
}
let mut bv = [_mm256_setzero_si256(); NR];
for j in 0..nr {
bv[j] = load_i8x16_to_i16(b.as_ptr().add((c0 + j) * k + t));
}
for i in 0..mr {
for j in 0..nr {
let prod = _mm256_madd_epi16(av[i], bv[j]);
acc[i][j] = _mm256_add_epi32(acc[i][j], prod);
}
}
t += 16;
}
for i in 0..mr {
for j in 0..nr {
let mut s = hsum_i32_avx2(acc[i][j]);
let arow = &a[(r0 + i) * k..(r0 + i) * k + k];
let brow = &b[(c0 + j) * k..(c0 + j) * k + k];
for tt in k16..k {
s += i32::from(arow[tt]) * i32::from(brow[tt]);
}
out[(r0 + i) * n + (c0 + j)] += s;
}
}
c0 += nr;
}
r0 += mr;
}
}
#[target_feature(enable = "avx2")]
pub unsafe fn igemm_u8s8_avx2(
a: &[u8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
let k16 = k - (k % 16);
let mut r0 = 0;
while r0 < m {
let mr = MR.min(m - r0);
let mut c0 = 0;
while c0 < n {
let nr = NR.min(n - c0);
let mut acc = [[_mm256_setzero_si256(); NR]; MR];
let mut t = 0;
while t < k16 {
let mut av = [_mm256_setzero_si256(); MR];
for i in 0..mr {
av[i] = load_u8x16_to_i16(a.as_ptr().add((r0 + i) * k + t));
}
let mut bv = [_mm256_setzero_si256(); NR];
for j in 0..nr {
bv[j] = load_i8x16_to_i16(b.as_ptr().add((c0 + j) * k + t));
}
for i in 0..mr {
for j in 0..nr {
let prod = _mm256_madd_epi16(av[i], bv[j]);
acc[i][j] = _mm256_add_epi32(acc[i][j], prod);
}
}
t += 16;
}
for i in 0..mr {
for j in 0..nr {
let mut s = hsum_i32_avx2(acc[i][j]);
let arow = &a[(r0 + i) * k..(r0 + i) * k + k];
let brow = &b[(c0 + j) * k..(c0 + j) * k + k];
for tt in k16..k {
s += i32::from(arow[tt]) * i32::from(brow[tt]);
}
out[(r0 + i) * n + (c0 + j)] += s;
}
}
c0 += nr;
}
r0 += mr;
}
}
}
#[cfg(target_arch = "x86_64")]
mod x86_avxvnni {
#![allow(clippy::needless_range_loop)]
use core::arch::x86_64::*;
const MR: usize = 2;
const NR: usize = 2;
#[inline]
#[target_feature(enable = "avx2")]
unsafe fn hsum_i32(v: __m256i) -> i32 {
let lo = _mm256_castsi256_si128(v);
let hi = _mm256_extracti128_si256::<1>(v);
let s = _mm_add_epi32(lo, hi);
let s = _mm_hadd_epi32(s, s);
let s = _mm_hadd_epi32(s, s);
_mm_cvtsi128_si32(s)
}
#[target_feature(enable = "avxvnni")]
pub unsafe fn igemm_u8s8_avxvnni(
a: &[u8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
let k32 = k - (k % 32);
let mut r0 = 0;
while r0 < m {
let mr = MR.min(m - r0);
let mut c0 = 0;
while c0 < n {
let nr = NR.min(n - c0);
let mut acc = [[_mm256_setzero_si256(); NR]; MR];
let mut t = 0;
while t < k32 {
let mut av = [_mm256_setzero_si256(); MR];
for i in 0..mr {
av[i] =
_mm256_loadu_si256(a.as_ptr().add((r0 + i) * k + t).cast::<__m256i>());
}
let mut bv = [_mm256_setzero_si256(); NR];
for j in 0..nr {
bv[j] =
_mm256_loadu_si256(b.as_ptr().add((c0 + j) * k + t).cast::<__m256i>());
}
for i in 0..mr {
for j in 0..nr {
acc[i][j] = _mm256_dpbusd_avx_epi32(acc[i][j], av[i], bv[j]);
}
}
t += 32;
}
for i in 0..mr {
for j in 0..nr {
let mut s = hsum_i32(acc[i][j]);
let arow = &a[(r0 + i) * k..(r0 + i) * k + k];
let brow = &b[(c0 + j) * k..(c0 + j) * k + k];
for tt in k32..k {
s += i32::from(arow[tt]) * i32::from(brow[tt]);
}
out[(r0 + i) * n + (c0 + j)] += s;
}
}
c0 += nr;
}
r0 += mr;
}
}
#[target_feature(enable = "avxvnni")]
pub unsafe fn igemm_s8s8_avxvnni(
a: &[i8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
let k32 = k - (k % 32);
let bias = _mm256_set1_epi8(-128i8); let mut r0 = 0;
while r0 < m {
let mr = MR.min(m - r0);
let mut c0 = 0;
while c0 < n {
let nr = NR.min(n - c0);
let mut acc = [[_mm256_setzero_si256(); NR]; MR];
let mut t = 0;
while t < k32 {
let mut av = [_mm256_setzero_si256(); MR];
for i in 0..mr {
let raw =
_mm256_loadu_si256(a.as_ptr().add((r0 + i) * k + t).cast::<__m256i>());
av[i] = _mm256_add_epi8(raw, bias);
}
let mut bv = [_mm256_setzero_si256(); NR];
for j in 0..nr {
bv[j] =
_mm256_loadu_si256(b.as_ptr().add((c0 + j) * k + t).cast::<__m256i>());
}
for i in 0..mr {
for j in 0..nr {
acc[i][j] = _mm256_dpbusd_avx_epi32(acc[i][j], av[i], bv[j]);
}
}
t += 32;
}
for i in 0..mr {
for j in 0..nr {
let arow = &a[(r0 + i) * k..(r0 + i) * k + k];
let brow = &b[(c0 + j) * k..(c0 + j) * k + k];
let mut s = hsum_i32(acc[i][j]);
let mut bsum_vec: i32 = 0;
for tt in 0..k32 {
bsum_vec += i32::from(brow[tt]);
}
s -= 128 * bsum_vec;
for tt in k32..k {
s += i32::from(arow[tt]) * i32::from(brow[tt]);
}
out[(r0 + i) * n + (c0 + j)] += s;
}
}
c0 += nr;
}
r0 += mr;
}
}
}
#[cfg(target_arch = "x86_64")]
mod x86_avx512vnni {
#![allow(clippy::needless_range_loop)]
use core::arch::x86_64::*;
const MR: usize = 2;
const NR: usize = 2;
#[inline]
#[target_feature(enable = "avx512f")]
unsafe fn hsum_i32_512(v: __m512i) -> i32 {
_mm512_reduce_add_epi32(v)
}
#[target_feature(enable = "avx512vnni,avx512bw,avx512f")]
pub unsafe fn igemm_u8s8_avx512vnni(
a: &[u8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
let k64 = k - (k % 64);
let mut r0 = 0;
while r0 < m {
let mr = MR.min(m - r0);
let mut c0 = 0;
while c0 < n {
let nr = NR.min(n - c0);
let mut acc = [[_mm512_setzero_si512(); NR]; MR];
let mut t = 0;
while t < k64 {
let mut av = [_mm512_setzero_si512(); MR];
for i in 0..mr {
av[i] =
_mm512_loadu_si512(a.as_ptr().add((r0 + i) * k + t).cast::<__m512i>());
}
let mut bv = [_mm512_setzero_si512(); NR];
for j in 0..nr {
bv[j] =
_mm512_loadu_si512(b.as_ptr().add((c0 + j) * k + t).cast::<__m512i>());
}
for i in 0..mr {
for j in 0..nr {
acc[i][j] = _mm512_dpbusd_epi32(acc[i][j], av[i], bv[j]);
}
}
t += 64;
}
for i in 0..mr {
for j in 0..nr {
let mut s = hsum_i32_512(acc[i][j]);
let arow = &a[(r0 + i) * k..(r0 + i) * k + k];
let brow = &b[(c0 + j) * k..(c0 + j) * k + k];
for tt in k64..k {
s += i32::from(arow[tt]) * i32::from(brow[tt]);
}
out[(r0 + i) * n + (c0 + j)] += s;
}
}
c0 += nr;
}
r0 += mr;
}
}
#[target_feature(enable = "avx512vnni,avx512bw,avx512f")]
pub unsafe fn igemm_s8s8_avx512vnni(
a: &[i8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
let k64 = k - (k % 64);
let bias = _mm512_set1_epi8(-128i8);
let mut r0 = 0;
while r0 < m {
let mr = MR.min(m - r0);
let mut c0 = 0;
while c0 < n {
let nr = NR.min(n - c0);
let mut acc = [[_mm512_setzero_si512(); NR]; MR];
let mut t = 0;
while t < k64 {
let mut av = [_mm512_setzero_si512(); MR];
for i in 0..mr {
let raw =
_mm512_loadu_si512(a.as_ptr().add((r0 + i) * k + t).cast::<__m512i>());
av[i] = _mm512_add_epi8(raw, bias); }
let mut bv = [_mm512_setzero_si512(); NR];
for j in 0..nr {
bv[j] =
_mm512_loadu_si512(b.as_ptr().add((c0 + j) * k + t).cast::<__m512i>());
}
for i in 0..mr {
for j in 0..nr {
acc[i][j] = _mm512_dpbusd_epi32(acc[i][j], av[i], bv[j]);
}
}
t += 64;
}
for i in 0..mr {
for j in 0..nr {
let arow = &a[(r0 + i) * k..(r0 + i) * k + k];
let brow = &b[(c0 + j) * k..(c0 + j) * k + k];
let mut s = hsum_i32_512(acc[i][j]);
let mut bsum_vec: i32 = 0;
for tt in 0..k64 {
bsum_vec += i32::from(brow[tt]);
}
s -= 128 * bsum_vec;
for tt in k64..k {
s += i32::from(arow[tt]) * i32::from(brow[tt]);
}
out[(r0 + i) * n + (c0 + j)] += s;
}
}
c0 += nr;
}
r0 += mr;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
struct Rng(u64);
impl Rng {
fn new(seed: u64) -> Self {
Rng(seed ^ 0x9E37_79B9_7F4A_7C15)
}
fn next_u64(&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_u64() & 0xFF) as u8 as i8
}
fn u8(&mut self) -> u8 {
(self.next_u64() & 0xFF) as u8
}
}
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 c in 0..n {
let mut acc: i64 = 0;
for t in 0..k {
acc += i64::from(a[r * k + t]) * i64::from(b[c * k + t]);
}
out[r * n + c] = acc as i32;
}
}
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 c in 0..n {
let mut acc: i64 = 0;
for t in 0..k {
acc += i64::from(a[r * k + t]) * i64::from(b[c * k + t]);
}
out[r * n + c] = acc as i32;
}
}
out
}
fn rand_a_s8(rng: &mut Rng, len: usize) -> Vec<i8> {
(0..len).map(|_| rng.i8()).collect()
}
fn rand_a_u8(rng: &mut Rng, len: usize) -> Vec<u8> {
(0..len).map(|_| rng.u8()).collect()
}
fn rand_b(rng: &mut Rng, len: usize) -> Vec<i8> {
(0..len).map(|_| rng.i8()).collect()
}
#[test]
fn scalar_s8s8_matches_i64_oracle() {
let mut rng = Rng::new(1);
for &(m, k, n) in &[(1, 1, 1), (2, 16, 3), (3, 17, 4), (4, 33, 5), (2, 6848, 2)] {
let a = rand_a_s8(&mut rng, m * k);
let b = rand_b(&mut rng, n * k);
let mut out = vec![0i32; m * n];
scalar_s8s8(&a, &b, m, k, n, &mut out);
assert_eq!(out, oracle_s8s8(&a, &b, m, k, n), "m={m} k={k} n={n}");
}
}
#[test]
fn scalar_u8s8_matches_i64_oracle() {
let mut rng = Rng::new(2);
for &(m, k, n) in &[(1, 1, 1), (2, 16, 3), (3, 17, 4), (4, 33, 5), (2, 6848, 2)] {
let a = rand_a_u8(&mut rng, m * k);
let b = rand_b(&mut rng, n * k);
let mut out = vec![0i32; m * n];
scalar_u8s8(&a, &b, m, k, n, &mut out);
assert_eq!(out, oracle_u8s8(&a, &b, m, k, n), "m={m} k={k} n={n}");
}
}
#[test]
fn accumulation_adds_into_out() {
let a = vec![1i8, 2, 3];
let b = vec![1i8, 1, 1, 2, 0, 1]; let mut out = vec![100i32, 200];
scalar_s8s8(&a, &b, 1, 3, 2, &mut out);
assert_eq!(out, vec![106, 205]);
}
#[test]
fn dispatch_s8s8_matches_oracle_random_and_adversarial() {
let mut rng = Rng::new(3);
let shapes = [
(1, 1, 1),
(2, 16, 2),
(3, 31, 4),
(5, 64, 3),
(2, 100, 2),
(1, 6848, 1),
];
for &(m, k, n) in &shapes {
let a = rand_a_s8(&mut rng, m * k);
let b = rand_b(&mut rng, n * k);
let mut out = vec![0i32; m * n];
igemm_s8s8(&a, &b, m, k, n, &mut out);
assert_eq!(
out,
oracle_s8s8(&a, &b, m, k, n),
"random m={m} k={k} n={n}"
);
}
for &(m, k, n) in &[(2, 6848, 2), (1, 6848, 3)] {
let a = vec![127i8; m * k];
let b = vec![127i8; n * k];
let mut out = vec![0i32; m * n];
igemm_s8s8(&a, &b, m, k, n, &mut out);
assert_eq!(out, oracle_s8s8(&a, &b, m, k, n), "+127 m={m} k={k} n={n}");
let a = vec![-128i8; m * k];
let b = vec![-128i8; n * k];
let mut out = vec![0i32; m * n];
igemm_s8s8(&a, &b, m, k, n, &mut out);
assert_eq!(out, oracle_s8s8(&a, &b, m, k, n), "-128 m={m} k={k} n={n}");
let a = vec![-128i8; m * k];
let b = vec![127i8; n * k];
let mut out = vec![0i32; m * n];
igemm_s8s8(&a, &b, m, k, n, &mut out);
assert_eq!(out, oracle_s8s8(&a, &b, m, k, n), "mixed m={m} k={k} n={n}");
}
}
#[test]
fn dispatch_u8s8_matches_oracle_random_and_adversarial() {
let mut rng = Rng::new(4);
let shapes = [
(1, 1, 1),
(2, 16, 2),
(3, 31, 4),
(5, 64, 3),
(2, 100, 2),
(1, 6848, 1),
];
for &(m, k, n) in &shapes {
let a = rand_a_u8(&mut rng, m * k);
let b = rand_b(&mut rng, n * k);
let mut out = vec![0i32; m * n];
igemm_u8s8(&a, &b, m, k, n, &mut out);
assert_eq!(
out,
oracle_u8s8(&a, &b, m, k, n),
"random m={m} k={k} n={n}"
);
}
for &(m, k, n) in &[(2, 6848, 2), (1, 6848, 3)] {
let a = vec![255u8; m * k];
let b = vec![127i8; n * k];
let mut out = vec![0i32; m * n];
igemm_u8s8(&a, &b, m, k, n, &mut out);
assert_eq!(
out,
oracle_u8s8(&a, &b, m, k, n),
"255*127 m={m} k={k} n={n}"
);
let b = vec![-128i8; n * k];
let mut out = vec![0i32; m * n];
igemm_u8s8(&a, &b, m, k, n, &mut out);
assert_eq!(
out,
oracle_u8s8(&a, &b, m, k, n),
"255*-128 m={m} k={k} n={n}"
);
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avx2_tiers_bit_identical_to_scalar() {
if !is_x86_feature_detected!("avx2") {
eprintln!("[skip] avx2 not present on this host");
return;
}
let mut rng = Rng::new(10);
let shapes = [
(1, 1, 1),
(2, 15, 3),
(3, 16, 2),
(4, 17, 5),
(2, 64, 4),
(3, 6848, 2),
];
for &(m, k, n) in &shapes {
let a = rand_a_s8(&mut rng, m * k);
let b = rand_b(&mut rng, n * k);
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avx2::igemm_s8s8_avx2(&a, &b, m, k, n, &mut got) };
scalar_s8s8(&a, &b, m, k, n, &mut want);
assert_eq!(got, want, "avx2 s8s8 m={m} k={k} n={n}");
let au = rand_a_u8(&mut rng, m * k);
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avx2::igemm_u8s8_avx2(&au, &b, m, k, n, &mut got) };
scalar_u8s8(&au, &b, m, k, n, &mut want);
assert_eq!(got, want, "avx2 u8s8 m={m} k={k} n={n}");
}
let (m, k, n) = (2, 6848, 2);
let a = vec![127i8; m * k];
let b = vec![127i8; n * k];
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avx2::igemm_s8s8_avx2(&a, &b, m, k, n, &mut got) };
scalar_s8s8(&a, &b, m, k, n, &mut want);
assert_eq!(got, want, "avx2 s8s8 adversarial");
let au = vec![255u8; m * k];
let b = vec![127i8; n * k];
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avx2::igemm_u8s8_avx2(&au, &b, m, k, n, &mut got) };
scalar_u8s8(&au, &b, m, k, n, &mut want);
assert_eq!(got, want, "avx2 u8s8 adversarial");
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avxvnni_tiers_bit_identical_to_scalar() {
if !is_x86_feature_detected!("avxvnni") {
eprintln!("[skip] avxvnni not present on this host");
return;
}
let mut rng = Rng::new(11);
let shapes = [
(1, 1, 1),
(2, 31, 3),
(3, 32, 2),
(4, 33, 5),
(2, 96, 4),
(3, 6848, 2),
];
for &(m, k, n) in &shapes {
let a = rand_a_s8(&mut rng, m * k);
let b = rand_b(&mut rng, n * k);
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avxvnni::igemm_s8s8_avxvnni(&a, &b, m, k, n, &mut got) };
scalar_s8s8(&a, &b, m, k, n, &mut want);
assert_eq!(got, want, "avxvnni s8s8 m={m} k={k} n={n}");
let au = rand_a_u8(&mut rng, m * k);
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avxvnni::igemm_u8s8_avxvnni(&au, &b, m, k, n, &mut got) };
scalar_u8s8(&au, &b, m, k, n, &mut want);
assert_eq!(got, want, "avxvnni u8s8 m={m} k={k} n={n}");
}
let (m, k, n) = (2, 6848, 2);
let a = vec![-128i8; m * k];
let b = vec![127i8; n * k];
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avxvnni::igemm_s8s8_avxvnni(&a, &b, m, k, n, &mut got) };
scalar_s8s8(&a, &b, m, k, n, &mut want);
assert_eq!(got, want, "avxvnni s8s8 adversarial -128*127");
let au = vec![255u8; m * k];
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avxvnni::igemm_u8s8_avxvnni(&au, &b, m, k, n, &mut got) };
scalar_u8s8(&au, &b, m, k, n, &mut want);
assert_eq!(got, want, "avxvnni u8s8 adversarial 255*127");
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avx512vnni_tiers_bit_identical_to_scalar() {
if !(is_x86_feature_detected!("avx512vnni")
&& is_x86_feature_detected!("avx512bw")
&& is_x86_feature_detected!("avx512f"))
{
eprintln!("[skip] avx512vnni/bw/f not present on this host");
return;
}
let mut rng = Rng::new(12);
let shapes = [
(1, 1, 1),
(2, 63, 3),
(3, 64, 2),
(4, 65, 5),
(2, 192, 4),
(3, 6848, 2),
];
for &(m, k, n) in &shapes {
let a = rand_a_s8(&mut rng, m * k);
let b = rand_b(&mut rng, n * k);
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avx512vnni::igemm_s8s8_avx512vnni(&a, &b, m, k, n, &mut got) };
scalar_s8s8(&a, &b, m, k, n, &mut want);
assert_eq!(got, want, "avx512vnni s8s8 m={m} k={k} n={n}");
let au = rand_a_u8(&mut rng, m * k);
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avx512vnni::igemm_u8s8_avx512vnni(&au, &b, m, k, n, &mut got) };
scalar_u8s8(&au, &b, m, k, n, &mut want);
assert_eq!(got, want, "avx512vnni u8s8 m={m} k={k} n={n}");
}
let (m, k, n) = (2, 6848, 2);
let a = vec![-128i8; m * k];
let b = vec![127i8; n * k];
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avx512vnni::igemm_s8s8_avx512vnni(&a, &b, m, k, n, &mut got) };
scalar_s8s8(&a, &b, m, k, n, &mut want);
assert_eq!(got, want, "avx512vnni s8s8 adversarial");
let au = vec![255u8; m * k];
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
unsafe { super::x86_avx512vnni::igemm_u8s8_avx512vnni(&au, &b, m, k, n, &mut got) };
scalar_u8s8(&au, &b, m, k, n, &mut want);
assert_eq!(got, want, "avx512vnni u8s8 adversarial");
}
}