#[must_use]
pub fn smmla_packed_len(rows: usize, k: usize) -> usize {
rows.div_ceil(2) * k.div_ceil(8) * 16
}
#[must_use]
pub fn smmla_pack_panels(
src: &[i8],
base_row: usize,
rows: usize,
k: usize,
src_k: usize,
) -> (Vec<i8>, usize, usize) {
let row_pairs = rows.div_ceil(2);
let kb = k.div_ceil(8);
let mut packed = vec![0i8; row_pairs * kb * 16];
for p in 0..row_pairs {
for block in 0..kb {
let panel = (p * kb + block) * 16;
let kcol = block * 8;
let kvalid = (k - kcol).min(8); for sub in 0..2 {
let row = p * 2 + sub;
if row >= rows {
continue; }
let src_off = (base_row + row) * src_k + kcol;
let dst_off = panel + sub * 8;
packed[dst_off..dst_off + kvalid].copy_from_slice(&src[src_off..src_off + kvalid]);
}
}
}
(packed, row_pairs, kb)
}
pub fn smmla_unpack_panels(packed: &[i8], rows: usize, k: usize) -> Result<Vec<i8>, String> {
let row_pairs = rows.div_ceil(2);
let kb = k.div_ceil(8);
let want = row_pairs * kb * 16;
if packed.len() != want {
return Err(format!(
"SMMLA panel stream is {} bytes; [{rows}, {k}] requires {want}",
packed.len()
));
}
let mut out = vec![0i8; rows * k];
for p in 0..row_pairs {
for block in 0..kb {
let panel = (p * kb + block) * 16;
let kcol = block * 8;
let kvalid = (k - kcol).min(8);
for sub in 0..2 {
let row = p * 2 + sub;
if row >= rows {
continue;
}
let dst_off = row * k + kcol;
out[dst_off..dst_off + kvalid]
.copy_from_slice(&packed[panel + sub * 8..panel + sub * 8 + kvalid]);
}
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
fn filled(rows: usize, k: usize) -> Vec<i8> {
(0..rows * k)
.map(|i| ((i as i64 * 73 + 11) % 255 - 127) as i8)
.collect()
}
#[test]
fn pack_unpack_round_trips_byte_exact() {
for (rows, k) in [
(8, 16),
(7, 16),
(8, 13),
(5, 3),
(1, 1),
(2, 8),
(64, 1280),
] {
let src = filled(rows, k);
let (packed, pairs, kb) = smmla_pack_panels(&src, 0, rows, k, k);
assert_eq!(packed.len(), pairs * kb * 16);
assert_eq!(packed.len(), smmla_packed_len(rows, k));
let back = smmla_unpack_panels(&packed, rows, k).expect("len agrees");
assert_eq!(back, src, "round-trip must be lossless for [{rows}, {k}]");
println!(
r#"{{"check":"smmla_pack_roundtrip","rows":{rows},"k":{k},"packed_len":{},"result":"pass"}}"#,
packed.len()
);
}
}
#[test]
fn packing_is_deterministic() {
let src = filled(16, 64);
let (a, _, _) = smmla_pack_panels(&src, 0, 16, 64, 64);
let (b, _, _) = smmla_pack_panels(&src, 0, 16, 64, 64);
assert_eq!(a, b);
}
#[test]
fn padding_lanes_are_zero() {
let (rows, k) = (3, 5); let src = vec![7i8; rows * k];
let (packed, pairs, kb) = smmla_pack_panels(&src, 0, rows, k, k);
assert_eq!((pairs, kb), (2, 1));
let panel = kb * 16; assert!(packed[panel + 8..panel + 16].iter().all(|&b| b == 0));
assert!(packed[panel + k..panel + 8].iter().all(|&b| b == 0));
assert!(packed[panel..panel + k].iter().all(|&b| b == 7));
}
#[test]
fn real_decoder_shapes_pack_without_padding() {
for (n, k) in [
(1280usize, 1280usize),
(6848, 1280),
(256, 6848),
(3072, 1024),
(1024, 2816),
(1600, 960),
(960, 2560),
(3072, 768),
(768, 3072),
] {
assert_eq!(
smmla_packed_len(n, k),
n * k,
"[{n}, {k}] must tile cleanly (n even, k % 8 == 0)"
);
}
}
#[test]
fn even_base_region_pack_is_a_slice_of_the_full_pack() {
let (rows, k) = (32, 24);
let src = filled(rows, k);
let (full, _pairs, kb) = smmla_pack_panels(&src, 0, rows, k, k);
for (base, cnt) in [(0usize, 8usize), (8, 8), (16, 16), (24, 8), (2, 30)] {
let (region, rpairs, rkb) = smmla_pack_panels(&src, base, cnt, k, k);
assert_eq!(rkb, kb);
let off = (base / 2) * kb * 16;
assert_eq!(
region,
full[off..off + rpairs * kb * 16],
"even-base region [{base}, {base}+{cnt}) must alias the full pack"
);
}
}
}