use onnx_runtime_ep_api::{Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::Node;
use super::check_arity;
use super::elementwise::erf;
use crate::dtype::{to_dense_f32_widen, write_dense_f32_narrow};
const FRAC_1_SQRT_2: f64 = std::f64::consts::FRAC_1_SQRT_2;
pub(crate) fn exact_gelu(x: f32) -> f32 {
let xf = x as f64;
(0.5 * xf * (1.0 + erf(xf * FRAC_1_SQRT_2))) as f32
}
const SQRT_2_OVER_PI: f64 = 0.797_884_560_802_865_4;
pub(crate) fn tanh_gelu(x: f32) -> f32 {
let xf = x as f64;
let inner = SQRT_2_OVER_PI * (xf + 0.044_715 * xf * xf * xf);
(0.5 * xf * (1.0 + inner.tanh())) as f32
}
pub struct GeluKernel;
pub struct GeluFactory;
impl KernelFactory for GeluFactory {
fn create(&self, _node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
Ok(Box::new(GeluKernel))
}
}
impl Kernel for GeluKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
check_arity("Gelu", inputs, outputs, 1, 1, 1)?;
let x = to_dense_f32_widen("Gelu", &inputs[0])?;
let y: Vec<f32> = x.iter().map(|&v| exact_gelu(v)).collect();
write_dense_f32_narrow("Gelu", &mut outputs[0], &y)
}
fn supports_strided_input(&self, _input_idx: usize) -> bool {
true
}
}
pub struct StdGeluKernel {
tanh: bool,
}
pub struct StdGeluFactory;
impl KernelFactory for StdGeluFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let tanh = node
.attr("approximate")
.and_then(|a| a.as_str())
.map(|s| s == "tanh")
.unwrap_or(false);
Ok(Box::new(StdGeluKernel { tanh }))
}
}
impl Kernel for StdGeluKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
check_arity("Gelu", inputs, outputs, 1, 1, 1)?;
let x = to_dense_f32_widen("Gelu", &inputs[0])?;
let y: Vec<f32> = if self.tanh {
x.iter().map(|&v| tanh_gelu(v)).collect()
} else {
x.iter().map(|&v| exact_gelu(v)).collect()
};
write_dense_f32_narrow("Gelu", &mut outputs[0], &y)
}
fn supports_strided_input(&self, _input_idx: usize) -> bool {
true
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernels::testutil::Owned;
fn reference(x: f32) -> f32 {
let xf = x as f64;
(0.5 * xf * (1.0 + erf(xf / std::f64::consts::SQRT_2))) as f32
}
#[test]
fn gelu_known_values() {
let xs = [-3.0f32, -1.0, -0.5, 0.0, 0.5, 1.0, 2.0, 3.0];
let a = Owned::f32(&[xs.len()], &xs);
let mut out = Owned::zeros_f32(&[xs.len()]);
GeluKernel
.execute(&[a.view()], &mut [out.view_mut()])
.unwrap();
let got = out.to_f32();
for (&x, &g) in xs.iter().zip(got.iter()) {
assert!((g - reference(x)).abs() <= 1e-6, "gelu({x}) = {g}");
}
assert_eq!(got[xs.iter().position(|&v| v == 0.0).unwrap()], 0.0);
}
#[test]
fn gelu_reads_strided_view() {
let a = Owned::f32(&[2, 2], &[-1.0, 1.0, 2.0, -2.0]).with_view(&[2, 2], &[1, 2]);
let mut out = Owned::zeros_f32(&[2, 2]);
GeluKernel
.execute(&[a.view()], &mut [out.view_mut()])
.unwrap();
let expect: Vec<f32> = [-1.0f32, 2.0, 1.0, -2.0]
.iter()
.map(|&v| reference(v))
.collect();
assert_eq!(out.to_f32(), expect);
}
fn reference_tanh(x: f32) -> f32 {
let xf = x as f64;
let inner = (2.0f64 / std::f64::consts::PI).sqrt() * (xf + 0.044_715 * xf * xf * xf);
(0.5 * xf * (1.0 + inner.tanh())) as f32
}
#[test]
fn std_gelu_approximate_none_matches_exact() {
let xs = [-3.0f32, -1.0, -0.5, 0.0, 0.5, 1.0, 2.0, 3.0];
let a = Owned::f32(&[xs.len()], &xs);
let mut out = Owned::zeros_f32(&[xs.len()]);
StdGeluKernel { tanh: false }
.execute(&[a.view()], &mut [out.view_mut()])
.unwrap();
for (&x, &g) in xs.iter().zip(out.to_f32().iter()) {
assert!((g - reference(x)).abs() <= 1e-6, "gelu({x}) = {g}");
}
}
#[test]
fn std_gelu_approximate_tanh_matches_reference() {
let xs = [-3.0f32, -1.0, -0.5, 0.0, 0.5, 1.0, 2.0, 3.0];
let a = Owned::f32(&[xs.len()], &xs);
let mut out = Owned::zeros_f32(&[xs.len()]);
StdGeluKernel { tanh: true }
.execute(&[a.view()], &mut [out.view_mut()])
.unwrap();
let got = out.to_f32();
for (&x, &g) in xs.iter().zip(got.iter()) {
assert!(
(g - reference_tanh(x)).abs() <= 1e-6,
"gelu_tanh({x}) = {g}"
);
}
assert_eq!(got[3], 0.0);
assert!(
(got[5] - 0.841_192).abs() < 1e-4,
"gelu_tanh(1) = {}",
got[5]
);
}
#[test]
fn std_gelu_factory_reads_approximate_attr() {
use onnx_runtime_ir::{Attribute, Node, NodeId};
let mut node = Node::new(NodeId(0), "Gelu", vec![], vec![]);
node.attributes.insert(
"approximate".to_string(),
Attribute::String(b"tanh".to_vec()),
);
let k = StdGeluFactory.create(&node, &[]).unwrap();
let a = Owned::f32(&[1], &[1.0]);
let mut out = Owned::zeros_f32(&[1]);
k.execute(&[a.view()], &mut [out.view_mut()]).unwrap();
assert!((out.to_f32()[0] - reference_tanh(1.0)).abs() < 1e-6);
}
}