use onnx_runtime_ep_api::{EpError, Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::{Node, broadcast_shapes, compute_contiguous_strides};
use rayon::prelude::*;
use super::check_arity;
use crate::backend::CpuBackend;
use crate::dtype::{to_dense_f32_widen, write_dense_f32_narrow};
use crate::strided::{next_index, numel};
pub struct MatMulKernel;
pub struct MatMulFactory;
impl KernelFactory for MatMulFactory {
fn create(&self, _node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
Ok(Box::new(MatMulKernel))
}
}
fn gemm(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) -> Result<()> {
match CpuBackend::auto_detect() {
#[cfg(feature = "onednn")]
CpuBackend::OneDnn => crate::kernels::onednn::sgemm(a, b, c, m, k, n),
_ => {
gemm_generic(a, b, c, m, k, n);
Ok(())
}
}
}
const MR: usize = 4;
const NR: usize = 4;
const KC: usize = 256;
const MC: usize = 64;
fn gemm_generic(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) {
if m == 0 || n == 0 {
return;
}
c.par_chunks_mut(MC * n)
.enumerate()
.for_each(|(blk, c_block)| {
let i0 = blk * MC;
let rows = c_block.len() / n; let a_block = &a[i0 * k..i0 * k + rows * k];
gemm_block(a_block, b, c_block, rows, k, n);
});
}
fn gemm_block(a: &[f32], b: &[f32], c: &mut [f32], rows: usize, k: usize, n: usize) {
for v in c.iter_mut() {
*v = 0.0;
}
let mut kk = 0;
while kk < k {
let kc = KC.min(k - kk);
let mut i = 0;
while i < rows {
let mr = MR.min(rows - i);
let mut j = 0;
while j < n {
let nr = NR.min(n - j);
micro_kernel(a, b, c, k, n, i, j, kk, kc, mr, nr);
j += NR;
}
i += MR;
}
kk += KC;
}
}
#[inline]
#[allow(clippy::too_many_arguments)]
fn micro_kernel(
a: &[f32],
b: &[f32],
c: &mut [f32],
k: usize,
n: usize,
i: usize,
j: usize,
kk: usize,
kc: usize,
mr: usize,
nr: usize,
) {
let mut acc = [[0.0f32; NR]; MR];
for p in kk..kk + kc {
let brow = &b[p * n + j..p * n + j + nr];
for (ii, acc_row) in acc.iter_mut().enumerate().take(mr) {
let aik = a[(i + ii) * k + p];
for (jj, acc_v) in acc_row.iter_mut().enumerate().take(nr) {
*acc_v += aik * brow[jj];
}
}
}
for (ii, acc_row) in acc.iter().enumerate().take(mr) {
let c_row = &mut c[(i + ii) * n + j..(i + ii) * n + j + nr];
for (jj, cv) in c_row.iter_mut().enumerate().take(nr) {
*cv += acc_row[jj];
}
}
}
impl Kernel for MatMulKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
check_arity("MatMul", inputs, outputs, 2, 2, 1)?;
let out = matmul_dense(&inputs[0], &inputs[1])?;
write_dense_f32_narrow("MatMul", &mut outputs[0], &out)
}
fn supports_strided_input(&self, _input_idx: usize) -> bool {
true
}
fn estimated_flops(&self) -> Option<u64> {
None
}
}
pub(crate) fn matmul_dense(a: &TensorView, b: &TensorView) -> Result<Vec<f32>> {
let a_dense = to_dense_f32_widen("MatMul", a)?;
let b_dense = to_dense_f32_widen("MatMul", b)?;
let a_raw = a.shape;
let b_raw = b.shape;
let a_1d = a_raw.len() == 1;
let b_1d = b_raw.len() == 1;
let a_shape: Vec<usize> = if a_1d {
vec![1, a_raw[0]]
} else {
a_raw.to_vec()
};
let b_shape: Vec<usize> = if b_1d {
vec![b_raw[0], 1]
} else {
b_raw.to_vec()
};
if a_shape.len() < 2 || b_shape.len() < 2 {
return Err(EpError::KernelFailed(
"MatMul: operands must be at least 1-D".into(),
));
}
let m = a_shape[a_shape.len() - 2];
let k = a_shape[a_shape.len() - 1];
let k2 = b_shape[b_shape.len() - 2];
let n = b_shape[b_shape.len() - 1];
if k != k2 {
return Err(EpError::KernelFailed(format!(
"MatMul: inner dims disagree ({k} vs {k2})"
)));
}
let a_batch = &a_shape[..a_shape.len() - 2];
let b_batch = &b_shape[..b_shape.len() - 2];
let batch_shape = broadcast_shapes(a_batch, b_batch)?;
let batch_count = numel(&batch_shape);
let a_batch_strides = compute_contiguous_strides(a_batch);
let b_batch_strides = compute_contiguous_strides(b_batch);
let a_mat = m * k;
let b_mat = k * n;
let c_mat = m * n;
let mut out = vec![0.0f32; batch_count * c_mat];
if out.is_empty() {
return Ok(out);
}
if batch_shape.is_empty() {
gemm(&a_dense, &b_dense, &mut out, m, k, n)?;
} else {
let mut bidx = vec![0usize; batch_shape.len()];
let mut b_out = 0usize;
loop {
let a_off = broadcast_offset(&bidx, a_batch, &a_batch_strides) * a_mat;
let b_off = broadcast_offset(&bidx, b_batch, &b_batch_strides) * b_mat;
gemm(
&a_dense[a_off..a_off + a_mat],
&b_dense[b_off..b_off + b_mat],
&mut out[b_out * c_mat..b_out * c_mat + c_mat],
m,
k,
n,
)?;
b_out += 1;
if !next_index(&batch_shape, &mut bidx) {
break;
}
}
}
Ok(out)
}
fn broadcast_offset(bidx: &[usize], batch: &[usize], batch_strides: &[i64]) -> usize {
let out_rank = bidx.len();
let mut off = 0i64;
for axis in 0..batch.len() {
let out_axis = axis + (out_rank - batch.len());
let i = if batch[axis] == 1 { 0 } else { bidx[out_axis] };
off += batch_strides[axis] * i as i64;
}
off as usize
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernels::testutil::Owned;
#[test]
fn matmul_zero_batch_returns_empty_without_panicking() {
let a = Owned::f32(&[0, 1, 1], &[]);
let b = Owned::f32(&[0, 1, 1], &[]);
let out = matmul_dense(&a.view(), &b.view()).unwrap();
assert!(out.is_empty());
}
#[test]
fn matmul_2x3_times_3x2() {
let a = Owned::f32(&[2, 3], &[1., 2., 3., 4., 5., 6.]);
let b = Owned::f32(&[3, 2], &[7., 8., 9., 10., 11., 12.]);
let mut out = Owned::zeros_f32(&[2, 2]);
MatMulKernel
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![58., 64., 139., 154.]);
}
#[test]
fn matmul_with_transposed_b_view() {
let a = Owned::f32(&[2, 3], &[1., 2., 3., 4., 5., 6.]);
let b = Owned::f32(&[2, 3], &[7., 9., 11., 8., 10., 12.]).with_view(&[3, 2], &[1, 3]);
let mut out = Owned::zeros_f32(&[2, 2]);
MatMulKernel
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![58., 64., 139., 154.]);
}
#[test]
fn matmul_batched() {
let a = Owned::f32(&[2, 2, 2], &[1., 2., 3., 4., 5., 6., 7., 8.]);
let b = Owned::f32(&[2, 2, 2], &[1., 0., 0., 1., 2., 0., 0., 2.]);
let mut out = Owned::zeros_f32(&[2, 2, 2]);
MatMulKernel
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![1., 2., 3., 4., 10., 12., 14., 16.]);
}
#[test]
fn matmul_broadcast_batch() {
let a = Owned::f32(&[2, 2, 2], &[1., 2., 3., 4., 5., 6., 7., 8.]);
let b = Owned::f32(&[2, 2], &[1., 0., 0., 1.]); let mut out = Owned::zeros_f32(&[2, 2, 2]);
MatMulKernel
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![1., 2., 3., 4., 5., 6., 7., 8.]);
}
#[test]
fn matmul_vector_times_matrix() {
let a = Owned::f32(&[3], &[1., 2., 3.]);
let b = Owned::f32(&[3, 2], &[7., 8., 9., 10., 11., 12.]);
let mut out = Owned::zeros_f32(&[2]);
MatMulKernel
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![58., 64.]);
}
#[test]
fn matmul_f16_accumulates_in_f32() {
let a = Owned::f16(&[2, 3], &[1., 2., 3., 4., 5., 6.]);
let b = Owned::f16(&[3, 2], &[7., 8., 9., 10., 11., 12.]);
let mut out = Owned::zeros(onnx_runtime_ir::DataType::Float16, &[2, 2]);
MatMulKernel
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f16_as_f32(), vec![58., 64., 139., 154.]);
}
#[test]
fn matmul_bf16_batched() {
let a = Owned::bf16(&[2, 2, 2], &[1., 2., 3., 4., 5., 6., 7., 8.]);
let b = Owned::bf16(&[2, 2, 2], &[1., 0., 0., 1., 2., 0., 0., 2.]);
let mut out = Owned::zeros(onnx_runtime_ir::DataType::BFloat16, &[2, 2, 2]);
MatMulKernel
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(
out.to_bf16_as_f32(),
vec![1., 2., 3., 4., 10., 12., 14., 16.]
);
}
#[test]
fn matmul_rejects_integer_dtype_with_rule1() {
let a = Owned::i32(&[2, 2], &[1, 2, 3, 4]);
let b = Owned::i32(&[2, 2], &[1, 0, 0, 1]);
let mut out = Owned::zeros(onnx_runtime_ir::DataType::Int32, &[2, 2]);
let err = MatMulKernel
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap_err();
assert!(format!("{err}").contains("WHAT"));
}
#[test]
#[allow(clippy::needless_range_loop)]
fn matmul_generic_block_boundaries_match_naive_reference() {
const SHAPES: &[(usize, usize, usize)] = &[
(65, 257, 70),
(128, 300, 200),
(100, 64, 4),
(4, 256, 4),
(1, 512, 1),
(200, 1, 200),
];
const ABS_TOLERANCE: f32 = 1e-3;
let mut state = 0x1234_5678_u32;
let mut next_f32 = || {
state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
((state >> 8) as f32 / 16_777_216.0 - 0.5) * 0.25
};
let mut overall_max_abs_error = 0.0f32;
for &(m, k, n) in SHAPES {
let a_data: Vec<f32> = (0..m * k).map(|_| next_f32()).collect();
let b_data: Vec<f32> = (0..k * n).map(|_| next_f32()).collect();
let mut reference = vec![0.0f32; m * n];
for row in 0..m {
for col in 0..n {
let mut sum = 0.0f32;
for depth in 0..k {
sum += a_data[row * k + depth] * b_data[depth * n + col];
}
reference[row * n + col] = sum;
}
}
let a = Owned::f32(&[m, k], &a_data);
let b = Owned::f32(&[k, n], &b_data);
let mut out = Owned::zeros_f32(&[m, n]);
MatMulKernel
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
let actual = out.to_f32();
let max_abs_error = actual
.iter()
.zip(&reference)
.map(|(actual, expected)| (actual - expected).abs())
.fold(0.0f32, f32::max);
overall_max_abs_error = overall_max_abs_error.max(max_abs_error);
assert!(
max_abs_error <= ABS_TOLERANCE,
"{m}x{k} @ {k}x{n}: max abs error {max_abs_error} exceeds {ABS_TOLERANCE}"
);
}
println!("generic MatMul max abs error: {overall_max_abs_error}");
}
}