const MR: usize = 4;
const NR: usize = 8;
const PANEL_BYTES: usize = 256 * 1024;
pub fn linear_packed(
x: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
m: usize,
k: usize,
n: usize,
out: &mut [f32],
) {
assert_eq!(x.len(), m * k, "x must be [m, k]");
assert_eq!(weight.len(), n * k, "weight must be [n, k]");
assert_eq!(out.len(), m * n, "out must be [m, n]");
if let Some(bias) = bias {
assert_eq!(bias.len(), n, "bias must be [n]");
}
unsafe {
linear_packed_range(x, weight, bias, m, k, n, 0, n, out.as_mut_ptr());
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) unsafe fn linear_packed_range(
x: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
m: usize,
k: usize,
n: usize,
col_start: usize,
col_end: usize,
out: *mut f32,
) {
for row in 0..m {
for column in col_start..col_end {
unsafe {
*out.add(row * n + column) = bias.map_or(0.0, |values| values[column]);
}
}
}
if m == 0 || k == 0 || col_start >= col_end {
return; }
let columns = col_end - col_start;
let m_full = m - m % MR;
let n_full = col_start + columns - columns % NR;
let panel_columns = {
let fitting = PANEL_BYTES / (k.max(1) * size_of::<f32>());
(fitting / NR).max(1) * NR
};
thread_local! {
static PANEL_SCRATCH: std::cell::RefCell<Vec<f32>> =
const { std::cell::RefCell::new(Vec::new()) };
}
PANEL_SCRATCH.with(|scratch| {
let mut panel_guard = scratch.borrow_mut();
if panel_guard.len() < k * NR {
panel_guard.resize(k * NR, 0.0);
}
let panel = &mut panel_guard[..k * NR];
let mut jc = col_start;
while jc < n_full {
let jc_end = (jc + panel_columns).min(n_full);
let mut j0 = jc;
while j0 < jc_end {
for (column, offset) in (j0..j0 + NR).enumerate() {
let source = &weight[offset * k..offset * k + k];
for (depth, &value) in source.iter().enumerate() {
panel[depth * NR + column] = value;
}
}
let mut i0 = 0;
while i0 < m_full {
unsafe { accumulate_tile::<MR>(x, panel, out, i0, j0, k, n) };
i0 += MR;
}
for row in m_full..m {
unsafe { accumulate_tile::<1>(x, panel, out, row, j0, k, n) };
}
j0 += NR;
}
jc += panel_columns;
}
for row in 0..m {
let x_row = &x[row * k..row * k + k];
for column in n_full..col_end {
let w_row = &weight[column * k..column * k + k];
let mut sum = 0.0_f32;
for depth in 0..k {
sum += x_row[depth] * w_row[depth];
}
unsafe { *out.add(row * n + column) += sum };
}
}
});
}
#[inline]
unsafe fn accumulate_tile<const ROWS: usize>(
x: &[f32],
panel: &[f32],
out: *mut f32,
i0: usize,
j0: usize,
k: usize,
n: usize,
) {
let mut acc = [[0.0_f32; NR]; ROWS];
for depth in 0..k {
let weights = &panel[depth * NR..depth * NR + NR];
for (row, slots) in acc.iter_mut().enumerate() {
let value = x[(i0 + row) * k + depth];
for (slot, &weight) in slots.iter_mut().zip(weights) {
*slot += value * weight;
}
}
}
for (row, slots) in acc.iter().enumerate() {
let base = (i0 + row) * n + j0;
for (column, &value) in slots.iter().enumerate() {
unsafe { *out.add(base + column) += value };
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn scalar_reference(
x: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
m: usize,
k: usize,
n: usize,
) -> Vec<f32> {
let mut out = vec![0.0_f32; m * n];
for row in 0..m {
for column in 0..n {
let mut sum = 0.0_f32;
for depth in 0..k {
sum += x[row * k + depth] * weight[column * k + depth];
}
out[row * n + column] = bias.map_or(sum, |b| sum + b[column]);
}
}
out
}
fn deterministic(count: usize, seed: u64) -> Vec<f32> {
let mut state = seed | 1;
(0..count)
.map(|_| {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
((state >> 40) as f32 / 2048.0) - 0.5
})
.collect()
}
#[test]
fn packed_matches_scalar_bit_for_bit() {
let shapes = [
(1, 1, 1),
(1, 16, 8),
(3, 5, 7),
(4, 8, 8),
(5, 9, 9),
(8, 64, 16),
(7, 128, 13),
(16, 512, 32),
(2, 0, 4),
(4, 1, 8),
(9, 1024, 24),
];
for (index, &(m, k, n)) in shapes.iter().enumerate() {
let x = deterministic(m * k, 0x51ED_0000 + index as u64);
let weight = deterministic(n * k, 0xA113_0000 + index as u64);
let bias = deterministic(n, 0xB1A5_0000 + index as u64);
for carry_bias in [None, Some(&bias[..])] {
let expected = scalar_reference(&x, &weight, carry_bias, m, k, n);
let mut actual = vec![0.0_f32; m * n];
linear_packed(&x, &weight, carry_bias, m, k, n, &mut actual);
assert_eq!(
actual.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
expected.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"m={m} k={k} n={n} bias={}: packed GEMM diverged from the scalar reference",
carry_bias.is_some()
);
}
}
}
#[test]
fn every_partition_count_reproduces_the_serial_bits() {
let shapes = [
(32, 7168, 1536),
(72, 512, 512),
(48, 512, 1024),
(17, 96, 40),
];
for (index, &(m, k, n)) in shapes.iter().enumerate() {
let x = deterministic(m * k, 0x9E11_0000 + index as u64);
let weight = deterministic(n * k, 0x7A31_0000 + index as u64);
let bias = deterministic(n, 0x1CE5_0000 + index as u64);
let mut serial = vec![0.0_f32; m * n];
linear_packed(&x, &weight, Some(&bias), m, k, n, &mut serial);
for partitions in [1, 2, 3, 5, 6, 8] {
let mut parallel = vec![0.0_f32; m * n];
let chunk = n.div_ceil(partitions).next_multiple_of(NR);
for worker in 0..partitions {
let start = (worker * chunk).min(n);
let end = ((worker + 1) * chunk).min(n);
if start >= end {
continue;
}
unsafe {
linear_packed_range(
&x,
&weight,
Some(&bias),
m,
k,
n,
start,
end,
parallel.as_mut_ptr(),
);
}
}
assert_eq!(
parallel.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
serial.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"m={m} k={k} n={n} partitions={partitions}"
);
}
}
}
#[test]
fn column_partitions_are_bit_identical_to_the_whole() {
let (m, k, n) = (6, 96, 24);
let x = deterministic(m * k, 0xC0F1);
let weight = deterministic(n * k, 0xD00D);
let mut whole = vec![0.0_f32; m * n];
linear_packed(&x, &weight, None, m, k, n, &mut whole);
for split in [8, 16] {
let columns = split;
let slice: Vec<f32> = weight[..columns * k].to_vec();
let mut part = vec![0.0_f32; m * columns];
linear_packed(&x, &slice, None, m, k, columns, &mut part);
for row in 0..m {
for column in 0..columns {
assert_eq!(
part[row * columns + column].to_bits(),
whole[row * n + column].to_bits(),
"split={split} row={row} column={column}"
);
}
}
}
}
}