use crate::streaming::Stream;
#[inline]
pub unsafe fn tile_16x16_i16_i32(
_stream: &Stream,
a_pack: *const i16,
b_pack: *const i16,
c: *mut i32,
ldc: usize,
kc: usize,
accumulate: bool,
) {
debug_assert!(kc > 0);
let ldc_bytes = ldc * 4;
if accumulate {
tile_16x16_i16_accum(a_pack, b_pack, c, ldc_bytes, kc);
} else {
tile_16x16_i16_set(a_pack, b_pack, c, ldc_bytes, kc);
}
}
#[inline(never)]
unsafe fn tile_16x16_i16_set(
a_pack: *const i16,
b_pack: *const i16,
c: *mut i32,
ldc_bytes: usize,
kc: usize,
) {
core::arch::asm!(
".word 0x2518E3E0", ".word 0x2518E3E1", ".word 0xC00800FF",
"1:",
".word 0xA4A0A400", ".word 0xA4A0A421", "add x0, x0, #64",
"add x1, x1, #64",
".word 0xA0810008", "subs x2, x2, #1",
"b.ne 1b",
"mov w12, #0",
"2:",
".word 0xC0820002", ".word 0xE540E062", "add x3, x3, x4",
"add w12, w12, #1",
"cmp w12, #16",
"b.ne 2b",
inout("x0") a_pack => _,
inout("x1") b_pack => _,
inout("x2") kc => _,
inout("x3") c => _,
in("x4") ldc_bytes,
out("x12") _,
options(nostack),
);
}
#[inline(never)]
unsafe fn tile_16x16_i16_accum(
a_pack: *const i16,
b_pack: *const i16,
c: *mut i32,
ldc_bytes: usize,
kc: usize,
) {
core::arch::asm!(
".word 0x2518E3E0", ".word 0x2518E3E1", ".word 0xC00800FF",
"1:",
".word 0xA4A0A400", ".word 0xA4A0A421", "add x0, x0, #64",
"add x1, x1, #64",
".word 0xA0810008", "subs x2, x2, #1",
"b.ne 1b",
"mov w12, #0",
"2:",
".word 0xA540A064", ".word 0xC0820002", ".word 0x04A40044", ".word 0xE540E064", "add x3, x3, x4",
"add w12, w12, #1",
"cmp w12, #16",
"b.ne 2b",
inout("x0") a_pack => _,
inout("x1") b_pack => _,
inout("x2") kc => _,
inout("x3") c => _,
in("x4") ldc_bytes,
out("x12") _,
options(nostack),
);
}
#[cfg(test)]
mod tests {
use super::*;
fn ref_matmul_i16_i32(a: &[i16], b: &[i16], c: &mut [i32], m: usize, n: usize, k: usize) {
for i in 0..m {
for j in 0..n {
let mut acc: i32 = 0;
for p in 0..k {
acc = acc.wrapping_add((a[i * k + p] as i32) * (b[p * n + j] as i32));
}
c[i * n + j] = acc;
}
}
}
fn pack_a(a: &[i16], kc: usize) -> Vec<i16> {
let m = 16usize;
let k_full = 2 * kc;
let mut out = vec![0i16; kc * 32];
for p in 0..kc {
for i in 0..m {
for kk in 0..2 {
out[p * 32 + 2 * i + kk] = a[i * k_full + 2 * p + kk];
}
}
}
out
}
fn pack_b(b: &[i16], kc: usize) -> Vec<i16> {
let n = 16usize;
let mut out = vec![0i16; kc * 32];
for p in 0..kc {
for j in 0..n {
for kk in 0..2 {
out[p * 32 + 2 * j + kk] = b[(2 * p + kk) * n + j];
}
}
}
out
}
#[test]
fn smopa_int16_tile_set_correct() {
if !crate::probe::scan().has_sme2 {
eprintln!("skip: FEAT_SME2 not present");
return;
}
let m = 16usize;
let n = 16usize;
let kc = 4usize;
let k = 2 * kc;
let a: Vec<i16> = (0..m * k).map(|i| (((i as i32) % 11) - 5) as i16).collect();
let b: Vec<i16> = (0..k * n).map(|i| (((i as i32) % 13) - 6) as i16).collect();
let mut c_sme = vec![0i32; m * n];
let mut c_ref = vec![0i32; m * n];
let a_pack = pack_a(&a, kc);
let b_pack = pack_b(&b, kc);
let stream = Stream::new().unwrap();
unsafe {
tile_16x16_i16_i32(
&stream,
a_pack.as_ptr(),
b_pack.as_ptr(),
c_sme.as_mut_ptr(),
n,
kc,
false,
);
}
drop(stream);
ref_matmul_i16_i32(&a, &b, &mut c_ref, m, n, k);
for i in 0..m * n {
assert_eq!(
c_sme[i],
c_ref[i],
"tile mismatch at [{},{}]: sme={}, ref={}",
i / n,
i % n,
c_sme[i],
c_ref[i]
);
}
}
#[test]
fn smopa_int16_tile_accumulate_correct() {
if !crate::probe::scan().has_sme2 {
return;
}
let m = 16usize;
let n = 16usize;
let kc = 2usize;
let k = 2 * kc;
let a: Vec<i16> = (0..m * k).map(|i| ((i as i32) - 8) as i16).collect();
let b: Vec<i16> = (0..k * n).map(|i| ((i as i32) - 8) as i16).collect();
let mut c_sme: Vec<i32> = (0..m * n).map(|i| (i as i32) * 7 - 100).collect();
let mut c_ref: Vec<i32> = c_sme.clone();
let a_pack = pack_a(&a, kc);
let b_pack = pack_b(&b, kc);
let stream = Stream::new().unwrap();
unsafe {
tile_16x16_i16_i32(
&stream,
a_pack.as_ptr(),
b_pack.as_ptr(),
c_sme.as_mut_ptr(),
n,
kc,
true,
);
}
drop(stream);
for i in 0..m {
for j in 0..n {
let mut acc: i32 = 0;
for p in 0..k {
acc = acc.wrapping_add((a[i * k + p] as i32) * (b[p * n + j] as i32));
}
c_ref[i * n + j] = c_ref[i * n + j].wrapping_add(acc);
}
}
for i in 0..m * n {
assert_eq!(c_sme[i], c_ref[i], "accumulate mismatch at idx {i}");
}
}
}