use std::borrow::Cow;
use std::sync::OnceLock;
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};
#[path = "simd_gemm.rs"]
mod simd_gemm;
#[path = "bf16_gemm.rs"]
mod bf16_gemm;
#[derive(Default)]
pub(crate) struct MatMulPrepack {
constant_inputs: [bool; 2],
dense: [OnceLock<Vec<f32>>; 2],
#[cfg(feature = "mlas")]
packed_b: OnceLock<mlas_sys::PackedB>,
}
impl MatMulPrepack {
pub(crate) fn set_constant_inputs(&mut self, constant_inputs: &[bool]) {
for (index, is_constant) in self.constant_inputs.iter_mut().enumerate() {
*is_constant = constant_inputs.get(index).copied().unwrap_or(false);
}
}
fn dense<'a>(&'a self, index: usize, view: &'a TensorView<'_>) -> Result<Cow<'a, [f32]>> {
if !self.constant_inputs[index] {
return to_dense_f32_widen("MatMul", view);
}
if let Some(cached) = self.dense[index].get() {
return Ok(Cow::Borrowed(cached));
}
match to_dense_f32_widen("MatMul", view)? {
Cow::Borrowed(dense) => Ok(Cow::Borrowed(dense)),
Cow::Owned(dense) => {
let _ = self.dense[index].set(dense);
Ok(Cow::Borrowed(
self.dense[index]
.get()
.expect("constant MatMul prepack was just initialized"),
))
}
}
}
#[cfg(feature = "mlas")]
fn packed_b(&self, b: &[f32], k: usize, n: usize) -> Option<&mlas_sys::PackedB> {
self.constant_inputs[1].then(|| {
self.packed_b
.get_or_init(|| mlas_sys::PackedB::new(n, k, b))
})
}
}
#[derive(Default)]
pub struct MatMulKernel {
prepack: MatMulPrepack,
}
pub struct MatMulFactory;
impl KernelFactory for MatMulFactory {
fn create(&self, _node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
Ok(Box::new(MatMulKernel::default()))
}
}
pub(crate) fn gemm(
a: &[f32],
b: &[f32],
c: &mut [f32],
m: usize,
k: usize,
n: usize,
) -> Result<()> {
gemm_with_backend(CpuBackend::auto_detect(), a, b, c, m, k, n)
}
#[cfg(feature = "mlas")]
fn gemm_packed(
a: &[f32],
packed: &mlas_sys::PackedB,
c: &mut [f32],
m: usize,
k: usize,
n: usize,
) -> Result<()> {
assert_eq!(packed.dimensions(), (k, n));
mlas_sys::sgemm_nn_packed(m, a, packed, c);
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn gemm_with_backend(
backend: CpuBackend,
a: &[f32],
b: &[f32],
c: &mut [f32],
m: usize,
k: usize,
n: usize,
) -> Result<()> {
match backend {
#[cfg(feature = "mlas")]
CpuBackend::Mlas => {
mlas_sys::sgemm_nn(m, n, k, a, b, c);
Ok(())
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
CpuBackend::SimdX86 => {
simd_gemm::sgemm_simd(a, b, c, m, k, n);
Ok(())
}
_ => {
gemm_generic(a, b, c, m, k, n);
Ok(())
}
}
}
const MR: usize = 4;
const NR: usize = 4;
const KC: usize = 256;
const MAX_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;
}
let threads = rayon::current_num_threads();
let mc = if threads <= 1 {
MAX_MC.min(m)
} else {
let target_tasks = threads.saturating_mul(2);
let rows = m.div_ceil(target_tasks).clamp(1, MAX_MC);
if rows == 1 {
1
} else {
rows.div_ceil(MR).saturating_mul(MR).min(MAX_MC)
}
};
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 set_constant_inputs(&mut self, constant_inputs: &[bool]) {
self.prepack.set_constant_inputs(constant_inputs);
}
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
self.execute_with_backend(inputs, outputs, CpuBackend::auto_detect())
}
fn supports_strided_input(&self, _input_idx: usize) -> bool {
true
}
fn estimated_flops(&self) -> Option<u64> {
None
}
}
impl MatMulKernel {
fn execute_with_backend(
&self,
inputs: &[TensorView],
outputs: &mut [TensorMut],
backend: CpuBackend,
) -> Result<()> {
check_arity("MatMul", inputs, outputs, 2, 2, 1)?;
let geom = matmul_geometry(&inputs[0], &inputs[1])?;
crate::trace::record_kernel_metrics(inputs, outputs, || {
(numel(&geom.batch_shape) as u64)
.saturating_mul(geom.m as u64)
.saturating_mul(geom.n as u64)
.saturating_mul(geom.k as u64)
.saturating_mul(2)
});
if let Some(result) = try_matmul_bf16_native(&inputs[0], &inputs[1], &geom)? {
return write_dense_f32_narrow("MatMul", &mut outputs[0], &result);
}
if output_is_direct_f32_eligible(&inputs[0], &inputs[1], &outputs[0]) {
let out = &mut outputs[0];
out.validate()?;
let numel = out.numel();
if numel != geom.result_len {
return Err(EpError::KernelFailed(format!(
"MatMul: output element count {numel} does not match result length {}",
geom.result_len
)));
}
if numel == 0 {
return Ok(());
}
let ptr = out.data_ptr_mut::<f32>();
let out_slice = unsafe { std::slice::from_raw_parts_mut(ptr, numel) };
return matmul_dense_into_with_backend(
&self.prepack.dense(0, &inputs[0])?,
&self.prepack.dense(1, &inputs[1])?,
&geom,
backend,
#[cfg(feature = "mlas")]
Some(&self.prepack),
out_slice,
);
}
let out =
matmul_dense_prepacked_with_backend(&inputs[0], &inputs[1], &self.prepack, backend)?;
write_dense_f32_narrow("MatMul", &mut outputs[0], &out)
}
}
fn output_is_direct_f32_eligible(a: &TensorView, b: &TensorView, out: &TensorMut) -> bool {
use onnx_runtime_ir::DataType;
use onnx_runtime_ir::DeviceType;
if out.device.device_type != DeviceType::Cpu
|| out.dtype != DataType::Float32
|| !out.is_contiguous()
{
return false;
}
let out_origin = (out.data.0 as *const u8).wrapping_add(out.byte_offset) as usize;
let out_end = out_origin.saturating_add(out.byte_size());
!std::iter::once(a)
.chain(std::iter::once(b))
.any(|input| output_overlaps_input(out_origin, out_end, input, out.device))
}
fn output_overlaps_input(
out_origin: usize,
out_end: usize,
input: &TensorView,
out_device: onnx_runtime_ir::DeviceId,
) -> bool {
if input.is_absent() || input.device != out_device {
return false;
}
let in_start = input.data_ptr::<u8>() as usize;
let in_end = in_start.saturating_add(input.byte_size());
out_origin < in_end && in_start < out_end
}
#[allow(unused_variables)]
fn try_matmul_bf16_native(
a: &TensorView,
b: &TensorView,
geom: &MatMulGeometry,
) -> Result<Option<Vec<f32>>> {
#[cfg(target_arch = "x86_64")]
{
use onnx_runtime_ir::DataType;
if a.dtype != DataType::BFloat16
|| b.dtype != DataType::BFloat16
|| !a.is_contiguous()
|| !b.is_contiguous()
|| !bf16_gemm::native_available()
{
return Ok(None);
}
a.validate()?;
b.validate()?;
let mut out = vec![0.0f32; geom.result_len];
if out.is_empty() {
return Ok(Some(out));
}
let a_len = a.numel();
let b_len = b.numel();
let a_bits = unsafe { std::slice::from_raw_parts(a.data_ptr::<u16>(), a_len) };
let b_bits = unsafe { std::slice::from_raw_parts(b.data_ptr::<u16>(), b_len) };
bf16_native_dense_into(a_bits, b_bits, geom, &mut out);
Ok(Some(out))
}
#[cfg(not(target_arch = "x86_64"))]
{
Ok(None)
}
}
#[cfg(target_arch = "x86_64")]
fn bf16_native_dense_into(a: &[u16], b: &[u16], geom: &MatMulGeometry, out: &mut [f32]) {
let (m, k, n) = (geom.m, geom.k, geom.n);
let (a_mat, b_mat, c_mat) = (geom.a_mat, geom.b_mat, geom.c_mat);
if geom.batch_shape.is_empty() {
bf16_gemm::gemm(a, b, out, m, k, n);
return;
}
let mut bidx = vec![0usize; geom.batch_shape.len()];
let mut b_out = 0usize;
loop {
let a_off = broadcast_offset(&bidx, &geom.a_batch, &geom.a_batch_strides) * a_mat;
let b_off = broadcast_offset(&bidx, &geom.b_batch, &geom.b_batch_strides) * b_mat;
bf16_gemm::gemm(
&a[a_off..a_off + a_mat],
&b[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(&geom.batch_shape, &mut bidx) {
break;
}
}
}
pub(crate) fn matmul_dense(a: &TensorView, b: &TensorView) -> Result<Vec<f32>> {
let geom = matmul_geometry(a, b)?;
if let Some(result) = try_matmul_bf16_native(a, b, &geom)? {
return Ok(result);
}
matmul_dense_impl_with_backend(
a,
b,
to_dense_f32_widen("MatMul", a)?,
to_dense_f32_widen("MatMul", b)?,
CpuBackend::auto_detect(),
#[cfg(feature = "mlas")]
None,
)
}
pub(crate) fn matmul_dense_prepacked(
a: &TensorView,
b: &TensorView,
prepack: &MatMulPrepack,
) -> Result<Vec<f32>> {
matmul_dense_prepacked_with_backend(a, b, prepack, CpuBackend::auto_detect())
}
fn matmul_dense_prepacked_with_backend(
a: &TensorView,
b: &TensorView,
prepack: &MatMulPrepack,
backend: CpuBackend,
) -> Result<Vec<f32>> {
let geom = matmul_geometry(a, b)?;
if let Some(result) = try_matmul_bf16_native(a, b, &geom)? {
return Ok(result);
}
matmul_dense_impl_with_backend(
a,
b,
prepack.dense(0, a)?,
prepack.dense(1, b)?,
backend,
#[cfg(feature = "mlas")]
Some(prepack),
)
}
fn matmul_dense_impl_with_backend(
a: &TensorView,
b: &TensorView,
a_dense: Cow<'_, [f32]>,
b_dense: Cow<'_, [f32]>,
backend: CpuBackend,
#[cfg(feature = "mlas")] prepack: Option<&MatMulPrepack>,
) -> Result<Vec<f32>> {
let geom = matmul_geometry(a, b)?;
let mut out = vec![0.0f32; geom.result_len];
matmul_dense_into_with_backend(
&a_dense,
&b_dense,
&geom,
backend,
#[cfg(feature = "mlas")]
prepack,
&mut out,
)?;
Ok(out)
}
struct MatMulGeometry {
m: usize,
k: usize,
n: usize,
a_mat: usize,
b_mat: usize,
c_mat: usize,
a_batch: Vec<usize>,
b_batch: Vec<usize>,
a_batch_strides: Vec<i64>,
b_batch_strides: Vec<i64>,
batch_shape: Vec<usize>,
#[cfg_attr(not(feature = "mlas"), allow(dead_code))]
b_promoted_rank: usize,
result_len: usize,
}
fn matmul_geometry(a: &TensorView, b: &TensorView) -> Result<MatMulGeometry> {
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].to_vec();
let b_batch = b_shape[..b_shape.len() - 2].to_vec();
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;
Ok(MatMulGeometry {
m,
k,
n,
a_mat,
b_mat,
c_mat,
a_batch,
b_batch,
a_batch_strides,
b_batch_strides,
batch_shape,
b_promoted_rank: b_shape.len(),
result_len: batch_count * c_mat,
})
}
fn matmul_dense_into_with_backend(
a_dense: &[f32],
b_dense: &[f32],
geom: &MatMulGeometry,
backend: CpuBackend,
#[cfg(feature = "mlas")] prepack: Option<&MatMulPrepack>,
out: &mut [f32],
) -> Result<()> {
if out.len() != geom.result_len {
return Err(EpError::KernelFailed(format!(
"MatMul: output buffer length {} does not match result length {}",
out.len(),
geom.result_len
)));
}
if out.is_empty() {
return Ok(());
}
let (m, k, n) = (geom.m, geom.k, geom.n);
let (a_mat, b_mat, c_mat) = (geom.a_mat, geom.b_mat, geom.c_mat);
#[cfg(feature = "mlas")]
let packed_b = if backend == CpuBackend::Mlas && geom.b_promoted_rank == 2 {
prepack.and_then(|prepack| prepack.packed_b(b_dense, k, n))
} else {
None
};
if geom.batch_shape.is_empty() {
#[cfg(feature = "mlas")]
if let Some(packed_b) = packed_b {
gemm_packed(a_dense, packed_b, out, m, k, n)?;
} else {
gemm_with_backend(backend, a_dense, b_dense, out, m, k, n)?;
}
#[cfg(not(feature = "mlas"))]
gemm_with_backend(backend, a_dense, b_dense, out, m, k, n)?;
} else {
let mut bidx = vec![0usize; geom.batch_shape.len()];
let mut b_out = 0usize;
loop {
let a_off = broadcast_offset(&bidx, &geom.a_batch, &geom.a_batch_strides) * a_mat;
let b_off = broadcast_offset(&bidx, &geom.b_batch, &geom.b_batch_strides) * b_mat;
let a_tile = &a_dense[a_off..a_off + a_mat];
let c_tile = &mut out[b_out * c_mat..b_out * c_mat + c_mat];
#[cfg(feature = "mlas")]
if let Some(packed_b) = packed_b {
gemm_packed(a_tile, packed_b, c_tile, m, k, n)?;
} else {
gemm_with_backend(
backend,
a_tile,
&b_dense[b_off..b_off + b_mat],
c_tile,
m,
k,
n,
)?;
}
#[cfg(not(feature = "mlas"))]
gemm_with_backend(
backend,
a_tile,
&b_dense[b_off..b_off + b_mat],
c_tile,
m,
k,
n,
)?;
b_out += 1;
if !next_index(&geom.batch_shape, &mut bidx) {
break;
}
}
}
Ok(())
}
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::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![58., 64., 139., 154.]);
}
#[cfg(feature = "tracing")]
#[test]
fn matmul_populates_active_trace_span_metrics() {
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]);
let (trace, events) = onnx_runtime_tracer::TraceContext::in_memory();
{
let _span = trace.span("MatMul", "compute");
MatMulKernel::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
}
let events = events.events();
let args = events[0].args.as_ref().expect("MatMul trace args");
assert_eq!(args["bytes"], 64);
assert_eq!(args["flops"], 24);
}
#[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::default()
.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::default()
.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::default()
.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::default()
.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::default()
.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_f16_preserves_near_tie_argmax_after_f32_accumulation() {
let a = Owned::f16(&[1, 4], &[1., 1., 1., 1.]);
#[allow(clippy::excessive_precision)]
let b = Owned::f16(
&[4, 2],
&[5.00390625, 5.0, 5.00390625, 5.0, 5.00390625, 5.0, 5.0, 5.0],
);
let mut out = Owned::zeros(onnx_runtime_ir::DataType::Float16, &[1, 2]);
MatMulKernel::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f16_as_f32(), vec![20.015625, 20.0]);
}
#[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::default()
.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.]
);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn matmul_bf16_native_matches_f64_reference_within_bf16_tolerance() {
if !bf16_gemm::native_available() {
eprintln!("skipping: host lacks avx512_bf16");
return;
}
const SHAPES: &[(usize, usize, usize)] = &[
(1, 2048, 512), (1, 100, 40), (32, 256, 64), (17, 130, 50), (4, 33, 3), ];
let mut state = 0x9E37_79B9_u32;
let mut next = || {
state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
((state >> 8) as f32 / 16_777_216.0 - 0.5) * 2.0 };
let mut worst_native_rel = 0.0f64;
let mut worst_ratio = 0.0f64;
for &(m, k, n) in SHAPES {
let a_f32: Vec<f32> = (0..m * k).map(|_| next()).collect();
let b_f32: Vec<f32> = (0..k * n).map(|_| next()).collect();
let a_bf: Vec<half::bf16> = a_f32.iter().map(|&v| half::bf16::from_f32(v)).collect();
let b_bf: Vec<half::bf16> = b_f32.iter().map(|&v| half::bf16::from_f32(v)).collect();
let a_bits: Vec<u16> = a_bf.iter().map(|v| v.to_bits()).collect();
let b_bits: Vec<u16> = b_bf.iter().map(|v| v.to_bits()).collect();
let a_wide: Vec<f32> = a_bf.iter().map(|v| v.to_f32()).collect();
let b_wide: Vec<f32> = b_bf.iter().map(|v| v.to_f32()).collect();
let mut reference = vec![0.0f64; m * n];
let mut upcast = vec![0.0f32; m * n];
for row in 0..m {
for col in 0..n {
let mut acc64 = 0.0f64;
let mut acc32 = 0.0f32;
for depth in 0..k {
acc64 += a_wide[row * k + depth] as f64 * b_wide[depth * n + col] as f64;
acc32 += a_wide[row * k + depth] * b_wide[depth * n + col];
}
reference[row * n + col] = acc64;
upcast[row * n + col] = acc32;
}
}
let mut native = vec![0.0f32; m * n];
bf16_gemm::gemm(&a_bits, &b_bits, &mut native, m, k, n);
let rel = |got: f32, want: f64| -> f64 {
let denom = want.abs().max(1.0);
(got as f64 - want).abs() / denom
};
let mut max_native = 0.0f64;
let mut max_upcast = 0.0f64;
for idx in 0..m * n {
max_native = max_native.max(rel(native[idx], reference[idx]));
max_upcast = max_upcast.max(rel(upcast[idx], reference[idx]));
}
worst_native_rel = worst_native_rel.max(max_native);
let ratio = max_native / max_upcast.max(1e-9);
worst_ratio = worst_ratio.max(ratio);
assert!(
max_native <= max_upcast * 4.0 + 1e-4,
"{m}x{k}@{k}x{n}: native rel {max_native} worse than upcast {max_upcast}"
);
assert!(
max_native <= 5e-2,
"{m}x{k}@{k}x{n}: native rel {max_native} exceeds bf16 tolerance"
);
}
println!(
"native bf16 GEMM: worst native-vs-f64 rel {worst_native_rel:.3e}, \
worst native/upcast ratio {worst_ratio:.3}"
);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn matmul_bf16_native_handles_k_tail_and_kernel_matches_kernel_path() {
if !bf16_gemm::native_available() {
eprintln!("skipping: host lacks avx512_bf16");
return;
}
let (m, k, n) = (3usize, 70usize, 5usize); let mut state = 0x1357_9BDF_u32;
let mut next = || {
state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
((state >> 8) as f32 / 16_777_216.0 - 0.5) * 1.0
};
let a_f32: Vec<f32> = (0..m * k).map(|_| next()).collect();
let b_f32: Vec<f32> = (0..k * n).map(|_| next()).collect();
let a = Owned::bf16(&[m, k], &a_f32);
let b = Owned::bf16(&[k, n], &b_f32);
let mut out = Owned::zeros(onnx_runtime_ir::DataType::BFloat16, &[m, n]);
MatMulKernel::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
let a_w: Vec<f32> = a_f32
.iter()
.map(|&v| half::bf16::from_f32(v).to_f32())
.collect();
let b_w: Vec<f32> = b_f32
.iter()
.map(|&v| half::bf16::from_f32(v).to_f32())
.collect();
let got = out.to_bf16_as_f32();
for row in 0..m {
for col in 0..n {
let mut acc = 0.0f64;
for depth in 0..k {
acc += a_w[row * k + depth] as f64 * b_w[depth * n + col] as f64;
}
let want = half::bf16::from_f32(acc as f32).to_f32();
let denom = want.abs().max(1.0);
let rel = (got[row * n + col] - want).abs() / denom;
assert!(rel <= 3e-2, "K-tail mismatch at ({row},{col}): rel {rel}");
}
}
}
#[cfg(target_arch = "x86_64")]
#[test]
#[ignore = "microbench: run explicitly with --ignored --nocapture"]
fn bench_bf16_native_vs_upcast() {
use std::time::Instant;
if !bf16_gemm::native_available() {
eprintln!("skipping bench: host lacks avx512_bf16");
return;
}
const SHAPES: &[(usize, usize, usize)] = &[
(1, 4096, 4096), (1, 4096, 11008), (128, 4096, 4096), (256, 2048, 8192), ];
let mut state = 0x5DEE_CE66_u32;
let mut next = || {
state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
((state >> 8) as f32 / 16_777_216.0 - 0.5) * 2.0
};
let median3 = |mut f: Box<dyn FnMut() -> f64>| {
let mut t = [f(), f(), f()];
t.sort_by(|a, b| a.partial_cmp(b).unwrap());
t[1]
};
println!("bf16 GEMM microbench (native _mm512_dpbf16_ps vs widen-to-f32 + SGEMM)");
for &(m, k, n) in SHAPES {
let a_bits: Vec<u16> = (0..m * k)
.map(|_| half::bf16::from_f32(next()).to_bits())
.collect();
let b_bits: Vec<u16> = (0..k * n)
.map(|_| half::bf16::from_f32(next()).to_bits())
.collect();
let flops = 2.0 * m as f64 * k as f64 * n as f64;
let native_ms = {
let (a, b) = (a_bits.clone(), b_bits.clone());
median3(Box::new(move || {
let mut c = vec![0.0f32; m * n];
let t = Instant::now();
bf16_gemm::gemm(&a, &b, &mut c, m, k, n);
std::hint::black_box(&c);
t.elapsed().as_secs_f64() * 1e3
}))
};
let upcast_ms = {
let (a, b) = (a_bits.clone(), b_bits.clone());
median3(Box::new(move || {
let t = Instant::now();
let a_f: Vec<f32> = a
.iter()
.map(|&x| half::bf16::from_bits(x).to_f32())
.collect();
let b_f: Vec<f32> = b
.iter()
.map(|&x| half::bf16::from_bits(x).to_f32())
.collect();
let mut c = vec![0.0f32; m * n];
gemm(&a_f, &b_f, &mut c, m, k, n).unwrap();
std::hint::black_box(&c);
t.elapsed().as_secs_f64() * 1e3
}))
};
let g = |ms: f64| flops / (ms * 1e-3) / 1e9;
println!(
" {m:>4}x{k}x{n}: native {native_ms:>8.3} ms ({:>7.1} GFLOP/s) \
upcast {upcast_ms:>8.3} ms ({:>7.1} GFLOP/s) speedup {:.2}x",
g(native_ms),
g(upcast_ms),
upcast_ms / native_ms,
);
}
}
#[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::default()
.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::default()
.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}");
}
#[cfg(feature = "mlas")]
#[test]
fn default_gemm_backend_matches_generic_reference() {
const SHAPES: &[(usize, usize, usize)] = &[
(1, 2048, 2048), (1, 2304, 9216), (1, 9216, 2304), (5, 128, 256),
(32, 512, 512),
];
assert_eq!(CpuBackend::auto_detect(), CpuBackend::Mlas);
let mut state = 0x1234_abcd_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
};
for &(m, k, n) in SHAPES {
let a: Vec<f32> = (0..m * k).map(|_| next_f32()).collect();
let b: Vec<f32> = (0..k * n).map(|_| next_f32()).collect();
let mut expected = vec![0.0; m * n];
let mut actual = vec![0.0; m * n];
gemm_generic(&a, &b, &mut expected, m, k, n);
gemm(&a, &b, &mut actual, m, k, n).unwrap();
let max_error = actual
.iter()
.zip(&expected)
.map(|(actual, expected)| (actual - expected).abs())
.fold(0.0f32, f32::max);
assert!(
max_error <= 1e-3,
"{m}x{k} @ {k}x{n}: default-backend max error {max_error} exceeds tolerance"
);
}
}
#[cfg(feature = "mlas")]
#[test]
fn mlas_gemm_matches_generic_for_matrix_and_batched_vector_tiles() {
const SHAPES: &[(usize, usize, usize)] = &[
(1, 1, 1),
(7, 13, 5),
(32, 512, 512),
(97, 11, 3),
(1, 13, 5),
(3, 13, 1),
];
let mut state = 0x5eed_1234_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
};
for &(m, k, n) in SHAPES {
let a: Vec<f32> = (0..m * k).map(|_| next_f32()).collect();
let b: Vec<f32> = (0..k * n).map(|_| next_f32()).collect();
let mut expected = vec![0.0; m * n];
let mut actual = vec![0.0; m * n];
gemm_generic(&a, &b, &mut expected, m, k, n);
gemm_with_backend(CpuBackend::Mlas, &a, &b, &mut actual, m, k, n).unwrap();
let max_error = actual
.iter()
.zip(&expected)
.map(|(actual, expected)| (actual - expected).abs())
.fold(0.0f32, f32::max);
assert!(
max_error <= 1e-3,
"{m}x{k} @ {k}x{n}: MLAS max error {max_error} exceeds tolerance"
);
}
}
#[cfg(feature = "mlas")]
#[test]
fn mlas_constant_b_packed_kernel_matches_unpacked_and_generic() {
for (m, k, n) in [(5usize, 17usize, 9usize), (33, 64, 48)] {
let a_data: Vec<f32> = (0..m * k)
.map(|i| ((i as f32 * 0.037).sin()) * 0.25)
.collect();
let b_data: Vec<f32> = (0..k * n)
.map(|i| ((i as f32 * 0.021 + 0.3).cos()) * 0.25)
.collect();
let a = Owned::f32(&[m, k], &a_data);
let b = Owned::f32(&[k, n], &b_data);
let mut out = Owned::zeros_f32(&[m, n]);
let mut kernel = MatMulKernel::default();
kernel.set_constant_inputs(&[false, true]);
kernel
.execute_with_backend(
&[a.view(), b.view()],
&mut [out.view_mut()],
CpuBackend::Mlas,
)
.unwrap();
let mut unpacked = vec![0.0; m * n];
let mut generic = vec![0.0; m * n];
gemm_with_backend(CpuBackend::Mlas, &a_data, &b_data, &mut unpacked, m, k, n).unwrap();
gemm_with_backend(CpuBackend::Generic, &a_data, &b_data, &mut generic, m, k, n)
.unwrap();
let packed = out.to_f32();
for (index, ((packed, unpacked), generic)) in
packed.iter().zip(&unpacked).zip(&generic).enumerate()
{
assert!(
(packed - unpacked).abs() <= 1e-4,
"{m}x{k}x{n} packed/unpacked mismatch at {index}: {packed} vs {unpacked}"
);
assert!(
(packed - generic).abs() <= 1e-3,
"{m}x{k}x{n} packed/generic mismatch at {index}: {packed} vs {generic}"
);
}
assert!(kernel.prepack.packed_b.get().is_some());
}
}
#[cfg(feature = "mlas")]
#[test]
fn mlas_constant_b_packed_buffer_is_reused() {
let mut kernel = MatMulKernel::default();
kernel.set_constant_inputs(&[false, true]);
let weight_data: Vec<f32> = (0..17 * 9)
.map(|i| ((i as f32 * 0.031).sin()) * 0.5)
.collect();
let weight = Owned::f16(&[17, 9], &weight_data);
let a1_data: Vec<f32> = (0..5 * 17).map(|i| i as f32 * 0.01).collect();
let a1 = Owned::f32(&[5, 17], &a1_data);
let mut out1 = Owned::zeros_f32(&[5, 9]);
kernel
.execute_with_backend(
&[a1.view(), weight.view()],
&mut [out1.view_mut()],
CpuBackend::Mlas,
)
.unwrap();
let packed_ptr = kernel.prepack.packed_b.get().unwrap() as *const mlas_sys::PackedB;
let dense_ptr = kernel.prepack.dense[1].get().unwrap().as_ptr();
let a2_data: Vec<f32> = (0..5 * 17)
.map(|i| ((i as f32 * 0.07).cos()) * 0.2)
.collect();
let a2 = Owned::f32(&[5, 17], &a2_data);
let mut out2 = Owned::zeros_f32(&[5, 9]);
kernel
.execute_with_backend(
&[a2.view(), weight.view()],
&mut [out2.view_mut()],
CpuBackend::Mlas,
)
.unwrap();
assert_eq!(
kernel.prepack.packed_b.get().unwrap() as *const mlas_sys::PackedB,
packed_ptr
);
assert_eq!(kernel.prepack.dense[1].get().unwrap().as_ptr(), dense_ptr);
assert!(kernel.prepack.dense[0].get().is_none());
assert_ne!(out1.to_f32(), out2.to_f32());
}
#[cfg(feature = "mlas")]
#[test]
fn mlas_packed_cache_requires_mlas_constant_unbatched_b() {
let (m, k, n) = (5usize, 17usize, 9usize);
let a_data: Vec<f32> = (0..m * k).map(|i| i as f32 * 0.01).collect();
let b_data: Vec<f32> = (0..k * n)
.map(|i| ((i as f32 * 0.02).sin()) * 0.1)
.collect();
let a = Owned::f32(&[m, k], &a_data);
let b = Owned::f32(&[k, n], &b_data);
let mut out = Owned::zeros_f32(&[m, n]);
let mut kernel = MatMulKernel::default();
kernel.set_constant_inputs(&[false, false]);
kernel
.execute_with_backend(
&[a.view(), b.view()],
&mut [out.view_mut()],
CpuBackend::Mlas,
)
.unwrap();
let mut expected = vec![0.0; m * n];
gemm_generic(&a_data, &b_data, &mut expected, m, k, n);
assert!(kernel.prepack.packed_b.get().is_none());
for (actual, expected) in out.to_f32().iter().zip(&expected) {
assert!((actual - expected).abs() <= 1e-3);
}
let mut generic_kernel = MatMulKernel::default();
generic_kernel.set_constant_inputs(&[false, true]);
let mut generic_out = Owned::zeros_f32(&[m, n]);
generic_kernel
.execute_with_backend(
&[a.view(), b.view()],
&mut [generic_out.view_mut()],
CpuBackend::Generic,
)
.unwrap();
assert!(generic_kernel.prepack.packed_b.get().is_none());
assert_eq!(generic_out.to_f32(), expected);
let batched_b_data = [b_data.clone(), b_data].concat();
let batched_a_data = [a_data.clone(), a_data].concat();
let batched_a = Owned::f32(&[2, m, k], &batched_a_data);
let batched_b = Owned::f32(&[2, k, n], &batched_b_data);
let mut batched_out = Owned::zeros_f32(&[2, m, n]);
let mut batched_kernel = MatMulKernel::default();
batched_kernel.set_constant_inputs(&[false, true]);
batched_kernel
.execute_with_backend(
&[batched_a.view(), batched_b.view()],
&mut [batched_out.view_mut()],
CpuBackend::Mlas,
)
.unwrap();
assert!(batched_kernel.prepack.packed_b.get().is_none());
for (actual, expected) in batched_out.to_f32().iter().zip(expected.iter().cycle()) {
assert!((actual - expected).abs() <= 1e-3);
}
}
#[cfg(feature = "mlas")]
#[test]
fn mlas_selects_a_float_kernel_on_x86_64() {
assert_ne!(mlas_sys::selected_float_kernel(), 0);
}
#[test]
fn constant_weight_prepack_reuses_weight_and_keeps_activation_live() {
let mut kernel = MatMulKernel::default();
kernel.set_constant_inputs(&[false, true]);
let weight = Owned::f16(&[2, 2], &[2., 0., 0., 3.]);
let a1 = Owned::f32(&[1, 2], &[1., 2.]);
let mut out1 = Owned::zeros_f32(&[1, 2]);
kernel
.execute(&[a1.view(), weight.view()], &mut [out1.view_mut()])
.unwrap();
assert_eq!(out1.to_f32(), vec![2., 6.]);
assert!(kernel.prepack.dense[1].get().is_some());
assert!(kernel.prepack.dense[0].get().is_none());
let cached_weight = kernel.prepack.dense[1].get().unwrap().as_ptr();
let a2 = Owned::f32(&[1, 2], &[4., 5.]);
let mut out2 = Owned::zeros_f32(&[1, 2]);
kernel
.execute(&[a2.view(), weight.view()], &mut [out2.view_mut()])
.unwrap();
assert_eq!(out2.to_f32(), vec![8., 15.]);
assert_eq!(
kernel.prepack.dense[1].get().unwrap().as_ptr(),
cached_weight
);
}
fn naive_matmul(a: &[f32], b: &[f32], m: usize, k: usize, n: usize) -> Vec<f32> {
let mut c = vec![0.0f32; m * n];
for i in 0..m {
for p in 0..k {
let aip = a[i * k + p];
for j in 0..n {
c[i * n + j] += aip * b[p * n + j];
}
}
}
c
}
#[test]
fn direct_f32_eligible_for_contiguous_cpu_output() {
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]);
assert!(output_is_direct_f32_eligible(
&a.view(),
&b.view(),
&out.view_mut()
));
}
#[test]
fn direct_f32_rejects_non_f32_output() {
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(onnx_runtime_ir::DataType::Float16, &[2, 2]);
assert!(!output_is_direct_f32_eligible(
&a.view(),
&b.view(),
&out.view_mut()
));
}
#[test]
fn direct_f32_2d_nonsquare_matches_reference() {
let a_data = [1., 2., 3., 4., 5., 6.];
let b_data = [1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11., 12.];
let a = Owned::f32(&[2, 3], &a_data);
let b = Owned::f32(&[3, 4], &b_data);
let mut out = Owned::zeros_f32(&[2, 4]);
MatMulKernel::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), naive_matmul(&a_data, &b_data, 2, 3, 4));
}
#[test]
fn direct_f32_batched_and_broadcast_match_reference() {
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::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![1., 2., 3., 4., 10., 12., 14., 16.]);
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::default()
.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 direct_f32_matrix_times_vector() {
let a = Owned::f32(&[2, 3], &[1., 2., 3., 4., 5., 6.]);
let b = Owned::f32(&[3], &[7., 9., 11.]);
let mut out = Owned::zeros_f32(&[2]);
MatMulKernel::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![58., 139.]);
}
#[test]
fn direct_f32_vector_times_vector_scalar_result() {
let a = Owned::f32(&[3], &[1., 2., 3.]);
let b = Owned::f32(&[3], &[4., 5., 6.]);
let mut out = Owned::zeros_f32(&[]);
MatMulKernel::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![32.]);
}
#[test]
fn direct_f32_zero_sized_result_writes_nothing() {
let a = Owned::f32(&[0, 2, 3], &[]);
let b = Owned::f32(&[0, 3, 2], &[]);
let mut out = Owned::zeros_f32(&[0, 2, 2]);
assert!(output_is_direct_f32_eligible(
&a.view(),
&b.view(),
&out.view_mut()
));
MatMulKernel::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert!(out.to_f32().is_empty());
}
#[test]
fn strided_f32_output_takes_fallback_and_is_correct() {
let a = Owned::f32(&[2, 2], &[1., 2., 3., 4.]);
let b = Owned::f32(&[2, 2], &[5., 6., 7., 8.]);
let mut out = Owned::zeros_f32(&[2, 3]).with_view(&[2, 2], &[3, 1]);
assert!(!out.view_mut().is_contiguous());
assert!(!output_is_direct_f32_eligible(
&a.view(),
&b.view(),
&out.view_mut()
));
MatMulKernel::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![19., 22., 0., 43., 50., 0.]);
}
#[test]
fn mismatched_output_length_errors_before_write() {
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, 3]);
let err = MatMulKernel::default()
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap_err();
assert!(format!("{err}").contains("does not match result length"));
assert_eq!(out.to_f32(), vec![0.; 6]);
}
#[test]
fn output_overlaps_input_helper_detects_ranges() {
use onnx_runtime_ep_api::{DevicePtr, DevicePtrMut};
use onnx_runtime_ir::{DataType, DeviceId, DeviceType};
let buf = [0.0f32; 8];
let shape = [2usize, 2];
let strides = compute_contiguous_strides(&shape);
let base = buf.as_ptr() as usize;
let bytes = 4 * 4;
let input = TensorView::new(
DevicePtr(buf.as_ptr() as *const std::ffi::c_void),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
);
assert!(output_overlaps_input(
base + 8,
base + 8 + bytes,
&input,
DeviceId::cpu()
));
assert!(!output_overlaps_input(
base + bytes,
base + 2 * bytes,
&input,
DeviceId::cpu()
));
assert!(!output_overlaps_input(
base,
base + bytes,
&TensorView::absent(DataType::Float32),
DeviceId::cpu()
));
assert!(!output_overlaps_input(
base,
base + bytes,
&input,
DeviceId::new(DeviceType::Cuda, 0)
));
let _ = DevicePtrMut(std::ptr::null_mut());
}
#[test]
fn aliasing_output_takes_fallback_and_is_correct() {
use onnx_runtime_ep_api::{DevicePtr, DevicePtrMut};
use onnx_runtime_ir::{DataType, DeviceId};
let mut shared = vec![1.0f32, 2.0, 3.0, 4.0];
let b_buf = [0.0f32, 1.0, 1.0, 0.0];
let shape = vec![2usize, 2];
let strides = compute_contiguous_strides(&shape);
let a_ptr = shared.as_ptr() as *const std::ffi::c_void;
let c_ptr = shared.as_mut_ptr() as *mut std::ffi::c_void;
let a = TensorView::new(
DevicePtr(a_ptr),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
);
let b = TensorView::new(
DevicePtr(b_buf.as_ptr() as *const std::ffi::c_void),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
);
let c = TensorMut::new(
DevicePtrMut(c_ptr),
DataType::Float32,
&shape,
&strides,
DeviceId::cpu(),
);
assert!(!output_is_direct_f32_eligible(&a, &b, &c));
MatMulKernel::default().execute(&[a, b], &mut [c]).unwrap();
assert_eq!(shared, vec![2.0, 1.0, 4.0, 3.0]);
}
}