#[inline(always)]
fn acc_add(acc: i32, prod: i32) -> i32 {
acc + prod
}
#[track_caller]
pub(crate) fn checked_len(context: &str, lhs: usize, rhs: usize, expr: &str) -> usize {
let len = lhs.checked_mul(rhs);
assert!(len.is_some(), "{context}: {expr} overflow ({lhs} * {rhs})");
len.unwrap_or(0)
}
#[track_caller]
pub(crate) fn assert_gemm_shapes(
context: &str,
a_len: usize,
b_len: usize,
out_len: usize,
m: usize,
k: usize,
n: usize,
) {
let want_a = checked_len(context, m, k, "m*k");
let want_b = checked_len(context, n, k, "n*k");
let want_out = checked_len(context, m, n, "m*n");
assert_eq!(a_len, want_a, "{context}: a.len {a_len} != m*k {want_a}");
assert_eq!(b_len, want_b, "{context}: b.len {b_len} != n*k {want_b}");
assert_eq!(
out_len, want_out,
"{context}: out.len {out_len} != m*n {want_out}"
);
}
pub fn igemm_s8s8(a: &[i8], b: &[i8], m: usize, k: usize, n: usize, out: &mut [i32]) {
assert_gemm_shapes("igemm_s8s8", a.len(), b.len(), out.len(), m, k, n);
for i in 0..m {
let a_row = &a[i * k..i * k + k];
let out_row = &mut out[i * n..i * n + n];
for o in 0..n {
let b_row = &b[o * k..o * k + k];
let mut acc: i32 = 0;
for p in 0..k {
acc = acc_add(acc, i32::from(a_row[p]) * i32::from(b_row[p]));
}
out_row[o] += acc;
}
}
}
pub fn igemm_u8s8(a: &[u8], b: &[i8], m: usize, k: usize, n: usize, out: &mut [i32]) {
assert_gemm_shapes("igemm_u8s8", a.len(), b.len(), out.len(), m, k, n);
for i in 0..m {
let a_row = &a[i * k..i * k + k];
let out_row = &mut out[i * n..i * n + n];
for o in 0..n {
let b_row = &b[o * k..o * k + k];
let mut acc: i32 = 0;
for p in 0..k {
acc = acc_add(acc, i32::from(a_row[p]) * i32::from(b_row[p]));
}
out_row[o] += acc;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn s8s8_matches_hand_computed() {
let a: [i8; 6] = [1, 2, 3, 4, 5, 6];
let b: [i8; 6] = [1, 0, 1, 0, 1, 0];
let mut out = [0i32; 4];
igemm_s8s8(&a, &b, 2, 3, 2, &mut out);
assert_eq!(out, [4, 2, 10, 5]);
}
#[test]
fn s8s8_accumulates_into_out() {
let a: [i8; 3] = [1, 2, 3];
let b: [i8; 3] = [1, 1, 1];
let mut out = [100i32; 1];
igemm_s8s8(&a, &b, 1, 3, 1, &mut out); assert_eq!(out, [106]);
igemm_s8s8(&a, &b, 1, 3, 1, &mut out); assert_eq!(out, [112]);
}
#[test]
fn u8s8_matches_hand_computed() {
let a: [u8; 3] = [10, 20, 30];
let b: [i8; 3] = [2, -1, 1];
let mut out = [0i32; 1];
igemm_u8s8(&a, &b, 1, 3, 1, &mut out);
assert_eq!(out, [30]);
}
#[test]
fn s8s8_handles_negatives() {
let a: [i8; 2] = [-2, 3];
let b: [i8; 4] = [-1, -1, 4, -2];
let mut out = [0i32; 2];
igemm_s8s8(&a, &b, 1, 2, 2, &mut out);
assert_eq!(out, [-1, -14]);
}
#[test]
fn s8s8_no_overflow_at_k6848_all_max() {
const K: usize = 6848;
let a = vec![127i8; K];
let b = vec![127i8; K];
let mut out = [0i32; 1];
igemm_s8s8(&a, &b, 1, K, 1, &mut out);
assert_eq!(out[0], 110_451_392);
assert!(out[0] < i32::MAX);
}
#[test]
fn s8s8_no_overflow_at_k6848_all_neg128() {
const K: usize = 6848;
let a = vec![-128i8; K];
let b = vec![-128i8; K];
let mut out = [0i32; 1];
igemm_s8s8(&a, &b, 1, K, 1, &mut out);
assert_eq!(out[0], 112_197_632);
assert!(out[0] < i32::MAX);
}
#[test]
fn u8s8_no_overflow_at_k6848_all_max() {
const K: usize = 6848;
let a = vec![255u8; K];
let b = vec![127i8; K];
let mut out = [0i32; 1];
igemm_u8s8(&a, &b, 1, K, 1, &mut out);
assert_eq!(out[0], 221_772_480);
assert!(out[0] < i32::MAX);
}
#[test]
fn s8s8_matches_i64_oracle_randomized() {
let (m, k, n) = (3usize, 17usize, 5usize);
let a = pseudo_i8(m * k, 0x1234_5678);
let b = pseudo_i8(n * k, 0x9abc_def0);
let mut out = vec![0i32; m * n];
igemm_s8s8(&a, &b, m, k, n, &mut out);
for i in 0..m {
for o in 0..n {
let mut acc: i64 = 0;
for p in 0..k {
acc += i64::from(a[i * k + p]) * i64::from(b[o * k + p]);
}
assert_eq!(i64::from(out[i * n + o]), acc, "mismatch at ({i},{o})");
}
}
}
#[test]
fn u8s8_matches_i64_oracle_randomized() {
let (m, k, n) = (4usize, 13usize, 6usize);
let a = pseudo_u8(m * k, 0x0f0f_1234);
let b = pseudo_i8(n * k, 0xfeed_face);
let mut out = vec![0i32; m * n];
igemm_u8s8(&a, &b, m, k, n, &mut out);
for i in 0..m {
for o in 0..n {
let mut acc: i64 = 0;
for p in 0..k {
acc += i64::from(a[i * k + p]) * i64::from(b[o * k + p]);
}
assert_eq!(i64::from(out[i * n + o]), acc, "mismatch at ({i},{o})");
}
}
}
#[test]
#[should_panic(expected = "a.len")]
fn s8s8_rejects_bad_a_len() {
let mut out = [0i32; 1];
igemm_s8s8(&[1i8, 2], &[1i8], 1, 1, 1, &mut out);
}
#[test]
#[should_panic(expected = "igemm_s8s8: m*k overflow")]
fn s8s8_rejects_shape_product_overflow_before_len_checks() {
let mut out = [];
igemm_s8s8(&[], &[], usize::MAX, 2, 0, &mut out);
}
#[test]
#[should_panic(expected = "igemm_u8s8: m*n overflow")]
fn u8s8_rejects_output_shape_overflow_before_looping() {
let mut out = [];
igemm_u8s8(&[], &[], usize::MAX, 0, 2, &mut out);
}
fn xorshift(state: &mut u32) -> u32 {
let mut x = *state;
x ^= x << 13;
x ^= x >> 17;
x ^= x << 5;
*state = x;
x
}
fn pseudo_i8(len: usize, seed: u32) -> Vec<i8> {
let mut s = seed | 1;
(0..len)
.map(|_| (xorshift(&mut s) & 0xff) as u8 as i8)
.collect()
}
fn pseudo_u8(len: usize, seed: u32) -> Vec<u8> {
let mut s = seed | 1;
(0..len).map(|_| (xorshift(&mut s) & 0xff) as u8).collect()
}
}