use std::alloc::{alloc, dealloc, Layout};
use std::ffi::c_void;
use onnx_runtime_ep_api::{
Cost, DeviceBuffer, EpConfig, EpError, ExecutionProvider, Fence, Kernel, KernelMatch,
OpRegistry, Result,
};
use onnx_runtime_ir::{DeviceId, DeviceType, Node, Shape, TensorLayout};
use crate::kernels::build_cpu_registry;
pub struct CpuExecutionProvider {
device: DeviceId,
initialized: bool,
registry: OpRegistry,
}
impl std::fmt::Debug for CpuExecutionProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CpuExecutionProvider")
.field("device", &self.device)
.field("initialized", &self.initialized)
.field("registered_ops", &self.registry.len())
.finish()
}
}
impl Default for CpuExecutionProvider {
fn default() -> Self {
Self::new()
}
}
impl CpuExecutionProvider {
pub fn new() -> Self {
Self {
device: DeviceId::cpu(),
initialized: false,
registry: build_cpu_registry(),
}
}
pub fn registry(&self) -> &OpRegistry {
&self.registry
}
}
impl ExecutionProvider for CpuExecutionProvider {
fn name(&self) -> &str {
"cpu_ep"
}
fn device_type(&self) -> DeviceType {
DeviceType::Cpu
}
fn device_id(&self) -> DeviceId {
self.device
}
fn initialize(&mut self, _config: &EpConfig) -> Result<()> {
self.initialized = true;
Ok(())
}
fn shutdown(&mut self) -> Result<()> {
self.initialized = false;
Ok(())
}
fn supports_op(&self, op: &Node, shapes: &[Shape], _layouts: &[TensorLayout]) -> KernelMatch {
if !self.registry.supports(&op.op_type, &op.domain) {
return KernelMatch::Unsupported;
}
let output_layouts = vec![TensorLayout::contiguous(); op.outputs.len()];
let elems: u64 = shapes
.iter()
.map(|s| {
s.iter()
.map(|d| d.as_static().unwrap_or(1) as u64)
.product::<u64>()
})
.sum();
let cost = Cost::new(elems as f64, elems as f64, 0.0)
.with_launch_us(0.1)
.with_bytes_moved(elems.saturating_mul(4));
KernelMatch::Supported {
cost,
required_input_layouts: None,
output_layouts,
}
}
fn get_kernel(&self, op: &Node, shapes: &[Vec<usize>], opset: u64) -> Result<Box<dyn Kernel>> {
let factory = self
.registry
.lookup(&op.op_type, &op.domain, opset)
.ok_or_else(|| EpError::NoEpForOp {
op_type: op.op_type.clone(),
})?;
factory.create(op, shapes)
}
fn allocate(&self, size: usize, alignment: usize) -> Result<DeviceBuffer> {
if alignment == 0 || !alignment.is_power_of_two() {
return Err(EpError::AlignmentError);
}
let alloc_size = size.max(1);
let layout =
Layout::from_size_align(alloc_size, alignment).map_err(|_| EpError::AlignmentError)?;
let ptr = unsafe { alloc(layout) } as *mut c_void;
if ptr.is_null() {
return Err(EpError::OutOfMemory {
requested: size,
available: 0,
});
}
Ok(unsafe { DeviceBuffer::from_raw_parts(ptr, self.device, size, alignment) })
}
fn deallocate(&self, buffer: DeviceBuffer) -> Result<()> {
assert_eq!(
buffer.device(),
self.device,
"cpu_ep: refusing to deallocate a buffer from device {:?}",
buffer.device()
);
let size = buffer.len();
let align = buffer.alignment();
let ptr = buffer.into_raw() as *mut u8;
let layout = Layout::from_size_align(size.max(1), align)
.expect("cpu_ep: layout was valid at allocation time");
unsafe { dealloc(ptr, layout) };
Ok(())
}
fn copy(&self, src: &DeviceBuffer, dst: &mut DeviceBuffer, size: usize) -> Result<()> {
assert_eq!(src.device(), self.device, "cpu_ep::copy: foreign src buffer");
assert_eq!(dst.device(), self.device, "cpu_ep::copy: foreign dst buffer");
if size > src.len() || size > dst.len() {
return Err(EpError::KernelFailed(format!(
"cpu_ep::copy: size {size} exceeds src {} or dst {}",
src.len(),
dst.len()
)));
}
if size == 0 {
return Ok(());
}
let src_ptr = src.as_ptr() as *const u8;
let dst_ptr = dst.as_mut_ptr() as *mut u8;
unsafe {
std::ptr::copy_nonoverlapping(src_ptr, dst_ptr, size);
}
Ok(())
}
fn copy_async(&self, src: &DeviceBuffer, dst: &mut DeviceBuffer, size: usize) -> Result<Fence> {
self.copy(src, dst, size)?;
Ok(Fence::default())
}
fn sync(&self) -> Result<()> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn identity_and_lifecycle() {
let mut ep = CpuExecutionProvider::new();
assert_eq!(ep.name(), "cpu_ep");
assert_eq!(ep.device_type(), DeviceType::Cpu);
assert_eq!(ep.device_id(), DeviceId::cpu());
ep.initialize(&EpConfig::default()).unwrap();
assert!(ep.initialized);
ep.shutdown().unwrap();
assert!(!ep.initialized);
}
#[test]
fn allocate_deallocate_single_free_and_aligned() {
let ep = CpuExecutionProvider::new();
let buf = ep.allocate(256, 64).unwrap();
assert_eq!(buf.len(), 256);
assert_eq!(buf.alignment(), 64);
assert_eq!(buf.device(), DeviceId::cpu());
assert_eq!(buf.as_ptr() as usize % 64, 0);
ep.deallocate(buf).unwrap();
}
#[test]
fn allocate_zero_size_is_nonnull() {
let ep = CpuExecutionProvider::new();
let buf = ep.allocate(0, 16).unwrap();
assert_eq!(buf.len(), 0);
assert!(!buf.as_ptr().is_null());
ep.deallocate(buf).unwrap();
}
#[test]
fn allocate_rejects_bad_alignment() {
let ep = CpuExecutionProvider::new();
assert!(matches!(ep.allocate(16, 0), Err(EpError::AlignmentError)));
assert!(matches!(
ep.allocate(16, 24), Err(EpError::AlignmentError)
));
}
#[test]
fn copy_moves_bytes_and_checks_size() {
let ep = CpuExecutionProvider::new();
let mut src = ep.allocate(16, 16).unwrap();
let mut dst = ep.allocate(16, 16).unwrap();
unsafe {
let p = src.as_mut_ptr() as *mut u8;
for i in 0..16u8 {
*p.add(i as usize) = i;
}
}
ep.copy(&src, &mut dst, 16).unwrap();
unsafe {
let p = dst.as_ptr() as *const u8;
for i in 0..16u8 {
assert_eq!(*p.add(i as usize), i);
}
}
assert!(ep.copy(&src, &mut dst, 32).is_err());
ep.deallocate(src).unwrap();
ep.deallocate(dst).unwrap();
}
#[test]
#[should_panic(expected = "device")]
fn deallocate_rejects_cross_device_buffer() {
let ep = CpuExecutionProvider::new();
let boxed = vec![0u8; 8].into_boxed_slice();
let ptr = Box::into_raw(boxed) as *mut c_void;
let foreign = unsafe { DeviceBuffer::from_raw_parts(ptr, DeviceId::cuda(0), 8, 8) };
let _ = ep.deallocate(foreign); }
#[test]
fn get_kernel_dispatches_phase1_ops() {
let ep = CpuExecutionProvider::new();
for (i, op) in crate::kernels::PHASE1_OPS.iter().enumerate() {
let node = Node::new(onnx_runtime_ir::NodeId(i as u32), *op, vec![], vec![]);
assert!(ep.get_kernel(&node, &[], 17).is_ok(), "no kernel for {op}");
}
let bad = Node::new(onnx_runtime_ir::NodeId(99), "Conv", vec![], vec![]);
assert!(ep.get_kernel(&bad, &[], 17).is_err());
}
#[test]
fn supports_op_reports_phase1_only() {
let ep = CpuExecutionProvider::new();
let mm = Node::new(onnx_runtime_ir::NodeId(0), "MatMul", vec![], vec![]);
assert!(ep.supports_op(&mm, &[], &[]).is_supported());
let conv = Node::new(onnx_runtime_ir::NodeId(1), "Conv", vec![], vec![]);
assert!(!ep.supports_op(&conv, &[], &[]).is_supported());
}
#[test]
fn supports_fused_contrib_domain_layernorm() {
let ep = CpuExecutionProvider::new();
let mut fused = Node::new(onnx_runtime_ir::NodeId(0), "LayerNormalization", vec![], vec![]);
fused.domain = "com.microsoft".to_string();
assert!(ep.supports_op(&fused, &[], &[]).is_supported());
assert!(ep.get_kernel(&fused, &[], 1).is_ok());
let mut fmb = Node::new(onnx_runtime_ir::NodeId(1), "FusedMatMulBias", vec![], vec![]);
fmb.domain = "com.microsoft".to_string();
assert!(ep.supports_op(&fmb, &[], &[]).is_supported());
assert!(ep.get_kernel(&fmb, &[], 1).is_ok());
let mut fg = Node::new(onnx_runtime_ir::NodeId(2), "FusedGemm", vec![], vec![]);
fg.domain = "com.microsoft".to_string();
assert!(ep.supports_op(&fg, &[], &[]).is_supported());
assert!(ep.get_kernel(&fg, &[], 1).is_ok());
let mut fa = Node::new(onnx_runtime_ir::NodeId(4), "FusedAttention", vec![], vec![]);
fa.domain = "com.microsoft".to_string();
assert!(ep.supports_op(&fa, &[], &[]).is_supported());
fa.attributes.insert(
"scale".to_string(),
onnx_runtime_ir::Attribute::Float(0.5),
);
assert!(ep.get_kernel(&fa, &[], 1).is_ok());
let mut unknown = Node::new(onnx_runtime_ir::NodeId(3), "NotARealFusedOp", vec![], vec![]);
unknown.domain = "com.microsoft".to_string();
assert!(!ep.supports_op(&unknown, &[], &[]).is_supported());
}
}