use std::collections::HashMap;
use std::sync::Arc;
use onnx_runtime_ep_api::{
DeviceBuffer, DevicePtr, DevicePtrMut, ExecutionProvider, KernelMatch, TensorMut, TensorView,
};
use onnx_runtime_ep_cpu::strided::view_in_bounds;
use onnx_runtime_ep_cpu::CpuExecutionProvider;
use onnx_runtime_ir::{
as_static_shape, compute_contiguous_strides, DataType, Dim, Graph, Node, NodeId, Shape,
SymbolId, TensorLayout, ValueId,
};
use onnx_runtime_loader::WeightStore;
use crate::error::{Result, SessionError};
use crate::tensor::{host_bytes, write_host, Tensor};
#[derive(Debug)]
pub(crate) struct NodePlan {
pub node_id: NodeId,
pub inputs: Vec<ValueId>,
pub outputs: Vec<ValueId>,
pub input_dtypes: Vec<DataType>,
pub output_dtypes: Vec<DataType>,
}
#[derive(Clone, PartialEq, Eq, Hash, Debug)]
struct KernelKey {
node: u32,
shapes: Vec<Vec<usize>>,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct CacheStats {
pub entries: usize,
pub hits: u64,
pub misses: u64,
}
#[derive(Default)]
pub(crate) struct KernelCache {
entries: HashMap<KernelKey, Box<dyn onnx_runtime_ep_api::Kernel>>,
hits: u64,
misses: u64,
}
impl KernelCache {
fn stats(&self) -> CacheStats {
CacheStats {
entries: self.entries.len(),
hits: self.hits,
misses: self.misses,
}
}
fn get_or_create(
&mut self,
node_id: NodeId,
node: &Node,
input_shapes: &[Vec<usize>],
opset: u64,
ep: &CpuExecutionProvider,
) -> Result<&dyn onnx_runtime_ep_api::Kernel> {
let key = KernelKey {
node: node_id.0,
shapes: input_shapes.to_vec(),
};
if self.entries.contains_key(&key) {
self.hits += 1;
} else {
let shape_dims: Vec<Shape> = input_shapes
.iter()
.map(|s| s.iter().map(|&d| Dim::Static(d)).collect())
.collect();
let layouts = vec![TensorLayout::contiguous(); input_shapes.len()];
if !matches!(
ep.supports_op(node, &shape_dims, &layouts),
KernelMatch::Supported { .. }
) {
return Err(SessionError::UnsupportedOp {
op_type: node.op_type.clone(),
});
}
let kernel = ep.get_kernel(node, input_shapes, opset)?;
self.entries.insert(key.clone(), kernel);
self.misses += 1;
}
Ok(self.entries.get(&key).expect("just inserted").as_ref())
}
}
pub(crate) struct Executor {
graph: Graph,
_weights: Arc<WeightStore>,
ep: Arc<CpuExecutionProvider>,
buffers: HashMap<ValueId, DeviceBuffer>,
buffer_shapes: HashMap<ValueId, Vec<usize>>,
value_shapes: HashMap<ValueId, Shape>,
value_dtypes: HashMap<ValueId, DataType>,
plan: Vec<NodePlan>,
input_index: HashMap<String, ValueId>,
required_inputs: Vec<ValueId>,
has_symbols: bool,
cache: KernelCache,
}
fn view_bounds(
shape: &[usize],
strides: &[i64],
byte_offset: usize,
dtype: DataType,
buffer_len: usize,
) -> Result<()> {
let esize = dtype.byte_size();
if esize == 0 {
let numel: usize = shape.iter().product();
let need = byte_offset + dtype.storage_bytes(numel);
if need > buffer_len {
return Err(SessionError::from(
onnx_runtime_ep_api::EpError::InvalidTensorView {
reason: format!(
"sub-byte view needs {need} bytes but backing allocation is {buffer_len}"
),
},
));
}
return Ok(());
}
view_in_bounds(shape, strides, byte_offset, esize, buffer_len)?;
Ok(())
}
fn checked_numel(dims: &[usize], value: impl FnOnce() -> String) -> Result<usize> {
let mut acc = 1usize;
for &d in dims {
acc = match acc.checked_mul(d) {
Some(n) => n,
None => {
return Err(SessionError::ShapeOverflow {
value: value(),
dims: dims.to_vec(),
})
}
};
}
Ok(acc)
}
fn checked_storage_bytes(
dtype: DataType,
numel: usize,
value: impl FnOnce() -> String,
dims: &[usize],
) -> Result<usize> {
dtype
.checked_storage_bytes(numel)
.ok_or_else(|| SessionError::ShapeOverflow {
value: value(),
dims: dims.to_vec(),
})
}
fn effective_opset(graph: &Graph, node: &Node) -> u64 {
let domain = node.domain.as_str();
graph
.opset_imports
.get(domain)
.or_else(|| {
if domain.is_empty() {
graph.opset_imports.get("ai.onnx")
} else if domain == "ai.onnx" {
graph.opset_imports.get("")
} else {
None
}
})
.copied()
.unwrap_or(u64::MAX)
}
fn substitute(shape: &Shape, bindings: &HashMap<SymbolId, usize>) -> Option<Vec<usize>> {
shape
.iter()
.map(|d| match d {
Dim::Static(n) => Some(*n),
Dim::Symbolic(s) => bindings.get(s).copied(),
})
.collect()
}
fn buffer_as_i64(buffer: &DeviceBuffer, dtype: DataType) -> Option<Vec<i64>> {
let bytes = crate::tensor::host_bytes(buffer);
match dtype {
DataType::Int64 => Some(
bytes
.chunks_exact(8)
.map(|c| i64::from_le_bytes(c.try_into().unwrap()))
.collect(),
),
DataType::Int32 => Some(
bytes
.chunks_exact(4)
.map(|c| i32::from_le_bytes(c.try_into().unwrap()) as i64)
.collect(),
),
_ => None,
}
}
fn dynamic_output_shapes(
node: &Node,
input_shapes: &[Vec<usize>],
input_values: &[Option<Vec<i64>>],
) -> Option<Vec<Vec<usize>>> {
match node.op_type.as_str() {
"Slice" => {
let data_shape = input_shapes.first()?;
let starts = input_values.get(1)?.as_ref()?;
let ends = input_values.get(2)?.as_ref()?;
let (axes, steps) = onnx_runtime_ep_cpu::slice_axes_steps(
starts.len(),
input_values.get(3).and_then(|v| v.as_deref()),
input_values.get(4).and_then(|v| v.as_deref()),
);
let plan =
onnx_runtime_ep_cpu::slice_plan(data_shape, starts, ends, &axes, &steps).ok()?;
let count: Vec<usize> = plan.iter().map(|p| p.count).collect();
Some(vec![count])
}
_ => None,
}
}
impl Executor {
pub(crate) fn build(
graph: Graph,
weights: Arc<WeightStore>,
ep: Arc<CpuExecutionProvider>,
) -> Result<Self> {
let order = graph.topological_order()?;
let mut value_shapes: HashMap<ValueId, Shape> = HashMap::new();
let mut value_dtypes: HashMap<ValueId, DataType> = HashMap::new();
let mut buffers: HashMap<ValueId, DeviceBuffer> = HashMap::new();
let mut buffer_shapes: HashMap<ValueId, Vec<usize>> = HashMap::new();
for (&vid, weight) in &graph.initializers {
let dtype = weight.dtype();
let dims = weight.dims().to_vec();
let bytes = weights.bytes(weight).ok_or_else(|| {
SessionError::Internal(format!("weight bytes unavailable for value#{}", vid.0))
})?;
let mut buf = ep.allocate(bytes.len().max(1), TensorLayout::contiguous().alignment)?;
write_host(&mut buf, bytes)?;
value_dtypes.insert(vid, dtype);
value_shapes.insert(vid, dims.iter().map(|&d| Dim::Static(d)).collect());
buffer_shapes.insert(vid, dims);
buffers.insert(vid, buf);
}
for &vid in &graph.inputs {
value_shapes
.entry(vid)
.or_insert_with(|| graph.value(vid).shape.clone());
value_dtypes.entry(vid).or_insert(graph.value(vid).dtype);
}
for &nid in &order {
for &out in &graph.node(nid).outputs {
value_shapes
.entry(out)
.or_insert_with(|| graph.value(out).shape.clone());
value_dtypes.entry(out).or_insert(graph.value(out).dtype);
}
}
let has_symbols = value_shapes.values().any(|s| as_static_shape(s).is_none());
let mut plan = Vec::with_capacity(order.len());
for &nid in &order {
let node = graph.node(nid);
if onnx_runtime_loader::is_ep_context_op(&node.op_type, &node.domain) {
continue;
}
let inputs: Vec<ValueId> = node.input_values().collect();
let outputs: Vec<ValueId> = node.outputs.clone();
let input_dtypes: Vec<DataType> = inputs.iter().map(|v| value_dtypes[v]).collect();
let output_dtypes: Vec<DataType> = outputs.iter().map(|v| value_dtypes[v]).collect();
plan.push(NodePlan {
node_id: nid,
inputs,
outputs,
input_dtypes,
output_dtypes,
});
}
let mut input_index = HashMap::new();
let mut required_inputs = Vec::new();
for &vid in &graph.inputs {
if graph.initializers.contains_key(&vid) {
continue; }
required_inputs.push(vid);
if let Some(name) = &graph.value(vid).name {
input_index.insert(name.clone(), vid);
}
}
let mut exec = Self {
graph,
_weights: weights,
ep,
buffers,
buffer_shapes,
value_shapes,
value_dtypes,
plan,
input_index,
required_inputs,
has_symbols,
cache: KernelCache::default(),
};
if !exec.has_symbols {
let empty = HashMap::new();
let resolved = exec.resolve_all(&empty)?;
exec.size_buffers(&resolved)?;
exec.compile_all(&resolved)?;
}
Ok(exec)
}
fn ensure_buffer(&mut self, vid: ValueId, dtype: DataType, dims: &[usize]) -> Result<()> {
if self.buffer_shapes.get(&vid).map(|s| s.as_slice()) == Some(dims) {
return Ok(()); }
if let Some(old) = self.buffers.remove(&vid) {
self.ep.deallocate(old)?;
}
let numel = checked_numel(dims, || format!("value#{}", vid.0))?;
let size = checked_storage_bytes(dtype, numel, || format!("value#{}", vid.0), dims)?;
let buf = self
.ep
.allocate(size.max(1), TensorLayout::contiguous().alignment)?;
self.buffers.insert(vid, buf);
self.buffer_shapes.insert(vid, dims.to_vec());
Ok(())
}
fn resolve_all(
&self,
bindings: &HashMap<SymbolId, usize>,
) -> Result<HashMap<ValueId, Vec<usize>>> {
let mut resolved = HashMap::with_capacity(self.value_shapes.len());
for (&vid, shape) in &self.value_shapes {
match substitute(shape, bindings) {
Some(dims) => {
resolved.insert(vid, dims);
}
None => {
let value = self.graph.value(vid);
let name = value
.name
.clone()
.unwrap_or_else(|| format!("value#{}", vid.0));
let op = value
.producer
.map(|nid| self.graph.node(nid).op_type.clone())
.unwrap_or_else(|| "<graph input>".to_string());
return Err(SessionError::UnresolvedShape { value: name, op });
}
}
}
Ok(resolved)
}
fn resolve_soft(&self, bindings: &HashMap<SymbolId, usize>) -> HashMap<ValueId, Vec<usize>> {
let mut resolved = HashMap::with_capacity(self.value_shapes.len());
for (&vid, shape) in &self.value_shapes {
if let Some(dims) = substitute(shape, bindings) {
resolved.insert(vid, dims);
}
}
resolved
}
fn size_buffers(&mut self, resolved: &HashMap<ValueId, Vec<usize>>) -> Result<()> {
let vids: Vec<ValueId> = self.value_shapes.keys().copied().collect();
for vid in vids {
if self.graph.initializers.contains_key(&vid) {
continue;
}
let dtype = self.value_dtypes[&vid];
let Some(dims) = resolved.get(&vid).cloned() else {
continue;
};
self.ensure_buffer(vid, dtype, &dims)?;
}
Ok(())
}
fn node_input_shapes(
plan: &NodePlan,
resolved: &HashMap<ValueId, Vec<usize>>,
) -> Vec<Vec<usize>> {
plan.inputs.iter().map(|v| resolved[v].clone()).collect()
}
fn node_output_shapes(
plan: &NodePlan,
resolved: &HashMap<ValueId, Vec<usize>>,
) -> Vec<Vec<usize>> {
plan.outputs.iter().map(|v| resolved[v].clone()).collect()
}
fn compile_all(&mut self, resolved: &HashMap<ValueId, Vec<usize>>) -> Result<()> {
for i in 0..self.plan.len() {
let node_id = self.plan[i].node_id;
let input_shapes = Self::node_input_shapes(&self.plan[i], resolved);
let node = self.graph.node(node_id);
let opset = effective_opset(&self.graph, node);
self.cache
.get_or_create(node_id, node, &input_shapes, opset, &self.ep)?;
}
Ok(())
}
pub(crate) fn cache_stats(&self) -> CacheStats {
self.cache.stats()
}
pub(crate) fn graph(&self) -> &Graph {
&self.graph
}
pub(crate) fn weights(&self) -> &Arc<WeightStore> {
&self._weights
}
pub(crate) fn warmup(&mut self) -> Result<()> {
if self.has_symbols {
return Ok(());
}
let empty = HashMap::new();
let resolved = self.resolve_all(&empty)?;
self.compile_all(&resolved)
}
fn bind_symbols(
&self,
inputs: &[(&str, &Tensor)],
) -> Result<HashMap<SymbolId, usize>> {
let mut bindings: HashMap<SymbolId, usize> = HashMap::new();
for (name, tensor) in inputs {
let vid = *self
.input_index
.get(*name)
.ok_or_else(|| SessionError::InputNotFound {
name: (*name).to_string(),
})?;
let want_dtype = self.value_dtypes[&vid];
if tensor.dtype != want_dtype {
return Err(SessionError::DtypeMismatch {
name: (*name).to_string(),
expected: format!("{want_dtype:?}"),
got: format!("{:?}", tensor.dtype),
});
}
let decl = &self.value_shapes[&vid];
if decl.len() != tensor.shape.len() {
return Err(SessionError::RankMismatch {
name: (*name).to_string(),
expected: decl.len(),
got: tensor.shape.len(),
});
}
for (dim, &actual) in decl.iter().zip(&tensor.shape) {
match dim {
Dim::Static(n) => {
if *n != actual {
return Err(SessionError::ShapeMismatch {
name: (*name).to_string(),
expected: as_static_shape(decl).unwrap_or_default(),
got: tensor.shape.clone(),
});
}
}
Dim::Symbolic(s) => {
if let Some(&prev) = bindings.get(s) {
if prev != actual {
let sym = self
.symbol_name(*s)
.unwrap_or_else(|| format!("symbol#{}", s.0));
return Err(SessionError::SymbolConflict {
symbol: sym,
first: prev,
second: actual,
});
}
} else {
bindings.insert(*s, actual);
}
}
}
}
}
Ok(bindings)
}
fn symbol_name(&self, s: SymbolId) -> Option<String> {
self.graph
.symbol_constraints
.get(&s)
.and_then(|c| c.name.clone())
}
pub(crate) fn run(&mut self, inputs: &[(&str, &Tensor)]) -> Result<Vec<Tensor>> {
let bindings = self.bind_symbols(inputs)?;
let provided: Vec<ValueId> = inputs
.iter()
.filter_map(|(name, _)| self.input_index.get(*name).copied())
.collect();
for &vid in &self.required_inputs {
if !provided.contains(&vid) {
let name = self
.graph
.value(vid)
.name
.clone()
.unwrap_or_else(|| format!("value#{}", vid.0));
return Err(SessionError::InputNotFound { name });
}
}
let mut resolved = self.resolve_soft(&bindings);
self.size_buffers(&resolved)?;
for (name, tensor) in inputs {
let vid = self.input_index[*name];
let buf = self
.buffers
.get_mut(&vid)
.expect("input value has a buffer");
write_host(buf, tensor.as_bytes())?;
}
let graph = &self.graph;
let ep = self.ep.clone();
let cache = &mut self.cache;
let buffers = &mut self.buffers;
for np in &self.plan {
let input_shapes = Self::node_input_shapes(np, &resolved);
if np.outputs.iter().any(|v| !resolved.contains_key(v)) {
let input_values: Vec<Option<Vec<i64>>> = np
.inputs
.iter()
.enumerate()
.map(|(i, v)| {
buffers
.get(v)
.and_then(|b| buffer_as_i64(b, np.input_dtypes[i]))
})
.collect();
let node = graph.node(np.node_id);
let out_shapes = dynamic_output_shapes(node, &input_shapes, &input_values)
.ok_or_else(|| {
let vid = np
.outputs
.iter()
.find(|v| !resolved.contains_key(v))
.copied()
.unwrap_or(np.outputs[0]);
let value = graph.value(vid);
SessionError::UnresolvedShape {
value: value
.name
.clone()
.unwrap_or_else(|| format!("value#{}", vid.0)),
op: node.op_type.clone(),
}
})?;
if out_shapes.len() != np.outputs.len() {
return Err(SessionError::OutputShapeCountMismatch {
op: node.op_type.clone(),
expected: np.outputs.len(),
got: out_shapes.len(),
});
}
for (oi, &ovid) in np.outputs.iter().enumerate() {
let dims = out_shapes[oi].clone();
let numel = checked_numel(&dims, || format!("value#{}", ovid.0))?;
let need = checked_storage_bytes(
np.output_dtypes[oi],
numel,
|| format!("value#{}", ovid.0),
&dims,
)?
.max(1);
let fits = buffers.get(&ovid).map(|b| b.len() == need).unwrap_or(false);
if !fits {
if let Some(old) = buffers.remove(&ovid) {
ep.deallocate(old)?;
}
let buf = ep.allocate(need, TensorLayout::contiguous().alignment)?;
buffers.insert(ovid, buf);
}
resolved.insert(ovid, dims);
}
}
let output_shapes = Self::node_output_shapes(np, &resolved);
let in_strides: Vec<Vec<i64>> = input_shapes
.iter()
.map(|s| compute_contiguous_strides(s))
.collect();
let out_strides: Vec<Vec<i64>> = output_shapes
.iter()
.map(|s| compute_contiguous_strides(s))
.collect();
let mut in_ptrs: Vec<*const std::ffi::c_void> = Vec::with_capacity(np.inputs.len());
for (i, &vid) in np.inputs.iter().enumerate() {
let buf = buffers.get(&vid).ok_or_else(|| {
SessionError::Internal(format!("missing buffer for input value#{}", vid.0))
})?;
view_bounds(
&input_shapes[i],
&in_strides[i],
0,
np.input_dtypes[i],
buf.len(),
)?;
in_ptrs.push(buf.as_ptr());
}
let mut out_bufs: Vec<(ValueId, DeviceBuffer)> = Vec::with_capacity(np.outputs.len());
for &vid in &np.outputs {
let buf = buffers.remove(&vid).ok_or_else(|| {
SessionError::Internal(format!("missing buffer for output value#{}", vid.0))
})?;
out_bufs.push((vid, buf));
}
let mut views: Vec<TensorView> = Vec::with_capacity(np.inputs.len());
for i in 0..np.inputs.len() {
views.push(TensorView::new(
DevicePtr(in_ptrs[i]),
np.input_dtypes[i],
&input_shapes[i],
&in_strides[i],
onnx_runtime_ir::DeviceId::cpu(),
));
}
let mut outs: Vec<TensorMut> = Vec::with_capacity(out_bufs.len());
for (i, (_, buf)) in out_bufs.iter_mut().enumerate() {
view_bounds(
&output_shapes[i],
&out_strides[i],
0,
np.output_dtypes[i],
buf.len(),
)?;
let ptr = buf.as_mut_ptr();
outs.push(TensorMut::new(
DevicePtrMut(ptr),
np.output_dtypes[i],
&output_shapes[i],
&out_strides[i],
onnx_runtime_ir::DeviceId::cpu(),
));
}
let node = graph.node(np.node_id);
let opset = effective_opset(graph, node);
let kernel = cache.get_or_create(np.node_id, node, &input_shapes, opset, &ep)?;
kernel.execute(&views, &mut outs)?;
drop(views);
drop(outs);
for (vid, buf) in out_bufs {
buffers.insert(vid, buf);
}
}
let mut results = Vec::with_capacity(self.graph.outputs.len());
for &vid in &self.graph.outputs {
let dtype = self.value_dtypes[&vid];
let shape = resolved[&vid].clone();
let buf = self.buffers.get(&vid).ok_or_else(|| {
SessionError::Internal(format!("output value#{} not produced", vid.0))
})?;
let n = dtype.storage_bytes(shape.iter().product());
let bytes = &host_bytes(buf)[..n];
results.push(Tensor::from_raw_in(self.ep.clone(), dtype, shape, bytes)?);
}
Ok(results)
}
}
impl Drop for Executor {
fn drop(&mut self) {
for (_, buf) in self.buffers.drain() {
let _ = self.ep.deallocate(buf);
}
}
}
pub(crate) fn auto_detect_cpu_ep() -> Result<Arc<CpuExecutionProvider>> {
let mut ep = CpuExecutionProvider::new();
ep.initialize(&Default::default())?;
Ok(Arc::new(ep))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn view_bounds_rejects_out_of_bounds_view() {
let shape = [2usize, 3];
let strides = compute_contiguous_strides(&shape);
let err = view_bounds(&shape, &strides, 0, DataType::Float32, 16);
assert!(err.is_err(), "gate must reject an oversized view");
assert!(view_bounds(&shape, &strides, 0, DataType::Float32, 24).is_ok());
}
#[test]
fn view_bounds_rejects_offset_overrun() {
let shape = [4usize];
let strides = compute_contiguous_strides(&shape);
assert!(view_bounds(&shape, &strides, 8, DataType::Float32, 16).is_err());
assert!(view_bounds(&shape, &strides, 0, DataType::Float32, 16).is_ok());
}
#[test]
fn substitute_resolves_bound_symbols_only() {
let mut bindings = HashMap::new();
bindings.insert(SymbolId(0), 7usize);
let shape = vec![Dim::Symbolic(SymbolId(0)), Dim::Static(4)];
assert_eq!(substitute(&shape, &bindings), Some(vec![7, 4]));
let unbound = vec![Dim::Symbolic(SymbolId(1)), Dim::Static(4)];
assert_eq!(substitute(&unbound, &bindings), None);
}
#[test]
fn checked_numel_detects_overflow() {
assert_eq!(checked_numel(&[2, 3, 4], || "v".into()).unwrap(), 24);
assert_eq!(checked_numel(&[], || "v".into()).unwrap(), 1);
let huge = [usize::MAX, 2];
let err = checked_numel(&huge, || "value#9".into());
assert!(matches!(
err,
Err(SessionError::ShapeOverflow { .. })
));
}
#[test]
fn checked_storage_bytes_detects_byte_overflow() {
let numel = usize::MAX / 4;
let err = checked_storage_bytes(DataType::Float64, numel, || "value#9".into(), &[numel]);
assert!(matches!(err, Err(SessionError::ShapeOverflow { .. })));
assert_eq!(
checked_storage_bytes(DataType::Float32, 4, || "v".into(), &[4]).unwrap(),
16
);
}
#[test]
fn dynamic_output_shapes_slice_is_single_output() {
let node = Node::new(NodeId(0), "Slice", vec![], vec![]);
let input_shapes = vec![vec![4usize, 2]];
let input_values = vec![
None, Some(vec![1]), Some(vec![3]), Some(vec![0]), Some(vec![1]), ];
let out = dynamic_output_shapes(&node, &input_shapes, &input_values).unwrap();
assert_eq!(out.len(), 1, "Slice must resolve exactly one output shape");
assert_eq!(out[0], vec![2, 2]);
let other = Node::new(NodeId(1), "Conv", vec![], vec![]);
assert!(dynamic_output_shapes(&other, &input_shapes, &input_values).is_none());
}
#[test]
fn effective_opset_reads_graph_import() {
let mut graph = Graph::default();
graph.opset_imports.insert(String::new(), 12);
let node = Node::new(NodeId(0), "Softmax", vec![], vec![]);
assert_eq!(effective_opset(&graph, &node), 12);
let empty = Graph::default();
assert_eq!(effective_opset(&empty, &node), u64::MAX);
}
}