use crate::streaming::Stream;
#[inline]
pub unsafe fn tile_16x16_f32(
_stream: &Stream,
a_pack: *const f32,
b_pack: *const f32,
c: *mut f32,
ldc: usize,
kc: usize,
accumulate: bool,
) {
debug_assert!(kc > 0);
let ldc_bytes: usize = ldc * 4;
if accumulate {
tile_16x16_f32_accum(a_pack, b_pack, c, ldc_bytes, kc);
} else {
tile_16x16_f32_set(a_pack, b_pack, c, ldc_bytes, kc);
}
}
#[inline(never)]
unsafe fn tile_16x16_f32_set(
a_pack: *const f32,
b_pack: *const f32,
c: *mut f32,
ldc_bytes: usize,
kc: usize,
) {
core::arch::asm!(
".word 0x2598E3E0", ".word 0xC00800FF",
"1:",
".word 0xA540A000", ".word 0xA540A021", "add x0, x0, #64",
"add x1, x1, #64",
".word 0x80810000", "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_f32_accum(
a_pack: *const f32,
b_pack: *const f32,
c: *mut f32,
ldc_bytes: usize,
kc: usize,
) {
core::arch::asm!(
".word 0x2598E3E0", ".word 0xC00800FF",
"1:",
".word 0xA540A000", ".word 0xA540A021", "add x0, x0, #64",
"add x1, x1, #64",
".word 0x80810000", "subs x2, x2, #1",
"b.ne 1b",
"mov w12, #0",
"2:",
".word 0xA540A064", ".word 0xC0820002", ".word 0x65840044", ".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(a: &[f32], b: &[f32], c: &mut [f32], m: usize, n: usize, k: usize) {
for i in 0..m {
for j in 0..n {
let mut acc = 0.0f32;
for p in 0..k {
acc += a[i * k + p] * b[p * n + j];
}
c[i * n + j] = acc;
}
}
}
fn pack_for_tile(a: &[f32], b: &[f32], k: usize) -> (Vec<f32>, Vec<f32>) {
let mut a_pack = vec![0.0f32; k * 16];
let mut b_pack = vec![0.0f32; k * 16];
for p in 0..k {
for i in 0..16 {
a_pack[p * 16 + i] = a[i * k + p];
}
for j in 0..16 {
b_pack[p * 16 + j] = b[p * 16 + j];
}
}
(a_pack, b_pack)
}
#[test]
fn tile_16x16_set_correct() {
if !crate::probe::scan().has_sme {
return;
}
let m = 16usize;
let n = 16usize;
let k = 8usize;
let a: Vec<f32> = (0..m * k).map(|i| ((i % 7) as f32) * 0.1).collect();
let b: Vec<f32> = (0..k * n).map(|i| ((i % 11) as f32) * 0.1).collect();
let mut c_sme = vec![0.0f32; m * n];
let mut c_ref = vec![0.0f32; m * n];
let (a_pack, b_pack) = pack_for_tile(&a, &b, k);
let stream = Stream::new().unwrap();
unsafe {
tile_16x16_f32(
&stream,
a_pack.as_ptr(),
b_pack.as_ptr(),
c_sme.as_mut_ptr(),
n,
k,
false,
);
}
drop(stream);
ref_matmul(&a, &b, &mut c_ref, m, n, k);
let mut max_err = 0.0f32;
for i in 0..m * n {
let e = (c_sme[i] - c_ref[i]).abs();
if e > max_err {
max_err = e;
}
}
assert!(
max_err < 1e-4,
"tile mismatch: max_err={max_err}, c_sme[0..4]={:?}, c_ref[0..4]={:?}",
&c_sme[..4],
&c_ref[..4]
);
}
#[test]
fn tile_16x16_accumulate_correct() {
if !crate::probe::scan().has_sme {
return;
}
let m = 16usize;
let n = 16usize;
let k = 4usize;
let a: Vec<f32> = (0..m * k).map(|i| ((i % 5) as f32) * 0.2).collect();
let b: Vec<f32> = (0..k * n).map(|i| ((i % 7) as f32) * 0.2).collect();
let mut c_sme: Vec<f32> = (0..m * n).map(|i| (i % 3) as f32).collect();
let mut c_ref: Vec<f32> = c_sme.clone();
let (a_pack, b_pack) = pack_for_tile(&a, &b, k);
let stream = Stream::new().unwrap();
unsafe {
tile_16x16_f32(
&stream,
a_pack.as_ptr(),
b_pack.as_ptr(),
c_sme.as_mut_ptr(),
n,
k,
true,
);
}
drop(stream);
for i in 0..m {
for j in 0..n {
let mut acc = 0.0f32;
for p in 0..k {
acc += a[i * k + p] * b[p * n + j];
}
c_ref[i * n + j] += acc;
}
}
let mut max_err = 0.0f32;
for i in 0..m * n {
let e = (c_sme[i] - c_ref[i]).abs();
if e > max_err {
max_err = e;
}
}
assert!(max_err < 1e-4, "accumulate mismatch: max_err={max_err}");
}
}