use onnx_runtime_ep_api::{EpError, OpKey, OpRegistry, Result, TensorMut, TensorView};
use onnx_runtime_ir::DataType;
use crate::strided::{elem_offset, next_index, numel};
pub mod add;
pub mod cast;
pub mod constant;
pub mod elementwise;
pub mod expand;
pub mod fused_attention;
pub mod fused_gemm;
pub mod fused_matmul_bias;
pub mod gather;
pub mod gelu;
pub mod gemm;
pub mod layernorm;
pub mod matmul;
pub mod reduce;
pub mod relu;
pub mod reshape;
pub mod shape;
pub mod slice;
pub mod softmax;
pub mod transpose;
pub mod unsqueeze;
pub const PHASE1_OPS: &[&str] = &[
"MatMul",
"Add",
"Relu",
"Reshape",
"Transpose",
"Gather",
"LayerNormalization",
"Sub",
"Mul",
"Div",
"Pow",
"Min",
"Max",
"Sqrt",
"Erf",
"Tanh",
"Cast",
"ReduceMean",
"Softmax",
"Shape",
"Unsqueeze",
"Expand",
"Slice",
"Constant",
"Gemm",
];
pub fn is_phase1_op(op_type: &str) -> bool {
PHASE1_OPS.contains(&op_type)
}
pub fn build_cpu_registry() -> OpRegistry {
let mut reg = OpRegistry::new();
reg.register(
OpKey::new("MatMul", "", 1),
Box::new(matmul::MatMulFactory),
);
reg.register(OpKey::new("Add", "", 1), Box::new(add::AddFactory));
reg.register(OpKey::new("Relu", "", 1), Box::new(relu::ReluFactory));
reg.register(
OpKey::new("Reshape", "", 1),
Box::new(reshape::ReshapeFactory),
);
reg.register(
OpKey::new("Transpose", "", 1),
Box::new(transpose::TransposeFactory),
);
reg.register(OpKey::new("Gather", "", 1), Box::new(gather::GatherFactory));
reg.register(
OpKey::new("LayerNormalization", "", 1),
Box::new(layernorm::LayerNormFactory),
);
reg.register(
OpKey::new("LayerNormalization", "com.microsoft", 1),
Box::new(layernorm::LayerNormFactory),
);
reg.register(
OpKey::new("FusedMatMulBias", "com.microsoft", 1),
Box::new(fused_matmul_bias::FusedMatMulBiasFactory),
);
reg.register(
OpKey::new("FusedGemm", "com.microsoft", 1),
Box::new(fused_gemm::FusedGemmFactory),
);
reg.register(
OpKey::new("FusedAttention", "com.microsoft", 1),
Box::new(fused_attention::FusedAttentionFactory),
);
reg.register(
OpKey::new("Gelu", "com.microsoft", 1),
Box::new(gelu::GeluFactory),
);
reg.register(OpKey::new("Sub", "", 1), Box::new(elementwise::SubFactory));
reg.register(OpKey::new("Mul", "", 1), Box::new(elementwise::MulFactory));
reg.register(OpKey::new("Div", "", 1), Box::new(elementwise::DivFactory));
reg.register(OpKey::new("Pow", "", 1), Box::new(elementwise::PowFactory));
reg.register(OpKey::new("Min", "", 1), Box::new(elementwise::MinFactory));
reg.register(OpKey::new("Max", "", 1), Box::new(elementwise::MaxFactory));
reg.register(OpKey::new("Sqrt", "", 1), Box::new(elementwise::SqrtFactory));
reg.register(OpKey::new("Erf", "", 1), Box::new(elementwise::ErfFactory));
reg.register(OpKey::new("Tanh", "", 1), Box::new(elementwise::TanhFactory));
reg.register(OpKey::new("Cast", "", 1), Box::new(cast::CastFactory));
reg.register(
OpKey::new("ReduceMean", "", 1),
Box::new(reduce::ReduceMeanFactory),
);
reg.register(
OpKey::new("Softmax", "", 1),
Box::new(softmax::SoftmaxLegacyFactory),
);
reg.register(
OpKey::new("Softmax", "", 13),
Box::new(softmax::SoftmaxFactory),
);
reg.register(OpKey::new("Shape", "", 1), Box::new(shape::ShapeFactory));
reg.register(
OpKey::new("Unsqueeze", "", 1),
Box::new(unsqueeze::UnsqueezeFactory),
);
reg.register(OpKey::new("Expand", "", 1), Box::new(expand::ExpandFactory));
reg.register(OpKey::new("Slice", "", 1), Box::new(slice::SliceFactory));
reg.register(
OpKey::new("Constant", "", 1),
Box::new(constant::ConstantFactory),
);
reg.register(OpKey::new("Gemm", "", 1), Box::new(gemm::GemmFactory));
reg
}
pub fn to_dense_f32(view: &TensorView) -> Result<Vec<f32>> {
view.validate()?;
require_dtype(view.dtype, DataType::Float32, "f32 kernel input")?;
let n = numel(view.shape);
let origin = view.data_ptr::<f32>();
let mut out = Vec::with_capacity(n);
if n == 0 {
return Ok(out);
}
let mut idx = vec![0usize; view.shape.len()];
loop {
let off = elem_offset(view.strides, &idx);
out.push(unsafe { *origin.offset(off) });
if !next_index(view.shape, &mut idx) {
break;
}
}
Ok(out)
}
pub fn to_dense_i64(view: &TensorView) -> Result<Vec<i64>> {
view.validate()?;
let n = numel(view.shape);
let mut out = Vec::with_capacity(n);
if n == 0 {
return Ok(out);
}
let mut idx = vec![0usize; view.shape.len()];
match view.dtype {
DataType::Int64 => {
let origin = view.data_ptr::<i64>();
loop {
let off = elem_offset(view.strides, &idx);
out.push(unsafe { *origin.offset(off) });
if !next_index(view.shape, &mut idx) {
break;
}
}
}
DataType::Int32 => {
let origin = view.data_ptr::<i32>();
loop {
let off = elem_offset(view.strides, &idx);
out.push(unsafe { *origin.offset(off) } as i64);
if !next_index(view.shape, &mut idx) {
break;
}
}
}
other => {
return Err(EpError::InvalidTensorView {
reason: format!("index tensor must be Int64 or Int32, got {other:?}"),
});
}
}
Ok(out)
}
pub fn write_dense_f32(out: &mut TensorMut, data: &[f32]) -> Result<()> {
out.validate()?;
require_dtype(out.dtype, DataType::Float32, "f32 kernel output")?;
let n = numel(out.shape);
if data.len() != n {
return Err(EpError::KernelFailed(format!(
"output element count {n} does not match produced {}",
data.len()
)));
}
if n == 0 {
return Ok(());
}
let origin = out.data_ptr_mut::<f32>();
let strides = out.strides;
let shape = out.shape;
let mut idx = vec![0usize; shape.len()];
let mut i = 0usize;
loop {
let off = elem_offset(strides, &idx);
unsafe {
*origin.offset(off) = data[i];
}
i += 1;
if !next_index(shape, &mut idx) {
break;
}
}
Ok(())
}
pub fn elem_size(dtype: DataType) -> Result<usize> {
let size = dtype.byte_size();
if size == 0 {
return Err(EpError::InvalidTensorView {
reason: format!("dtype {dtype:?} has no fixed-width byte layout"),
});
}
Ok(size)
}
pub fn to_dense_bytes(view: &TensorView) -> Result<Vec<u8>> {
view.validate()?;
let esize = elem_size(view.dtype)?;
let n = numel(view.shape);
let mut out = vec![0u8; n * esize];
if n == 0 {
return Ok(out);
}
let origin = view.data_ptr::<u8>();
let mut idx = vec![0usize; view.shape.len()];
let mut w = 0usize;
loop {
let elem_off = elem_offset(view.strides, &idx);
let byte_off = elem_off * esize as isize;
unsafe {
std::ptr::copy_nonoverlapping(origin.offset(byte_off), out.as_mut_ptr().add(w), esize);
}
w += esize;
if !next_index(view.shape, &mut idx) {
break;
}
}
Ok(out)
}
pub fn write_dense_bytes(out: &mut TensorMut, data: &[u8]) -> Result<()> {
out.validate()?;
let esize = elem_size(out.dtype)?;
let n = numel(out.shape);
if data.len() != n * esize {
return Err(EpError::KernelFailed(format!(
"output byte count {} does not match produced {}",
n * esize,
data.len()
)));
}
if n == 0 {
return Ok(());
}
let origin = out.data_ptr_mut::<u8>();
let strides = out.strides;
let shape = out.shape;
let mut idx = vec![0usize; shape.len()];
let mut r = 0usize;
loop {
let elem_off = elem_offset(strides, &idx);
let byte_off = elem_off * esize as isize;
unsafe {
std::ptr::copy_nonoverlapping(data.as_ptr().add(r), origin.offset(byte_off), esize);
}
r += esize;
if !next_index(shape, &mut idx) {
break;
}
}
Ok(())
}
fn require_dtype(got: DataType, want: DataType, ctx: &str) -> Result<()> {
if got != want {
return Err(EpError::InvalidTensorView {
reason: format!("{ctx} requires {want:?}, got {got:?}"),
});
}
Ok(())
}
fn check_arity(
op: &str,
inputs: &[TensorView],
outputs: &[TensorMut],
min_inputs: usize,
max_inputs: usize,
outputs_wanted: usize,
) -> Result<()> {
if inputs.len() < min_inputs || inputs.len() > max_inputs {
return Err(EpError::KernelFailed(format!(
"{op}: expected {min_inputs}..={max_inputs} inputs, got {}",
inputs.len()
)));
}
if outputs.len() < outputs_wanted {
return Err(EpError::KernelFailed(format!(
"{op}: expected at least {outputs_wanted} output(s), got {}",
outputs.len()
)));
}
Ok(())
}
#[cfg(test)]
pub(crate) mod testutil {
use onnx_runtime_ep_api::{DevicePtr, DevicePtrMut, TensorMut, TensorView};
use onnx_runtime_ir::{compute_contiguous_strides, DataType, DeviceId};
pub struct Owned {
pub bytes: Vec<u8>,
pub shape: Vec<usize>,
pub strides: Vec<i64>,
pub dtype: DataType,
}
impl Owned {
pub fn f32(shape: &[usize], data: &[f32]) -> Self {
let strides = compute_contiguous_strides(shape);
let mut bytes = Vec::with_capacity(data.len() * 4);
for v in data {
bytes.extend_from_slice(&v.to_le_bytes());
}
Self {
bytes,
shape: shape.to_vec(),
strides,
dtype: DataType::Float32,
}
}
pub fn i64(shape: &[usize], data: &[i64]) -> Self {
let strides = compute_contiguous_strides(shape);
let mut bytes = Vec::with_capacity(data.len() * 8);
for v in data {
bytes.extend_from_slice(&v.to_le_bytes());
}
Self {
bytes,
shape: shape.to_vec(),
strides,
dtype: DataType::Int64,
}
}
pub fn i32(shape: &[usize], data: &[i32]) -> Self {
let strides = compute_contiguous_strides(shape);
let mut bytes = Vec::with_capacity(data.len() * 4);
for v in data {
bytes.extend_from_slice(&v.to_le_bytes());
}
Self {
bytes,
shape: shape.to_vec(),
strides,
dtype: DataType::Int32,
}
}
pub fn bool_(shape: &[usize], data: &[bool]) -> Self {
let strides = compute_contiguous_strides(shape);
let bytes = data.iter().map(|&b| b as u8).collect();
Self {
bytes,
shape: shape.to_vec(),
strides,
dtype: DataType::Bool,
}
}
pub fn zeros_f32(shape: &[usize]) -> Self {
let n: usize = shape.iter().product();
Self::f32(shape, &vec![0.0; n])
}
pub fn zeros(dtype: DataType, shape: &[usize]) -> Self {
let n: usize = shape.iter().product();
let strides = compute_contiguous_strides(shape);
let esize = dtype.byte_size();
Self {
bytes: vec![0u8; n * esize],
shape: shape.to_vec(),
strides,
dtype,
}
}
pub fn with_view(mut self, shape: &[usize], strides: &[i64]) -> Self {
self.shape = shape.to_vec();
self.strides = strides.to_vec();
self
}
pub fn view(&self) -> TensorView<'_> {
TensorView::new(
DevicePtr(self.bytes.as_ptr() as *const std::ffi::c_void),
self.dtype,
&self.shape,
&self.strides,
DeviceId::cpu(),
)
}
pub fn view_mut(&mut self) -> TensorMut<'_> {
TensorMut::new(
DevicePtrMut(self.bytes.as_mut_ptr() as *mut std::ffi::c_void),
self.dtype,
&self.shape,
&self.strides,
DeviceId::cpu(),
)
}
pub fn to_f32(&self) -> Vec<f32> {
self.bytes
.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect()
}
pub fn to_i64(&self) -> Vec<i64> {
self.bytes
.chunks_exact(8)
.map(|c| i64::from_le_bytes(c.try_into().unwrap()))
.collect()
}
pub fn to_i32(&self) -> Vec<i32> {
self.bytes
.chunks_exact(4)
.map(|c| i32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect()
}
pub fn to_bool(&self) -> Vec<bool> {
self.bytes.iter().map(|&b| b != 0).collect()
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::strided::view_in_bounds;
use testutil::Owned;
#[test]
fn dense_roundtrip_contiguous() {
let a = Owned::f32(&[2, 3], &[1., 2., 3., 4., 5., 6.]);
let v = a.view();
assert_eq!(to_dense_f32(&v).unwrap(), vec![1., 2., 3., 4., 5., 6.]);
}
#[test]
fn dense_reads_transposed_view() {
let a = Owned::f32(&[2, 3], &[1., 2., 3., 4., 5., 6.]).with_view(&[3, 2], &[1, 3]);
let v = a.view();
assert_eq!(to_dense_f32(&v).unwrap(), vec![1., 4., 2., 5., 3., 6.]);
}
#[test]
fn registry_has_all_phase1_ops() {
let reg = build_cpu_registry();
assert_eq!(reg.len(), PHASE1_OPS.len() + 6);
for op in PHASE1_OPS {
assert!(reg.lookup(op, "", 21).is_some(), "missing factory for {op}");
}
assert!(reg.lookup("Softmax", "", 12).is_some());
assert!(reg.lookup("Softmax", "", 13).is_some());
assert!(reg.lookup("Conv", "", 21).is_none());
assert!(reg.lookup("LayerNormalization", "com.microsoft", 1).is_some());
assert!(reg.supports("LayerNormalization", "com.microsoft"));
assert!(reg.supports("MatMul", "ai.onnx"));
assert!(reg.supports("FusedMatMulBias", "com.microsoft"));
assert!(reg.supports("FusedGemm", "com.microsoft"));
assert!(reg.lookup("FusedGemm", "com.microsoft", 1).is_some());
assert!(reg.supports("Gelu", "com.microsoft"));
assert!(reg.lookup("Gelu", "com.microsoft", 1).is_some());
assert!(reg.lookup("Gelu", "", 21).is_none());
}
#[test]
fn dense_read_stays_in_bounds() {
let a = Owned::f32(&[3, 2], &[1., 4., 2., 5., 3., 6.]);
let v = a.view();
view_in_bounds(v.shape, v.strides, v.byte_offset, 4, a.bytes.len()).unwrap();
}
}