use onnx_runtime_ep_api::{Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::Node;
use super::add::broadcast_apply;
use super::{check_arity, to_dense_f32, write_dense_f32};
use crate::strided::numel;
#[derive(Clone, Copy)]
enum BinOp {
Sub,
Mul,
Div,
Pow,
Min,
Max,
}
impl BinOp {
fn name(self) -> &'static str {
match self {
BinOp::Sub => "Sub",
BinOp::Mul => "Mul",
BinOp::Div => "Div",
BinOp::Pow => "Pow",
BinOp::Min => "Min",
BinOp::Max => "Max",
}
}
fn apply(self, acc: f32, v: f32) -> f32 {
match self {
BinOp::Sub => acc - v,
BinOp::Mul => acc * v,
BinOp::Div => acc / v,
BinOp::Pow => acc.powf(v),
BinOp::Min => {
if acc.is_nan() || v.is_nan() {
f32::NAN
} else {
acc.min(v)
}
}
BinOp::Max => {
if acc.is_nan() || v.is_nan() {
f32::NAN
} else {
acc.max(v)
}
}
}
}
}
pub struct BinaryKernel {
op: BinOp,
}
macro_rules! binary_factory {
($factory:ident, $variant:expr) => {
pub struct $factory;
impl KernelFactory for $factory {
fn create(&self, _node: &Node, _shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
Ok(Box::new(BinaryKernel { op: $variant }))
}
}
};
}
binary_factory!(SubFactory, BinOp::Sub);
binary_factory!(MulFactory, BinOp::Mul);
binary_factory!(DivFactory, BinOp::Div);
binary_factory!(PowFactory, BinOp::Pow);
binary_factory!(MinFactory, BinOp::Min);
binary_factory!(MaxFactory, BinOp::Max);
impl Kernel for BinaryKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
let (min_in, max_in) = match self.op {
BinOp::Min | BinOp::Max => (1, usize::MAX),
_ => (2, 2),
};
check_arity(self.op.name(), inputs, outputs, min_in, max_in, 1)?;
let out_shape = outputs[0].shape.to_vec();
let n = numel(&out_shape);
let mut out = vec![0.0f32; n];
let first = to_dense_f32(&inputs[0])?;
broadcast_apply(&first, inputs[0].shape, &out_shape, |i, v| out[i] = v)?;
for input in &inputs[1..] {
let rhs = to_dense_f32(input)?;
let op = self.op;
broadcast_apply(&rhs, input.shape, &out_shape, |i, v| {
out[i] = op.apply(out[i], v)
})?;
}
write_dense_f32(&mut outputs[0], &out)
}
fn supports_strided_input(&self, _input_idx: usize) -> bool {
true
}
}
#[derive(Clone, Copy)]
enum UnOp {
Sqrt,
Erf,
Tanh,
}
impl UnOp {
fn name(self) -> &'static str {
match self {
UnOp::Sqrt => "Sqrt",
UnOp::Erf => "Erf",
UnOp::Tanh => "Tanh",
}
}
fn apply(self, x: f32) -> f32 {
match self {
UnOp::Sqrt => x.sqrt(),
UnOp::Erf => erf(x as f64) as f32,
UnOp::Tanh => x.tanh(),
}
}
}
pub struct UnaryKernel {
op: UnOp,
}
macro_rules! unary_factory {
($factory:ident, $variant:expr) => {
pub struct $factory;
impl KernelFactory for $factory {
fn create(&self, _node: &Node, _shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
Ok(Box::new(UnaryKernel { op: $variant }))
}
}
};
}
unary_factory!(SqrtFactory, UnOp::Sqrt);
unary_factory!(ErfFactory, UnOp::Erf);
unary_factory!(TanhFactory, UnOp::Tanh);
impl Kernel for UnaryKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
check_arity(self.op.name(), inputs, outputs, 1, 1, 1)?;
let x = to_dense_f32(&inputs[0])?;
let y: Vec<f32> = x.iter().map(|&v| self.op.apply(v)).collect();
write_dense_f32(&mut outputs[0], &y)
}
fn supports_strided_input(&self, _input_idx: usize) -> bool {
true
}
}
pub(crate) fn erf(x: f64) -> f64 {
if x.is_nan() {
return f64::NAN;
}
let sign = if x < 0.0 { -1.0 } else { 1.0 };
let x = x.abs();
const A1: f64 = 0.254_829_592;
const A2: f64 = -0.284_496_736;
const A3: f64 = 1.421_413_741;
const A4: f64 = -1.453_152_027;
const A5: f64 = 1.061_405_429;
const P: f64 = 0.327_591_1;
let t = 1.0 / (1.0 + P * x);
let poly = ((((A5 * t + A4) * t + A3) * t + A2) * t + A1) * t;
let y = 1.0 - poly * (-x * x).exp();
sign * y
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernels::testutil::Owned;
fn run_bin(f: BinOp, a: &Owned, b: &Owned, out: &mut Owned) {
BinaryKernel { op: f }
.execute(&[a.view(), b.view()], &mut [out.view_mut()])
.unwrap();
}
#[test]
fn sub_same_shape() {
let a = Owned::f32(&[2, 2], &[10., 20., 30., 40.]);
let b = Owned::f32(&[2, 2], &[1., 2., 3., 4.]);
let mut out = Owned::zeros_f32(&[2, 2]);
run_bin(BinOp::Sub, &a, &b, &mut out);
assert_eq!(out.to_f32(), vec![9., 18., 27., 36.]);
}
#[test]
fn mul_broadcasts_2d_with_2d() {
let a = Owned::f32(&[3, 1], &[1., 2., 3.]);
let b = Owned::f32(&[1, 4], &[10., 20., 30., 40.]);
let mut out = Owned::zeros_f32(&[3, 4]);
run_bin(BinOp::Mul, &a, &b, &mut out);
assert_eq!(
out.to_f32(),
vec![
10., 20., 30., 40., 20., 40., 60., 80., 30., 60., 90., 120., ]
);
}
#[test]
fn div_broadcasts_scalar() {
let a = Owned::f32(&[2, 2], &[2., 4., 6., 8.]);
let b = Owned::f32(&[], &[2.]); let mut out = Owned::zeros_f32(&[2, 2]);
run_bin(BinOp::Div, &a, &b, &mut out);
assert_eq!(out.to_f32(), vec![1., 2., 3., 4.]);
}
#[test]
fn div_by_zero_is_inf_and_nan() {
let a = Owned::f32(&[2], &[1., 0.]);
let b = Owned::f32(&[2], &[0., 0.]);
let mut out = Owned::zeros_f32(&[2]);
run_bin(BinOp::Div, &a, &b, &mut out);
let r = out.to_f32();
assert!(r[0].is_infinite() && r[0] > 0.0);
assert!(r[1].is_nan());
}
#[test]
fn pow_square() {
let a = Owned::f32(&[3], &[2., 3., 4.]);
let b = Owned::f32(&[], &[2.]);
let mut out = Owned::zeros_f32(&[3]);
run_bin(BinOp::Pow, &a, &b, &mut out);
assert_eq!(out.to_f32(), vec![4., 9., 16.]);
}
#[test]
fn min_variadic_three_inputs_with_broadcast() {
let a = Owned::f32(&[2, 2], &[5., 1., 8., 2.]);
let b = Owned::f32(&[2, 2], &[3., 3., 3., 3.]);
let c = Owned::f32(&[1], &[4.]); let mut out = Owned::zeros_f32(&[2, 2]);
BinaryKernel { op: BinOp::Min }
.execute(
&[a.view(), b.view(), c.view()],
&mut [out.view_mut()],
)
.unwrap();
assert_eq!(out.to_f32(), vec![3., 1., 3., 2.]);
}
#[test]
fn min_propagates_nan() {
let a = Owned::f32(&[3], &[f32::NAN, 2.0, 5.0]);
let b = Owned::f32(&[3], &[1.0, f32::NAN, 3.0]);
let mut out = Owned::zeros_f32(&[3]);
run_bin(BinOp::Min, &a, &b, &mut out);
let r = out.to_f32();
assert!(r[0].is_nan(), "NaN in lhs must propagate");
assert!(r[1].is_nan(), "NaN in rhs must propagate");
assert_eq!(r[2], 3.0);
}
#[test]
fn max_propagates_nan_and_reduces() {
let a = Owned::f32(&[3], &[f32::NAN, 2.0, 5.0]);
let b = Owned::f32(&[3], &[1.0, f32::NAN, 3.0]);
let mut out = Owned::zeros_f32(&[3]);
run_bin(BinOp::Max, &a, &b, &mut out);
let r = out.to_f32();
assert!(r[0].is_nan(), "NaN in lhs must propagate");
assert!(r[1].is_nan(), "NaN in rhs must propagate");
assert_eq!(r[2], 5.0);
}
#[test]
fn max_variadic_three_inputs() {
let a = Owned::f32(&[2, 2], &[5., 1., 8., 2.]);
let b = Owned::f32(&[2, 2], &[3., 3., 3., 3.]);
let c = Owned::f32(&[1], &[4.]);
let mut out = Owned::zeros_f32(&[2, 2]);
BinaryKernel { op: BinOp::Max }
.execute(&[a.view(), b.view(), c.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![5., 4., 8., 4.]);
}
#[test]
fn sqrt_unary() {
let a = Owned::f32(&[3], &[4., 9., 16.]);
let mut out = Owned::zeros_f32(&[3]);
UnaryKernel { op: UnOp::Sqrt }
.execute(&[a.view()], &mut [out.view_mut()])
.unwrap();
assert_eq!(out.to_f32(), vec![2., 3., 4.]);
}
#[test]
fn tanh_known_values() {
let a = Owned::f32(&[3], &[0., 1., -1.]);
let mut out = Owned::zeros_f32(&[3]);
UnaryKernel { op: UnOp::Tanh }
.execute(&[a.view()], &mut [out.view_mut()])
.unwrap();
let r = out.to_f32();
assert!((r[0] - 0.0).abs() < 1e-6);
assert!((r[1] - 0.761_594_2).abs() < 1e-6);
assert!((r[2] + 0.761_594_2).abs() < 1e-6);
}
#[test]
fn erf_known_values() {
let a = Owned::f32(&[4], &[0., 1., -1., 2.]);
let mut out = Owned::zeros_f32(&[4]);
UnaryKernel { op: UnOp::Erf }
.execute(&[a.view()], &mut [out.view_mut()])
.unwrap();
let r = out.to_f32();
assert!((r[0] - 0.0).abs() < 1e-6);
assert!((r[1] - 0.842_700_8).abs() < 1e-6);
assert!((r[2] + 0.842_700_8).abs() < 1e-6);
assert!((r[3] - 0.995_322_3).abs() < 1e-6);
}
#[test]
fn erf_odd_symmetry_and_limits() {
assert!((erf(0.0)).abs() < 1e-6);
assert!((erf(6.0) - 1.0).abs() < 1e-6);
assert!((erf(-6.0) + 1.0).abs() < 1e-6);
assert!(erf(f64::NAN).is_nan());
}
}