use std::alloc::{Layout, alloc, dealloc};
use std::ffi::c_void;
use onnx_runtime_ep_api::{
Cost, DeviceBuffer, EpConfig, EpError, ExecutionProvider, Fence, Kernel, KernelMatch,
OpRegistry, Result, deny,
};
use onnx_runtime_ir::{DataType, DeviceId, DeviceType, Node, Shape, TensorLayout};
use crate::WeightOffloadHostCache;
use crate::kernels::{build_cpu_registry, build_cpu_registry_with_weight_offload_cache};
use crate::optimizer::cpu_optimization_passes;
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 with_weight_offload_host_cache(host_cache: WeightOffloadHostCache) -> Self {
Self {
device: DeviceId::cpu(),
initialized: false,
registry: build_cpu_registry_with_weight_offload_cache(host_cache),
}
}
pub fn initialized_with_weight_offload_host_cache(
host_cache: WeightOffloadHostCache,
) -> Result<Self> {
let mut ep = Self::with_weight_offload_host_cache(host_cache);
ep.initialize(&Default::default())?;
Ok(ep)
}
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<()> {
crate::kernels::matmul_nbits::bound_process_to_decode_budget();
self.initialized = true;
Ok(())
}
fn shutdown(&mut self) -> Result<()> {
self.initialized = false;
Ok(())
}
fn supports_op(
&self,
op: &Node,
opset: u64,
shapes: &[Shape],
input_dtypes: &[DataType],
_layouts: &[TensorLayout],
) -> KernelMatch {
let domain = if op.domain.is_empty() {
"ai.onnx"
} else {
&op.domain
};
if !self.registry.supports(&op.op_type, &op.domain, opset) {
if let Some(since) = self
.registry
.earliest_since_version(&op.op_type, &op.domain)
{
deny!(
"no handler for {}::{} at opset {} — this EP registers {} since opset {} (or: add a claim+handler)",
domain,
op.op_type,
opset,
op.op_type,
since
);
}
deny!(
"no handler for {}::{} at opset {} — add a claim+handler",
domain,
op.op_type,
opset
);
}
if op.op_type == "CompressedSparseAttention"
&& op.domain == "pkg.nxrt"
&& let Some(reason) = crate::kernels::compressed_sparse_attention::unsupported_reason(
op,
shapes,
input_dtypes,
)
{
return KernelMatch::unsupported(reason);
}
if op.op_type == "BlockQuantizedMoE"
&& op.domain == "pkg.nxrt"
&& let Some(reason) =
crate::kernels::block_quantized_moe::unsupported_reason(op, shapes, input_dtypes)
{
return KernelMatch::unsupported(reason);
}
if op.op_type == "IndexShare"
&& op.domain == "pkg.nxrt"
&& let Some(reason) =
crate::kernels::index_share::unsupported_reason(op, shapes, input_dtypes)
{
return KernelMatch::unsupported(reason);
}
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 {
domain: if op.domain.is_empty() {
"ai.onnx".to_string()
} else {
op.domain.clone()
},
op_type: op.op_type.clone(),
opset,
})?;
factory.create(op, shapes)
}
fn custom_passes(&self) -> Vec<Box<dyn onnx_runtime_optimizer::OptimizationPass>> {
cpu_optimization_passes()
}
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()
);
if buffer.is_borrowed() {
return Ok(());
}
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::*;
use onnx_runtime_ir::{Attribute, Graph, NodeId, static_shape};
fn stateful_csa_node(ratio: i64, input_count: usize, output_count: usize) -> Node {
let mut graph = Graph::new();
let inputs = (0..input_count)
.map(|index| {
Some(graph.create_named_value(
format!("input_{index}"),
DataType::Float32,
static_shape([]),
))
})
.collect();
let outputs = (0..output_count)
.map(|index| {
graph.create_named_value(
format!("output_{index}"),
DataType::Float32,
static_shape([]),
)
})
.collect();
let mut node = Node::new(NodeId(0), "CompressedSparseAttention", inputs, outputs);
node.domain = "pkg.nxrt".into();
node.attributes
.insert("num_heads".into(), Attribute::Int(1));
node.attributes
.insert("head_dim".into(), Attribute::Int(512));
node.attributes
.insert("qk_rope_head_dim".into(), Attribute::Int(64));
node.attributes
.insert("compression_ratio".into(), Attribute::Int(ratio));
if ratio == 4 {
node.attributes
.insert("index_num_heads".into(), Attribute::Int(1));
node.attributes
.insert("index_head_dim".into(), Attribute::Int(128));
node.attributes
.insert("index_topk".into(), Attribute::Int(1));
}
node
}
#[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 deallocate_borrowed_buffer_is_a_noop_free() {
let ep = CpuExecutionProvider::new();
let mut backing = vec![42u8; 128];
let ptr = backing.as_mut_ptr() as *mut c_void;
let buf = unsafe { DeviceBuffer::from_borrowed_parts(ptr, ep.device_id(), 128, 1) };
assert!(buf.is_borrowed());
ep.deallocate(buf).unwrap();
assert!(backing.iter().all(|&b| b == 42));
backing[0] = 1; assert_eq!(backing[0], 1);
}
#[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 mut node = Node::new(onnx_runtime_ir::NodeId(i as u32), *op, vec![], vec![]);
if *op == "BitShift" {
node.attributes
.insert("direction".into(), Attribute::String(b"RIGHT".to_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, 17, &[], &[], &[]).is_supported());
let conv = Node::new(onnx_runtime_ir::NodeId(1), "Conv", vec![], vec![]);
let rejected = ep.supports_op(&conv, 17, &[], &[], &[]);
#[cfg(feature = "mlas")]
assert!(rejected.is_supported());
#[cfg(not(feature = "mlas"))]
{
assert!(!rejected.is_supported());
let reason = rejected.reason().expect("unsupported reason");
assert!(reason.contains("Conv"), "{reason}");
assert!(
reason.contains("no handler for ai.onnx::Conv at opset 17"),
"{reason}"
);
assert!(reason.contains("add a claim+handler"), "{reason}");
}
}
#[test]
fn supports_op_is_opset_aware_for_standard_gelu() {
let ep = CpuExecutionProvider::new();
let gelu = Node::new(onnx_runtime_ir::NodeId(0), "Gelu", vec![], vec![]);
let rejected = ep.supports_op(&gelu, 19, &[], &[], &[]);
let reason = rejected.reason().expect("opset 19 must be declined");
assert!(
reason.contains("no handler for ai.onnx::Gelu at opset 19"),
"{reason}"
);
assert!(reason.contains("registers Gelu since opset 20"), "{reason}");
assert!(ep.supports_op(&gelu, 20, &[], &[], &[]).is_supported());
}
#[test]
fn supports_op_rejects_malformed_csa_ratio_specific_arity() {
let ep = CpuExecutionProvider::new();
let mut ratio4_missing_index = stateful_csa_node(4, 19, 5);
ratio4_missing_index.inputs[17] = None;
for (node, expected) in [
(
ratio4_missing_index,
"ratio-4 requires all eight positional index inputs (11..=18)",
),
(
stateful_csa_node(4, 19, 4),
"ratio-4 requires 5 or 6 outputs, got 4",
),
(
stateful_csa_node(128, 12, 3),
"ratio-4-only inputs (11..=18)",
),
(
stateful_csa_node(128, 11, 4),
"ratio-128 supports exactly 3 outputs, got 4",
),
] {
let rejected = ep.supports_op(&node, 1, &[], &[], &[]);
assert!(!rejected.is_supported());
let reason = rejected.reason().expect("CSA claim must be denied");
assert!(reason.contains(expected), "{reason}");
}
}
#[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, 1, &[], &[], &[]).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, 1, &[], &[], &[]).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, 1, &[], &[], &[]).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, 1, &[], &[], &[]).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, 1, &[], &[], &[]).is_supported());
}
}