use onnx_runtime_ep_api::{Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::Node;
use super::add::broadcast_apply;
use super::matmul::{MatMulPrepack, matmul_dense_prepacked};
use super::relu::relu_in_place;
use super::{check_arity, to_dense_f32, write_dense_f32};
#[derive(Default)]
pub struct FusedGemmKernel {
prepack: MatMulPrepack,
}
pub struct FusedGemmFactory;
impl KernelFactory for FusedGemmFactory {
fn create(&self, _node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
Ok(Box::new(FusedGemmKernel::default()))
}
}
impl Kernel for FusedGemmKernel {
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<()> {
check_arity("FusedGemm", inputs, outputs, 3, 3, 1)?;
let mut out = matmul_dense_prepacked(&inputs[0], &inputs[1], &self.prepack)?;
let bias = to_dense_f32(&inputs[2])?;
let bias_shape = inputs[2].shape;
let out_shape = outputs[0].shape.to_vec();
broadcast_apply(&bias, bias_shape, &out_shape, |i, v| out[i] += v)?;
relu_in_place(&mut out);
write_dense_f32(&mut outputs[0], &out)
}
fn supports_strided_input(&self, _input_idx: usize) -> bool {
true
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernels::testutil::Owned;
#[test]
fn relu_clamps_negative_prebias_sums() {
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 bias = Owned::f32(&[2], &[-60., 20.]);
let mut out = Owned::zeros_f32(&[2, 2]);
FusedGemmKernel::default()
.execute(&[a.view(), b.view(), bias.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![0., 84., 79., 174.]);
}
#[test]
fn matches_matmul_then_add_then_relu() {
use crate::kernels::matmul::MatMulKernel;
use crate::kernels::relu::ReluKernel;
let a = Owned::f32(&[2, 4], &[1., -2., 3., -4., 5., -6., 7., -8.]);
let b = Owned::f32(
&[4, 3],
&[1., -2., 3., -4., 5., -6., 7., -8., 9., -10., 11., -12.],
);
let bias = Owned::f32(&[3], &[0.5, -100.0, 2.0]);
let mut mm = Owned::zeros_f32(&[2, 3]);
MatMulKernel::default()
.execute(&[a.view(), b.view()], &mut [mm.view_mut()])
.unwrap();
let mut biased = mm.to_f32();
for row in 0..2 {
for col in 0..3 {
biased[row * 3 + col] += [0.5, -100.0, 2.0][col];
}
}
let biased_owned = Owned::f32(&[2, 3], &biased);
let mut expect = Owned::zeros_f32(&[2, 3]);
ReluKernel
.execute(&[biased_owned.view()], &mut [expect.view_mut()])
.unwrap();
let mut out = Owned::zeros_f32(&[2, 3]);
FusedGemmKernel::default()
.execute(&[a.view(), b.view(), bias.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), expect.to_f32());
assert!(expect.to_f32().contains(&0.0));
}
#[test]
fn batched_matmul_with_bias_and_relu() {
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 bias = Owned::f32(&[2], &[-5., 3.]);
let mut out = Owned::zeros_f32(&[2, 2, 2]);
FusedGemmKernel::default()
.execute(&[a.view(), b.view(), bias.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![0., 1., 0., 7., 0., 0., 0., 11.]);
}
}