use std::collections::{HashMap, HashSet};
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 onnx_runtime_shape_inference::{InferenceRegistry, MergePolicy};
use crate::error::{Result, SessionError};
use crate::sequence::{
concat_axis, split_axis, stack_new_axis, SeqTensor, SequenceValue,
};
use crate::tensor::{host_bytes, write_host, Tensor};
#[derive(Debug)]
pub(crate) struct NodePlan {
pub node_id: NodeId,
pub inputs: Vec<Option<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(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct ControlFlowStats {
pub subgraph_builds: u64,
pub subgraph_runs: 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::unsupported_op(node, node_id, opset, ep.name()));
}
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,
name_index: HashMap<String, ValueId>,
subgraph_execs: HashMap<(NodeId, String), CompiledSubgraph>,
control_flow_stats: ControlFlowStats,
views: HashMap<ValueId, ValueView>,
pinned: HashSet<ValueId>,
sequence_values: HashSet<ValueId>,
sequences: HashMap<ValueId, SequenceValue>,
seq_elem_values: HashMap<ValueId, Arc<SeqTensor>>,
}
#[derive(Clone, Debug)]
struct ValueView {
source: ValueId,
shape: Vec<usize>,
strides: Vec<i64>,
byte_offset: usize,
}
struct InInfo {
present: bool,
dtype: DataType,
shape: Vec<usize>,
strides: Vec<i64>,
byte_offset: usize,
base_ptr: *const std::ffi::c_void,
root_len: usize,
}
struct CompiledSubgraph {
exec: Executor,
input_names: Vec<String>,
built_shapes: Vec<Vec<usize>>,
}
struct PreparedSubgraph {
key: (NodeId, String),
formal_names: Vec<String>,
capture_names: Vec<String>,
captures: HashMap<String, Tensor>,
}
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 gather_view(
src: &[u8],
shape: &[usize],
strides: &[i64],
byte_offset: usize,
esize: usize,
) -> Vec<u8> {
let n: usize = shape.iter().product();
let mut out = vec![0u8; n * esize];
if n == 0 {
return out;
}
let rank = shape.len();
let mut idx = vec![0usize; rank];
let mut w = 0usize;
loop {
let mut off = byte_offset as i64;
for d in 0..rank {
off += strides[d] * idx[d] as i64 * esize as i64;
}
let s = off as usize;
out[w..w + esize].copy_from_slice(&src[s..s + esize]);
w += esize;
let mut carried = true;
for axis in (0..rank).rev() {
idx[axis] += 1;
if idx[axis] < shape[axis] {
carried = false;
break;
}
idx[axis] = 0;
}
if carried {
break;
}
}
out
}
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_else(|| {
unreachable!(
"internal invariant violated: node #{} ({}::{}) has no opset import",
node.id.0,
if node.domain.is_empty() {
"ai.onnx"
} else {
&node.domain
},
node.op_type
)
})
}
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>> {
bytes_as_i64(crate::tensor::host_bytes(buffer), dtype)
}
fn bytes_as_i64(bytes: &[u8], dtype: DataType) -> Option<Vec<i64>> {
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();
let init_align = TensorLayout::contiguous().alignment;
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 producer_less = graph.value(vid).producer.is_none();
let buf = if producer_less
&& !bytes.is_empty()
&& (bytes.as_ptr() as usize).is_multiple_of(init_align)
{
unsafe {
DeviceBuffer::from_borrowed_parts(
bytes.as_ptr() as *mut std::ffi::c_void,
ep.device_id(),
bytes.len(),
init_align,
)
}
} else {
let mut owned = ep.allocate(bytes.len().max(1), init_align)?;
write_host(&mut owned, bytes)?;
owned
};
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 sequence_values: HashSet<ValueId> = HashSet::new();
for &nid in &order {
let node = graph.node(nid);
if produces_sequence_output(&node.op_type, &node.domain) {
for &out in &node.outputs {
sequence_values.insert(out);
}
}
}
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 mut slots: Vec<Option<ValueId>> = node.inputs.clone();
while matches!(slots.last(), Some(None)) {
slots.pop();
}
let inputs = slots;
let outputs: Vec<ValueId> = node.outputs.clone();
let input_dtypes: Vec<DataType> = inputs
.iter()
.map(|v| v.map(|vid| value_dtypes[&vid]).unwrap_or(DataType::Float32))
.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 name_index = HashMap::new();
for (vid, value) in graph.values.iter() {
if let Some(name) = &value.name {
name_index.insert(name.clone(), vid);
}
}
let mut exec = Self {
graph,
weights,
ep,
buffers,
buffer_shapes,
value_shapes,
value_dtypes,
plan,
input_index,
required_inputs,
has_symbols,
cache: KernelCache::default(),
name_index,
subgraph_execs: HashMap::new(),
control_flow_stats: ControlFlowStats::default(),
views: HashMap::new(),
pinned: HashSet::new(),
sequence_values,
sequences: HashMap::new(),
seq_elem_values: HashMap::new(),
};
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 {
if self.sequence_values.contains(&vid) {
continue;
}
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;
}
if self.sequence_values.contains(&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| v.map(|vid| resolved[&vid].clone()).unwrap_or_default())
.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 node = self.graph.node(node_id);
if is_control_flow_op(&node.op_type, &node.domain) {
continue;
}
if is_sequence_op(&node.op_type, &node.domain) {
continue;
}
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 control_flow_stats(&self) -> ControlFlowStats {
self.control_flow_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>> {
self.run_scoped(inputs, &HashMap::new())
}
fn run_scoped(
&mut self,
inputs: &[(&str, &Tensor)],
outer_scope: &HashMap<String, Tensor>,
) -> Result<Vec<Tensor>> {
self.views.clear();
self.pinned.clear();
self.sequences.clear();
self.seq_elem_values.clear();
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())?;
}
for pi in 0..self.plan.len() {
let node_id = self.plan[pi].node_id;
let node = self.graph.node(node_id);
if is_control_flow_op(&node.op_type, &node.domain) {
self.exec_control_flow(pi, &mut resolved, outer_scope)?;
} else if is_sequence_op(&node.op_type, &node.domain) {
self.exec_sequence_node(pi, &mut resolved)?;
} else {
self.exec_kernel_node(pi, &mut resolved)?;
}
}
let mut results = Vec::with_capacity(self.graph.outputs.len());
for &vid in &self.graph.outputs {
if self.sequence_values.contains(&vid) {
let name = self
.graph
.try_value(vid)
.and_then(|v| v.name.clone())
.unwrap_or_else(|| format!("value#{}", vid.0));
return Err(SessionError::SequenceOp {
op: "<graph output>".to_string(),
reason: format!(
"graph output {name} is a Sequence value, which cannot be \
returned through the tensor `run` API. To fix: end the graph \
with ConcatFromSequence or SequenceAt to produce tensor \
output(s)"
),
});
}
let dtype = self.value_dtypes[&vid];
let shape = resolved[&vid].clone();
let bytes = self.contiguous_bytes(vid, &shape, dtype)?;
results.push(Tensor::from_raw_in(self.ep.clone(), dtype, shape, &bytes)?);
}
Ok(results)
}
fn exec_kernel_node(
&mut self,
pi: usize,
resolved: &mut HashMap<ValueId, Vec<usize>>,
) -> Result<()> {
let node_id = self.plan[pi].node_id;
let inputs = self.plan[pi].inputs.clone();
let outputs = self.plan[pi].outputs.clone();
let input_dtypes = self.plan[pi].input_dtypes.clone();
let output_dtypes = self.plan[pi].output_dtypes.clone();
let input_shapes: Vec<Vec<usize>> = inputs
.iter()
.map(|v| v.map(|vid| resolved[&vid].clone()).unwrap_or_default())
.collect();
if outputs.iter().any(|v| !resolved.contains_key(v)) {
let input_values: Vec<Option<Vec<i64>>> = inputs
.iter()
.enumerate()
.map(|(i, v)| {
v.and_then(|vid| self.input_i64(vid, &input_shapes[i], input_dtypes[i]))
})
.collect();
let node = self.graph.node(node_id);
let out_shapes = dynamic_output_shapes(node, &input_shapes, &input_values)
.ok_or_else(|| {
let vid = outputs
.iter()
.find(|v| !resolved.contains_key(v))
.copied()
.unwrap_or(outputs[0]);
let value = self.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() != outputs.len() {
return Err(SessionError::OutputShapeCountMismatch {
op: self.graph.node(node_id).op_type.clone(),
expected: outputs.len(),
got: out_shapes.len(),
});
}
for (oi, &ovid) in outputs.iter().enumerate() {
resolved.insert(ovid, out_shapes[oi].clone());
}
}
let output_shapes: Vec<Vec<usize>> =
outputs.iter().map(|v| resolved[v].clone()).collect();
let mut in_infos: Vec<InInfo> = Vec::with_capacity(inputs.len());
for (i, slot) in inputs.iter().enumerate() {
let Some(vid) = *slot else {
in_infos.push(InInfo {
present: false,
dtype: input_dtypes[i],
shape: Vec::new(),
strides: Vec::new(),
byte_offset: 0,
base_ptr: std::ptr::null(),
root_len: 0,
});
continue;
};
if let Some(elem) = self.seq_elem_values.get(&vid) {
let shape = input_shapes[i].clone();
let strides = compute_contiguous_strides(&shape);
let root_len = elem.data.len();
let base_ptr = elem.as_ptr() as *const std::ffi::c_void;
view_bounds(&shape, &strides, 0, input_dtypes[i], root_len)?;
in_infos.push(InInfo {
present: true,
dtype: input_dtypes[i],
shape,
strides,
byte_offset: 0,
base_ptr,
root_len,
});
continue;
}
let root = self.root_of(vid);
let buf = self.buffers.get(&root).ok_or_else(|| {
SessionError::Internal(format!("missing buffer for input value#{}", vid.0))
})?;
let root_len = buf.len();
let base_ptr = buf.as_ptr();
let (shape, strides, byte_offset) = match self.views.get(&vid) {
Some(view) => (view.shape.clone(), view.strides.clone(), view.byte_offset),
None => {
let shape = input_shapes[i].clone();
let strides = compute_contiguous_strides(&shape);
(shape, strides, 0)
}
};
view_bounds(&shape, &strides, byte_offset, input_dtypes[i], root_len)?;
in_infos.push(InInfo {
present: true,
dtype: input_dtypes[i],
shape,
strides,
byte_offset,
base_ptr,
root_len,
});
}
let ep = self.ep.clone();
let graph = &self.graph;
let cache = &mut self.cache;
let buffers = &mut self.buffers;
let buffer_shapes = &mut self.buffer_shapes;
let views_meta = &mut self.views;
let pinned = &mut self.pinned;
let mut views: Vec<TensorView> = Vec::with_capacity(in_infos.len());
for info in &in_infos {
if !info.present {
views.push(TensorView::absent(info.dtype));
continue;
}
views.push(
TensorView::new(
DevicePtr(info.base_ptr),
info.dtype,
&info.shape,
&info.strides,
onnx_runtime_ir::DeviceId::cpu(),
)
.with_byte_offset(info.byte_offset),
);
}
let node = graph.node(node_id);
let opset = effective_opset(graph, node);
let kernel = cache.get_or_create(node_id, node, &input_shapes, opset, &ep)?;
if let Some(specs) = kernel.view_outputs(&views, outputs.len()) {
drop(views);
if specs.len() != outputs.len() {
return Err(SessionError::Internal(format!(
"op '{}' returned {} view outputs for {} outputs",
node.op_type,
specs.len(),
outputs.len()
)));
}
for (oi, spec) in specs.into_iter().enumerate() {
let ovid = outputs[oi];
let Some(in_vid) = inputs.get(spec.input_index).copied().flatten() else {
return Err(SessionError::Internal(format!(
"op '{}' view output {} references invalid input index {}",
node.op_type, oi, spec.input_index
)));
};
let root = match views_meta.get(&in_vid) {
Some(v) => v.source,
None => in_vid,
};
let root_len = buffers.get(&root).map(|b| b.len()).ok_or_else(|| {
SessionError::Internal(format!(
"view source value#{} has no buffer",
root.0
))
})?;
view_bounds(
&spec.shape,
&spec.strides,
spec.byte_offset,
output_dtypes[oi],
root_len,
)?;
debug_assert!(
!pinned.contains(&ovid),
"value#{} is pinned as a live view source yet is being reproduced",
ovid.0
);
if let Some(old) = buffers.remove(&ovid) {
ep.deallocate(old)?;
}
buffer_shapes.remove(&ovid);
views_meta.insert(
ovid,
ValueView {
source: root,
shape: spec.shape.clone(),
strides: spec.strides,
byte_offset: spec.byte_offset,
},
);
pinned.insert(root);
resolved.insert(ovid, spec.shape);
}
return Ok(());
}
for (oi, &ovid) in outputs.iter().enumerate() {
let dims = &output_shapes[oi];
let numel = checked_numel(dims, || format!("value#{}", ovid.0))?;
let need = checked_storage_bytes(
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 {
debug_assert!(
!pinned.contains(&ovid),
"value#{} is pinned as a live view source yet is being resized",
ovid.0
);
if let Some(old) = buffers.remove(&ovid) {
ep.deallocate(old)?;
}
let buf = ep.allocate(need, TensorLayout::contiguous().alignment)?;
buffers.insert(ovid, buf);
}
}
let mut mat: Vec<Option<(Vec<u8>, Vec<i64>)>> = Vec::with_capacity(in_infos.len());
for (i, info) in in_infos.iter().enumerate() {
if !info.present {
mat.push(None);
continue;
}
let contiguous = onnx_runtime_ir::is_contiguous(&info.shape, &info.strides);
if contiguous || kernel.supports_strided_input(i) {
mat.push(None);
continue;
}
let esize = info.dtype.byte_size();
if esize == 0 {
return Err(SessionError::from(
onnx_runtime_ep_api::EpError::InvalidTensorView {
reason: format!(
"cannot materialize sub-byte strided input {i} of op '{}'",
node.op_type
),
},
));
}
let src = unsafe {
std::slice::from_raw_parts(info.base_ptr as *const u8, info.root_len)
};
let gathered = gather_view(src, &info.shape, &info.strides, info.byte_offset, esize);
let strides = compute_contiguous_strides(&info.shape);
mat.push(Some((gathered, strides)));
}
drop(views);
let mut views: Vec<TensorView> = Vec::with_capacity(in_infos.len());
for (i, info) in in_infos.iter().enumerate() {
if !info.present {
views.push(TensorView::absent(info.dtype));
continue;
}
match &mat[i] {
Some((buf, strides)) => views.push(TensorView::new(
DevicePtr(buf.as_ptr() as *const std::ffi::c_void),
info.dtype,
&info.shape,
strides,
onnx_runtime_ir::DeviceId::cpu(),
)),
None => views.push(
TensorView::new(
DevicePtr(info.base_ptr),
info.dtype,
&info.shape,
&info.strides,
onnx_runtime_ir::DeviceId::cpu(),
)
.with_byte_offset(info.byte_offset),
),
}
}
let out_strides: Vec<Vec<i64>> = output_shapes
.iter()
.map(|s| compute_contiguous_strides(s))
.collect();
let mut out_bufs: Vec<(ValueId, DeviceBuffer)> = Vec::with_capacity(outputs.len());
for &vid in &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 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,
output_dtypes[i],
buf.len(),
)?;
let ptr = buf.as_mut_ptr();
outs.push(TensorMut::new(
DevicePtrMut(ptr),
output_dtypes[i],
&output_shapes[i],
&out_strides[i],
onnx_runtime_ir::DeviceId::cpu(),
));
}
kernel.execute(&views, &mut outs)?;
drop(views);
drop(outs);
for (vid, buf) in out_bufs {
buffers.insert(vid, buf);
}
Ok(())
}
fn input_i64(&self, vid: ValueId, shape: &[usize], dtype: DataType) -> Option<Vec<i64>> {
if self.views.contains_key(&vid) || self.seq_elem_values.contains_key(&vid) {
let bytes = self.contiguous_bytes(vid, shape, dtype).ok()?;
bytes_as_i64(&bytes, dtype)
} else {
self.buffers.get(&vid).and_then(|b| buffer_as_i64(b, dtype))
}
}
}
impl Executor {
fn exec_sequence_node(
&mut self,
pi: usize,
resolved: &mut HashMap<ValueId, Vec<usize>>,
) -> Result<()> {
let node_id = self.plan[pi].node_id;
let inputs = self.plan[pi].inputs.clone();
let outputs = self.plan[pi].outputs.clone();
let op = self.graph.node(node_id).op_type.clone();
match op.as_str() {
"SequenceEmpty" => {
let dtype_attr = self.graph.node(node_id).attr("dtype").and_then(|a| a.as_int());
let dtype = match dtype_attr {
None => DataType::Float32, Some(raw) => DataType::from_onnx(raw as i32).ok_or_else(|| {
SessionError::SequenceOp {
op: op.clone(),
reason: format!(
"attribute 'dtype' = {raw} is not a known ONNX \
TensorProto.DataType. To fix: use a valid element \
dtype id (e.g. 1=float32, 7=int64)"
),
}
})?,
};
self.sequences
.insert(outputs[0], SequenceValue::empty(dtype));
Ok(())
}
"SequenceConstruct" => {
let mut items = Vec::with_capacity(inputs.len());
for slot in &inputs {
let vid = slot.ok_or_else(|| self.seq_missing_input(&op))?;
items.push(self.read_seq_element(vid, resolved)?);
}
let seq = SequenceValue::construct(items).map_err(seq_err)?;
self.sequences.insert(outputs[0], seq);
Ok(())
}
"SequenceInsert" => {
let seq = self.get_sequence(inputs.first().copied().flatten(), &op)?;
let tvid = inputs.get(1).copied().flatten().ok_or_else(|| {
self.seq_missing_input(&op)
})?;
let tensor = self.read_seq_element(tvid, resolved)?;
let position = match inputs.get(2).copied().flatten() {
Some(pvid) => Some(self.read_scalar_i64(pvid, resolved, &op)?),
None => None,
};
let out = seq.insert(tensor, position).map_err(seq_err)?;
self.sequences.insert(outputs[0], out);
Ok(())
}
"SequenceErase" => {
let seq = self.get_sequence(inputs.first().copied().flatten(), &op)?;
let position = match inputs.get(1).copied().flatten() {
Some(pvid) => Some(self.read_scalar_i64(pvid, resolved, &op)?),
None => None,
};
let out = seq.erase(position).map_err(seq_err)?;
self.sequences.insert(outputs[0], out);
Ok(())
}
"SequenceAt" => {
let seq = self.get_sequence(inputs.first().copied().flatten(), &op)?;
let pvid = inputs.get(1).copied().flatten().ok_or_else(|| {
SessionError::SequenceOp {
op: op.clone(),
reason: "requires a 'position' input. To fix: supply the \
index tensor of the element to read"
.to_string(),
}
})?;
let pos = self.read_scalar_i64(pvid, resolved, &op)?;
let elem = seq.at(pos).map_err(seq_err)?;
self.store_seq_element_output(outputs[0], elem, resolved)
}
"SequenceLength" => {
let seq = self.get_sequence(inputs.first().copied().flatten(), &op)?;
let len = seq.len() as i64;
self.store_raw_tensor_output(
outputs[0],
DataType::Int64,
Vec::new(),
&len.to_le_bytes(),
resolved,
)
}
"SplitToSequence" => self.exec_split_to_sequence(&op, &inputs, &outputs, resolved),
"ConcatFromSequence" => {
self.exec_concat_from_sequence(node_id, &op, &inputs, &outputs, resolved)
}
other => Err(SessionError::SequenceOp {
op: other.to_string(),
reason: "unrecognized Sequence op (executor routing bug)".to_string(),
}),
}
}
fn exec_split_to_sequence(
&mut self,
op: &str,
inputs: &[Option<ValueId>],
outputs: &[ValueId],
resolved: &mut HashMap<ValueId, Vec<usize>>,
) -> Result<()> {
let node = self.graph.node(self.plan_node_of(outputs[0]));
let axis_attr = node.attr("axis").and_then(|a| a.as_int()).unwrap_or(0);
let keepdims = node.attr("keepdims").and_then(|a| a.as_int()).unwrap_or(1) != 0;
let ivid = inputs.first().copied().flatten().ok_or_else(|| self.seq_missing_input(op))?;
let dtype = self.value_dtypes[&ivid];
let esize = dtype.byte_size();
if esize == 0 {
return Err(SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"sub-byte dtype {dtype:?} is not supported for SplitToSequence. \
To fix: Cast to a byte-addressable dtype before splitting"
),
});
}
let shape = resolved
.get(&ivid)
.cloned()
.ok_or_else(|| self.seq_unresolved(op, ivid))?;
let rank = shape.len();
if rank == 0 {
return Err(SessionError::SequenceOp {
op: op.to_string(),
reason: "cannot split a scalar (rank-0) tensor. To fix: split a \
tensor with at least one dimension"
.to_string(),
});
}
let axis = normalize_axis(axis_attr, rank).ok_or_else(|| SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"attribute 'axis' = {axis_attr} is out of range for a rank-{rank} \
input (valid range is [{}, {}])",
-(rank as i64),
rank as i64 - 1
),
})?;
let axis_dim = shape[axis];
let bytes = self.contiguous_bytes(ivid, &shape, dtype)?;
let mut squeeze = false;
let sizes: Vec<usize> = match inputs.get(1).copied().flatten() {
None => {
squeeze = !keepdims;
vec![1; axis_dim]
}
Some(svid) => {
let sshape = resolved.get(&svid).cloned().unwrap_or_default();
let svals = self.read_i64_vec(svid, &sshape, op)?;
let is_scalar = sshape.is_empty();
if is_scalar {
let chunk = *svals.first().ok_or_else(|| SessionError::SequenceOp {
op: op.to_string(),
reason: "'split' scalar is empty".to_string(),
})?;
if chunk <= 0 {
return Err(SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"'split' chunk size {chunk} must be positive"
),
});
}
let chunk = chunk as usize;
let mut v = Vec::new();
let mut rem = axis_dim;
while rem > 0 {
let k = rem.min(chunk);
v.push(k);
rem -= k;
}
v
} else {
let v: Vec<usize> = svals.iter().map(|&x| x.max(0) as usize).collect();
let sum: usize = v.iter().sum();
if sum != axis_dim {
return Err(SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"'split' sizes {v:?} sum to {sum} but axis {axis} has \
extent {axis_dim}. To fix: make the split sizes sum to \
the axis length"
),
});
}
v
}
}
};
let parts = split_axis(&bytes, &shape, axis, &sizes, esize);
let items: Vec<std::sync::Arc<SeqTensor>> = parts
.into_iter()
.map(|(mut sh, data)| {
if squeeze {
sh.remove(axis);
}
SeqTensor::shared(dtype, sh, data)
})
.collect();
self.sequences.insert(
outputs[0],
SequenceValue {
elem_dtype: dtype,
items,
},
);
Ok(())
}
fn exec_concat_from_sequence(
&mut self,
node_id: NodeId,
op: &str,
inputs: &[Option<ValueId>],
outputs: &[ValueId],
resolved: &mut HashMap<ValueId, Vec<usize>>,
) -> Result<()> {
let node = self.graph.node(node_id);
let axis_attr = node.attr("axis").and_then(|a| a.as_int()).ok_or_else(|| {
SessionError::SequenceOp {
op: op.to_string(),
reason: "requires the mandatory 'axis' attribute. To fix: set 'axis'"
.to_string(),
}
})?;
let new_axis = node.attr("new_axis").and_then(|a| a.as_int()).unwrap_or(0) != 0;
let seq = self.get_sequence(inputs.first().copied().flatten(), op)?;
if seq.is_empty() {
return Err(SessionError::SequenceOp {
op: op.to_string(),
reason: "cannot concatenate an empty sequence (output shape is \
undefined). To fix: guard with SequenceLength"
.to_string(),
});
}
let dtype = seq.elem_dtype;
let esize = dtype.byte_size();
if esize == 0 {
return Err(SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"sub-byte dtype {dtype:?} is not supported for ConcatFromSequence"
),
});
}
let elem_shapes: Vec<Vec<usize>> = seq.items.iter().map(|t| t.shape.clone()).collect();
let elem_datas: Vec<&[u8]> = seq.items.iter().map(|t| t.data.as_slice()).collect();
let rank = elem_shapes[0].len();
let (oshape, out) = if new_axis {
for (i, s) in elem_shapes.iter().enumerate() {
if s != &elem_shapes[0] {
return Err(SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"ConcatFromSequence(new_axis=1) requires identical element \
shapes, but element {i} has shape {s:?} vs {:?}",
elem_shapes[0]
),
});
}
}
let axis = normalize_axis(axis_attr, rank + 1).ok_or_else(|| {
SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"'axis' = {axis_attr} is out of range for new_axis=1 stacking \
of rank-{rank} elements (valid range is [{}, {}])",
-(rank as i64) - 1,
rank as i64
),
}
})?;
stack_new_axis(&elem_datas, &elem_shapes[0], axis, esize)
} else {
let axis = normalize_axis(axis_attr, rank).ok_or_else(|| SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"'axis' = {axis_attr} is out of range for rank-{rank} elements \
(valid range is [{}, {}])",
-(rank as i64),
rank as i64 - 1
),
})?;
for (i, s) in elem_shapes.iter().enumerate() {
let mismatch = s.len() != rank
|| s.iter().enumerate().any(|(d, &v)| d != axis && v != elem_shapes[0][d]);
if mismatch {
return Err(SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"ConcatFromSequence requires elements to match on all axes \
except {axis}, but element {i} has shape {s:?} vs {:?}",
elem_shapes[0]
),
});
}
}
concat_axis(&elem_datas, &elem_shapes, axis, esize)
};
drop(seq);
self.store_raw_tensor_output(outputs[0], dtype, oshape, &out, resolved)
}
fn read_seq_element(
&self,
vid: ValueId,
resolved: &HashMap<ValueId, Vec<usize>>,
) -> Result<std::sync::Arc<SeqTensor>> {
if let Some(elem) = self.seq_elem_values.get(&vid) {
return Ok(std::sync::Arc::clone(elem)); }
let dtype = self.value_dtypes[&vid];
let shape = resolved
.get(&vid)
.cloned()
.ok_or_else(|| self.seq_unresolved("Sequence", vid))?;
let bytes = self.contiguous_bytes(vid, &shape, dtype)?;
Ok(SeqTensor::shared(dtype, shape, bytes))
}
fn get_sequence(&self, vid: Option<ValueId>, op: &str) -> Result<SequenceValue> {
let vid = vid.ok_or_else(|| self.seq_missing_input(op))?;
self.sequences.get(&vid).cloned().ok_or_else(|| SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"input value#{} is not a live sequence. To fix: ensure it is produced \
by a Sequence-producing op (SequenceEmpty/Construct/Insert/Erase/\
SplitToSequence)",
vid.0
),
})
}
fn read_scalar_i64(
&self,
vid: ValueId,
resolved: &HashMap<ValueId, Vec<usize>>,
op: &str,
) -> Result<i64> {
let shape = resolved.get(&vid).cloned().unwrap_or_default();
let dtype = self.value_dtypes[&vid];
let vals = self.input_i64(vid, &shape, dtype).ok_or_else(|| SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"position input has dtype {dtype:?}, expected an integer (int32/int64). \
To fix: provide an int64 scalar index"
),
})?;
vals.first().copied().ok_or_else(|| SessionError::SequenceOp {
op: op.to_string(),
reason: "position input is empty; expected a single scalar index".to_string(),
})
}
fn read_i64_vec(&self, vid: ValueId, shape: &[usize], op: &str) -> Result<Vec<i64>> {
let dtype = self.value_dtypes[&vid];
self.input_i64(vid, shape, dtype).ok_or_else(|| SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"'split' input has dtype {dtype:?}, expected int32/int64. To fix: \
provide integer split sizes"
),
})
}
fn store_seq_element_output(
&mut self,
vid: ValueId,
elem: std::sync::Arc<SeqTensor>,
resolved: &mut HashMap<ValueId, Vec<usize>>,
) -> Result<()> {
if let Some(old) = self.buffers.remove(&vid) {
self.ep.deallocate(old)?;
}
self.buffer_shapes.remove(&vid);
self.views.remove(&vid);
resolved.insert(vid, elem.shape.clone());
self.value_dtypes.insert(vid, elem.dtype);
self.seq_elem_values.insert(vid, elem);
Ok(())
}
fn store_raw_tensor_output(
&mut self,
vid: ValueId,
dtype: DataType,
dims: Vec<usize>,
bytes: &[u8],
resolved: &mut HashMap<ValueId, Vec<usize>>,
) -> Result<()> {
self.seq_elem_values.remove(&vid);
self.views.remove(&vid);
let need = bytes.len().max(1);
let fits = self.buffers.get(&vid).map(|b| b.len() == need).unwrap_or(false);
if !fits {
if let Some(old) = self.buffers.remove(&vid) {
self.ep.deallocate(old)?;
}
let buf = self.ep.allocate(need, TensorLayout::contiguous().alignment)?;
self.buffers.insert(vid, buf);
}
let buf = self.buffers.get_mut(&vid).expect("just ensured");
write_host(buf, bytes)?;
self.value_dtypes.insert(vid, dtype);
self.buffer_shapes.insert(vid, dims.clone());
resolved.insert(vid, dims);
Ok(())
}
fn plan_node_of(&self, vid: ValueId) -> NodeId {
self.graph
.value(vid)
.producer
.expect("sequence op output has a producer")
}
fn seq_missing_input(&self, op: &str) -> SessionError {
SessionError::SequenceOp {
op: op.to_string(),
reason: "a required input is missing (omitted None slot). To fix: connect \
all required inputs of this Sequence op"
.to_string(),
}
}
fn seq_unresolved(&self, op: &str, vid: ValueId) -> SessionError {
let name = self
.graph
.try_value(vid)
.and_then(|v| v.name.clone())
.unwrap_or_else(|| format!("value#{}", vid.0));
SessionError::SequenceOp {
op: op.to_string(),
reason: format!(
"input {name} has no resolved shape yet. To fix: ensure its producer \
runs before this Sequence op"
),
}
}
}
fn seq_err(e: crate::sequence::SeqOpError) -> SessionError {
SessionError::SequenceOp {
op: e.op.to_string(),
reason: e.reason,
}
}
fn normalize_axis(axis: i64, rank: usize) -> Option<usize> {
let r = rank as i64;
let a = if axis < 0 { axis + r } else { axis };
if a < 0 || a >= r {
None
} else {
Some(a as usize)
}
}
fn is_control_flow_op(op_type: &str, domain: &str) -> bool {
(domain.is_empty() || domain == "ai.onnx") && matches!(op_type, "If" | "Loop" | "Scan")
}
fn is_sequence_op(op_type: &str, domain: &str) -> bool {
(domain.is_empty() || domain == "ai.onnx")
&& matches!(
op_type,
"SequenceEmpty"
| "SequenceConstruct"
| "SequenceInsert"
| "SequenceErase"
| "SequenceAt"
| "SequenceLength"
| "SplitToSequence"
| "ConcatFromSequence"
)
}
fn produces_sequence_output(op_type: &str, domain: &str) -> bool {
(domain.is_empty() || domain == "ai.onnx")
&& matches!(
op_type,
"SequenceEmpty"
| "SequenceConstruct"
| "SequenceInsert"
| "SequenceErase"
| "SplitToSequence"
)
}
fn tensor_scalar_i64(t: &Tensor) -> Option<i64> {
match t.dtype {
DataType::Int64 => t
.as_bytes()
.get(..8)
.map(|c| i64::from_le_bytes(c.try_into().unwrap())),
DataType::Int32 => t
.as_bytes()
.get(..4)
.map(|c| i32::from_le_bytes(c.try_into().unwrap()) as i64),
_ => None,
}
}
fn tensor_scalar_bool(t: &Tensor) -> Option<bool> {
if t.dtype != DataType::Bool {
return None;
}
t.as_bytes().first().map(|&b| b != 0)
}
fn scalar_i64_tensor(v: i64) -> Result<Tensor> {
Tensor::from_raw(DataType::Int64, vec![], &v.to_le_bytes())
}
fn scalar_bool_tensor(v: bool) -> Result<Tensor> {
Tensor::from_raw(DataType::Bool, vec![], &[u8::from(v)])
}
fn missing_capture_error(attr_key: &str, name: &str) -> SessionError {
SessionError::Internal(format!(
"control-flow body '{attr_key}' captures free variable '{name}', but it is not \
available in the enclosing scope. RULES #1: a subgraph may only reference outer \
values that are graph inputs, initializers, or produced by an upstream node in an \
enclosing graph; '{name}' matches none of these"
))
}
fn required_outer_names(graph: &Graph) -> HashSet<String> {
let formal_set: HashSet<ValueId> = graph.inputs.iter().copied().collect();
let local_names: HashSet<&str> = graph
.values
.iter()
.filter_map(|(_, value)| value.name.as_deref())
.collect();
let mut required = HashSet::new();
for (vid, value) in graph.values.iter() {
if value.producer.is_none()
&& !formal_set.contains(&vid)
&& !graph.initializers.contains_key(&vid)
&& let Some(name) = &value.name
{
required.insert(name.clone());
}
}
for nested in graph.subgraphs.values() {
for name in required_outer_names(nested) {
if !local_names.contains(name.as_str()) {
required.insert(name);
}
}
}
required
}
impl Executor {
fn value_tensor(
&self,
vid: ValueId,
resolved: &HashMap<ValueId, Vec<usize>>,
) -> Result<Tensor> {
let dtype = self.value_dtypes[&vid];
let shape = resolved.get(&vid).cloned().ok_or_else(|| {
let name = self
.graph
.try_value(vid)
.and_then(|v| v.name.clone())
.unwrap_or_else(|| format!("value#{}", vid.0));
SessionError::UnresolvedShape {
value: name,
op: "<control-flow input>".to_string(),
}
})?;
let bytes = self.contiguous_bytes(vid, &shape, dtype)?;
Tensor::from_raw_in(self.ep.clone(), dtype, shape, &bytes)
}
fn root_of(&self, vid: ValueId) -> ValueId {
match self.views.get(&vid) {
Some(v) => v.source,
None => vid,
}
}
fn contiguous_bytes(
&self,
vid: ValueId,
shape: &[usize],
dtype: DataType,
) -> Result<Vec<u8>> {
let numel: usize = shape.iter().product();
let n = dtype.storage_bytes(numel);
if let Some(elem) = self.seq_elem_values.get(&vid) {
return Ok(elem.data[..n.min(elem.data.len())].to_vec());
}
if let Some(view) = self.views.get(&vid) {
let buf = self.buffers.get(&view.source).ok_or_else(|| {
SessionError::Internal(format!(
"view value#{} aliases missing source buffer value#{}",
vid.0, view.source.0
))
})?;
let esize = dtype.byte_size();
if esize == 0 {
return Err(SessionError::Internal(format!(
"cannot materialize sub-byte view value#{}",
vid.0
)));
}
Ok(gather_view(
host_bytes(buf),
&view.shape,
&view.strides,
view.byte_offset,
esize,
))
} else {
let buf = self.buffers.get(&vid).ok_or_else(|| {
SessionError::Internal(format!("value#{} not produced", vid.0))
})?;
Ok(host_bytes(buf)[..n].to_vec())
}
}
fn store_output_tensor(
&mut self,
vid: ValueId,
tensor: &Tensor,
resolved: &mut HashMap<ValueId, Vec<usize>>,
) -> Result<()> {
self.store_output_bytes(
vid,
tensor.dtype,
tensor.shape.clone(),
tensor.as_bytes(),
resolved,
)
}
fn store_output_bytes(
&mut self,
vid: ValueId,
dtype: DataType,
dims: Vec<usize>,
bytes: &[u8],
resolved: &mut HashMap<ValueId, Vec<usize>>,
) -> Result<()> {
let numel = checked_numel(&dims, || format!("value#{}", vid.0))?;
let need = checked_storage_bytes(dtype, numel, || format!("value#{}", vid.0), &dims)?
.max(1);
let fits = self.buffers.get(&vid).map(|b| b.len() == need).unwrap_or(false);
if !fits {
if let Some(old) = self.buffers.remove(&vid) {
self.ep.deallocate(old)?;
}
let buf = self
.ep
.allocate(need, TensorLayout::contiguous().alignment)?;
self.buffers.insert(vid, buf);
}
let buf = self.buffers.get_mut(&vid).expect("just ensured");
write_host(buf, bytes)?;
self.value_dtypes.insert(vid, dtype);
self.buffer_shapes.insert(vid, dims.clone());
resolved.insert(vid, dims);
Ok(())
}
fn prepare_subgraph(
&self,
node_id: NodeId,
attr_key: &str,
resolved: &HashMap<ValueId, Vec<usize>>,
outer_scope: &HashMap<String, Tensor>,
) -> Result<PreparedSubgraph> {
let key = (node_id, attr_key.to_string());
let body = self.graph.subgraphs.get(&key).ok_or_else(|| {
SessionError::Internal(format!(
"control-flow node #{} references missing subgraph '{attr_key}'",
node_id.0
))
})?;
let formal_names: Vec<String> = body
.inputs
.iter()
.map(|&vid| {
body.value(vid)
.name
.clone()
.unwrap_or_else(|| format!("value#{}", vid.0))
})
.collect();
let formal_set: HashSet<ValueId> = body.inputs.iter().copied().collect();
let mut capture_names = Vec::new();
for (vid, value) in body.values.iter() {
if value.producer.is_none()
&& !formal_set.contains(&vid)
&& !body.initializers.contains_key(&vid)
&& let Some(name) = &value.name
{
capture_names.push(name.clone());
}
}
capture_names.sort();
let mut scope_names = required_outer_names(body);
scope_names.extend(capture_names.iter().cloned());
let mut captures = HashMap::with_capacity(scope_names.len());
for name in scope_names {
let tensor = if let Some(&vid) = self.name_index.get(&name) {
let materialized = self.buffers.contains_key(&vid)
|| self.views.contains_key(&vid)
|| self.seq_elem_values.contains_key(&vid);
if resolved.contains_key(&vid) && materialized {
self.value_tensor(vid, resolved)?
} else {
outer_scope.get(&name).cloned().ok_or_else(|| {
missing_capture_error(attr_key, &name)
})?
}
} else {
outer_scope.get(&name).cloned().ok_or_else(|| {
missing_capture_error(attr_key, &name)
})?
};
captures.insert(name, tensor);
}
Ok(PreparedSubgraph {
key,
formal_names,
capture_names,
captures,
})
}
fn build_subgraph_exec(
&self,
prepared: &PreparedSubgraph,
externals: &[&Tensor],
) -> Result<CompiledSubgraph> {
let key = &prepared.key;
let body = self.graph.subgraphs.get(key).ok_or_else(|| {
SessionError::Internal(format!(
"control-flow node #{} has no registered subgraph '{}'",
key.0 .0, key.1
))
})?;
let mut g = body.clone();
let mut body_names: HashMap<String, ValueId> = HashMap::new();
for (vid, value) in g.values.iter() {
if let Some(n) = &value.name {
body_names.insert(n.clone(), vid);
}
}
for cname in &prepared.capture_names {
let vid = *body_names.get(cname).ok_or_else(|| {
SessionError::Internal(format!(
"control-flow body '{}' lost capture value '{cname}'",
key.1
))
})?;
if !g.inputs.contains(&vid) {
g.add_input(vid);
}
}
let all_names = prepared
.formal_names
.iter()
.chain(prepared.capture_names.iter());
for (name, tensor) in all_names.zip(externals.iter()) {
let vid = *body_names.get(name).ok_or_else(|| {
SessionError::Internal(format!(
"control-flow body '{}' missing formal/captured input '{name}'",
key.1
))
})?;
let v = g.value_mut(vid);
v.dtype = tensor.dtype;
v.shape = tensor.shape.iter().map(|&d| Dim::Static(d)).collect();
}
let registry = InferenceRegistry::default_registry();
let opset_imports = self.graph.opset_imports.clone();
registry.infer_graph(&mut g, &opset_imports, MergePolicy::Permissive)?;
let exec = Executor::build(g, self.weights.clone(), self.ep.clone())?;
Ok(CompiledSubgraph {
exec,
input_names: prepared
.formal_names
.iter()
.chain(prepared.capture_names.iter())
.cloned()
.collect(),
built_shapes: externals.iter().map(|t| t.shape.clone()).collect(),
})
}
fn run_subgraph(
&mut self,
prepared: &PreparedSubgraph,
formal_inputs: &[&Tensor],
) -> Result<Vec<Tensor>> {
if prepared.formal_names.len() != formal_inputs.len() {
return Err(SessionError::Internal(format!(
"control-flow body '{}' expects {} formal input(s) but {} were supplied",
prepared.key.1,
prepared.formal_names.len(),
formal_inputs.len()
)));
}
let mut externals: Vec<&Tensor> =
Vec::with_capacity(formal_inputs.len() + prepared.capture_names.len());
externals.extend_from_slice(formal_inputs);
for name in &prepared.capture_names {
externals.push(prepared.captures.get(name).expect("prepared capture must be present"));
}
let rebuild = match self.subgraph_execs.get(&prepared.key) {
Some(cs) => {
cs.built_shapes.len() != externals.len()
|| cs
.built_shapes
.iter()
.zip(externals.iter())
.any(|(built, tensor)| built != &tensor.shape)
}
None => true,
};
if rebuild {
let child = self.build_subgraph_exec(prepared, &externals)?;
self.subgraph_execs.insert(prepared.key.clone(), child);
self.control_flow_stats.subgraph_builds += 1;
}
self.control_flow_stats.subgraph_runs += 1;
let cs = self.subgraph_execs.get_mut(&prepared.key).expect("child present");
let inputs: Vec<(&str, &Tensor)> = cs
.input_names
.iter()
.map(String::as_str)
.zip(externals)
.collect();
cs.exec.run_scoped(&inputs, &prepared.captures)
}
fn exec_control_flow(
&mut self,
pi: usize,
resolved: &mut HashMap<ValueId, Vec<usize>>,
outer_scope: &HashMap<String, Tensor>,
) -> Result<()> {
let node = self.graph.node(self.plan[pi].node_id).clone();
match node.op_type.as_str() {
"If" => self.exec_if(&node, resolved, outer_scope),
"Loop" => self.exec_loop(&node, resolved, outer_scope),
"Scan" => self.exec_scan(&node, resolved, outer_scope),
other => Err(SessionError::Internal(format!(
"exec_control_flow reached non-control-flow op {other:?}"
))),
}
}
fn exec_if(
&mut self,
node: &Node,
resolved: &mut HashMap<ValueId, Vec<usize>>,
outer_scope: &HashMap<String, Tensor>,
) -> Result<()> {
let cond_vid = node.inputs.first().and_then(|s| *s).ok_or_else(|| {
SessionError::Internal("If node is missing its required 'cond' input".to_string())
})?;
let cond_t = self.value_tensor(cond_vid, resolved)?;
let cond = tensor_scalar_bool(&cond_t).ok_or_else(|| SessionError::Internal(format!(
"If: 'cond' must be a BOOL scalar, got dtype {:?} shape {:?}",
cond_t.dtype, cond_t.shape
)))?;
let attr_key = if cond { "then_branch" } else { "else_branch" };
let prepared = self.prepare_subgraph(node.id, attr_key, resolved, outer_scope)?;
let outs = self.run_subgraph(&prepared, &[])?;
if outs.len() != node.outputs.len() {
return Err(SessionError::OutputShapeCountMismatch {
op: format!("If/{attr_key}"),
expected: node.outputs.len(),
got: outs.len(),
});
}
for (vid, t) in node.outputs.iter().zip(outs.iter()) {
self.store_output_tensor(*vid, t, resolved)?;
}
Ok(())
}
fn exec_loop(
&mut self,
node: &Node,
resolved: &mut HashMap<ValueId, Vec<usize>>,
outer_scope: &HashMap<String, Tensor>,
) -> Result<()> {
let m: Option<i64> = match node.inputs.first().and_then(|s| *s) {
Some(vid) => {
let t = self.value_tensor(vid, resolved)?;
let m = tensor_scalar_i64(&t).ok_or_else(|| SessionError::Internal(format!(
"Loop: trip-count 'M' must be an INT64/INT32 scalar, got dtype {:?}",
t.dtype
)))?;
Some(m)
}
None => None,
};
let mut cond: Option<bool> = match node.inputs.get(1).and_then(|s| *s) {
Some(vid) => {
let t = self.value_tensor(vid, resolved)?;
Some(tensor_scalar_bool(&t).ok_or_else(|| SessionError::Internal(format!(
"Loop: 'cond' must be a BOOL scalar, got dtype {:?}",
t.dtype
)))?)
}
None => None,
};
let mut carried: Vec<Tensor> = Vec::new();
for slot in node.inputs.iter().skip(2) {
let vid = slot.ok_or_else(|| SessionError::Internal(
"Loop: an interior loop-carried input is omitted (empty), which ONNX does not \
allow — every v_initial must be provided".to_string(),
))?;
carried.push(self.value_tensor(vid, resolved)?);
}
let num_carried = carried.len();
let num_outputs = node.outputs.len();
if num_outputs < num_carried {
return Err(SessionError::Internal(format!(
"Loop: node declares {num_outputs} output(s) but has {num_carried} loop-carried \
dependency(ies); outputs must be carried-finals followed by scan-outputs"
)));
}
let num_scan = num_outputs - num_carried;
let expected_iterations = m.and_then(|n| usize::try_from(n).ok());
let mut scan_acc: Vec<TensorStackAccumulator> = (0..num_scan)
.map(|_| TensorStackAccumulator::new(expected_iterations))
.collect();
let prepared = self.prepare_subgraph(node.id, "body", resolved, outer_scope)?;
let mut iter_tensor = scalar_i64_tensor(0)?;
let mut cond_tensor = scalar_bool_tensor(cond.unwrap_or(true))?;
let mut iter: i64 = 0;
loop {
if let Some(m) = m
&& iter >= m
{
break;
}
if cond == Some(false) {
break;
}
iter_tensor.overwrite_bytes(&iter.to_le_bytes())?;
cond_tensor.overwrite_bytes(&[u8::from(cond.unwrap_or(true))])?;
let mut formal: Vec<&Tensor> = Vec::with_capacity(2 + num_carried);
formal.push(&iter_tensor);
formal.push(&cond_tensor);
formal.extend(carried.iter());
let outs = self.run_subgraph(&prepared, &formal)?;
drop(formal);
let expected = 1 + num_carried + num_scan;
if outs.len() != expected {
return Err(SessionError::OutputShapeCountMismatch {
op: "Loop/body".to_string(),
expected,
got: outs.len(),
});
}
let mut it = outs.into_iter();
let cond_out = it.next().expect("cond_out present");
cond = Some(tensor_scalar_bool(&cond_out).ok_or_else(|| SessionError::Internal(
format!(
"Loop: body's first output 'cond_out' must be a BOOL scalar, got dtype {:?}",
cond_out.dtype
),
))?);
carried.clear();
carried.extend((&mut it).take(num_carried));
for acc in scan_acc.iter_mut() {
acc.push(it.next().expect("scan output present"))?;
}
iter += 1;
}
for (i, t) in carried.iter().enumerate() {
self.store_output_tensor(node.outputs[i], t, resolved)?;
}
for (s, acc) in scan_acc.into_iter().enumerate() {
let (dtype, shape, bytes) = acc.finish();
self.store_output_bytes(
node.outputs[num_carried + s],
dtype,
shape,
&bytes,
resolved,
)?;
}
Ok(())
}
fn exec_scan(
&mut self,
node: &Node,
resolved: &mut HashMap<ValueId, Vec<usize>>,
outer_scope: &HashMap<String, Tensor>,
) -> Result<()> {
let num_scan_inputs = node
.attr("num_scan_inputs")
.and_then(|a| a.as_int())
.ok_or_else(|| SessionError::Internal(
"Scan: required attribute 'num_scan_inputs' is missing or not an INT".to_string(),
))? as usize;
for attr in ["scan_input_axes", "scan_output_axes"] {
if let Some(a) = node.attr(attr)
&& let Some(axes) = a.as_ints()
&& axes.iter().any(|&ax| ax != 0)
{
return Err(SessionError::Internal(format!(
"Scan: attribute '{attr}' = {axes:?} requests a non-zero scan axis, \
which this runtime does not yet support. Expected axis 0 for every \
scan input/output; re-export with axis 0 or wait for full Scan-axis \
support"
)));
}
}
for attr in ["scan_input_directions", "scan_output_directions"] {
if let Some(a) = node.attr(attr)
&& let Some(dirs) = a.as_ints()
&& dirs.iter().any(|&d| d != 0)
{
return Err(SessionError::Internal(format!(
"Scan: attribute '{attr}' = {dirs:?} requests reverse iteration, which \
this runtime does not yet support (forward only). Re-export forward or \
wait for reverse-Scan support"
)));
}
}
let total_inputs = node.inputs.len();
if total_inputs < num_scan_inputs {
return Err(SessionError::Internal(format!(
"Scan: node has {total_inputs} input(s) but num_scan_inputs={num_scan_inputs}"
)));
}
let num_state = total_inputs - num_scan_inputs;
let mut state: Vec<Tensor> = Vec::with_capacity(num_state);
for slot in node.inputs.iter().take(num_state) {
let vid = slot.ok_or_else(|| SessionError::Internal(
"Scan: an initial-state input is omitted (empty), which ONNX does not allow"
.to_string(),
))?;
state.push(self.value_tensor(vid, resolved)?);
}
let mut scan_inputs: Vec<Tensor> = Vec::with_capacity(num_scan_inputs);
for slot in node.inputs.iter().skip(num_state) {
let vid = slot.ok_or_else(|| SessionError::Internal(
"Scan: a scan input is omitted (empty), which ONNX does not allow".to_string(),
))?;
scan_inputs.push(self.value_tensor(vid, resolved)?);
}
let seq_len = scan_inputs
.first()
.and_then(|t| t.shape.first().copied())
.ok_or_else(|| SessionError::Internal(
"Scan: requires at least one scan input with rank >= 1".to_string(),
))?;
for (i, t) in scan_inputs.iter().enumerate() {
let this = t.shape.first().copied().unwrap_or(0);
if this != seq_len {
return Err(SessionError::Internal(format!(
"Scan: scan input #{i} has scan-axis length {this} but the first scan input has \
{seq_len}; all scan inputs must share the same scan-axis length"
)));
}
}
let num_outputs = node.outputs.len();
if num_outputs < num_state {
return Err(SessionError::Internal(format!(
"Scan: declares {num_outputs} output(s) but has {num_state} state variable(s); \
outputs must be final-state followed by scan-outputs"
)));
}
let num_scan_out = num_outputs - num_state;
let mut scan_acc: Vec<TensorStackAccumulator> = (0..num_scan_out)
.map(|_| TensorStackAccumulator::new(Some(seq_len)))
.collect();
let prepared = self.prepare_subgraph(node.id, "body", resolved, outer_scope)?;
let mut scan_slices = Vec::with_capacity(num_scan_inputs);
if seq_len != 0 {
for t in &scan_inputs {
let (shape, bytes) = leading_slice(t, 0)?;
scan_slices.push(Tensor::from_raw_in(
self.ep.clone(),
t.dtype,
shape,
bytes,
)?);
}
}
for step in 0..seq_len {
if step != 0 {
for (source, slice) in scan_inputs.iter().zip(scan_slices.iter_mut()) {
let (_, bytes) = leading_slice(source, step)?;
slice.overwrite_bytes(bytes)?;
}
}
let mut formal: Vec<&Tensor> = Vec::with_capacity(num_state + num_scan_inputs);
formal.extend(state.iter());
formal.extend(scan_slices.iter());
let outs = self.run_subgraph(&prepared, &formal)?;
drop(formal);
let expected = num_state + num_scan_out;
if outs.len() != expected {
return Err(SessionError::OutputShapeCountMismatch {
op: "Scan/body".to_string(),
expected,
got: outs.len(),
});
}
let mut it = outs.into_iter();
state.clear();
state.extend((&mut it).take(num_state));
for acc in scan_acc.iter_mut() {
acc.push(it.next().expect("scan output present"))?;
}
}
for (i, t) in state.iter().enumerate() {
self.store_output_tensor(node.outputs[i], t, resolved)?;
}
for (s, acc) in scan_acc.into_iter().enumerate() {
let (dtype, shape, bytes) = acc.finish();
self.store_output_bytes(node.outputs[num_state + s], dtype, shape, &bytes, resolved)?;
}
Ok(())
}
}
fn leading_slice(t: &Tensor, index: usize) -> Result<(Vec<usize>, &[u8])> {
if t.shape.is_empty() {
return Err(SessionError::Internal(
"Scan: cannot slice a scalar scan input along axis 0".to_string(),
));
}
let outer = t.shape[0];
if index >= outer {
return Err(SessionError::Internal(format!(
"Scan: slice index {index} out of range for scan-axis length {outer}"
)));
}
let inner_shape = t.shape[1..].to_vec();
let inner_numel: usize = inner_shape.iter().product();
let esize = t.dtype.byte_size();
if esize == 0 {
return Err(SessionError::Internal(format!(
"Scan: sub-byte dtype {:?} scan inputs are not supported",
t.dtype
)));
}
let slice_bytes = inner_numel * esize;
let start = index * slice_bytes;
let bytes = &t.as_bytes()[start..start + slice_bytes];
Ok((inner_shape, bytes))
}
struct TensorStackAccumulator {
expected_len: Option<usize>,
dtype: Option<DataType>,
elem_shape: Vec<usize>,
len: usize,
bytes: Vec<u8>,
}
impl TensorStackAccumulator {
fn new(expected_len: Option<usize>) -> Self {
Self {
expected_len,
dtype: None,
elem_shape: Vec::new(),
len: 0,
bytes: Vec::new(),
}
}
fn push(&mut self, tensor: Tensor) -> Result<()> {
if let Some(dtype) = self.dtype {
if tensor.shape != self.elem_shape || tensor.dtype != dtype {
return Err(SessionError::Internal(format!(
"Loop/Scan: scan output slice {} has shape {:?} dtype {:?} but the first slice \
is shape {:?} dtype {:?}; every iteration's scan output must match",
self.len, tensor.shape, tensor.dtype, self.elem_shape, dtype
)));
}
} else {
if tensor.dtype.byte_size() == 0 {
return Err(SessionError::Internal(format!(
"Loop/Scan: sub-byte dtype {:?} scan outputs are not supported",
tensor.dtype
)));
}
self.dtype = Some(tensor.dtype);
self.elem_shape = tensor.shape.clone();
if let Some(expected) = self.expected_len {
self.bytes.reserve(expected.saturating_mul(tensor.as_bytes().len()));
}
}
self.bytes.extend_from_slice(tensor.as_bytes());
self.len += 1;
Ok(())
}
fn finish(self) -> (DataType, Vec<usize>, Vec<u8>) {
if self.len == 0 {
return (DataType::Float32, vec![0], Vec::new());
}
let dtype = self.dtype.expect("non-empty accumulator has dtype");
let mut shape = Vec::with_capacity(1 + self.elem_shape.len());
shape.push(self.len);
shape.extend(self.elem_shape);
(dtype, shape, self.bytes)
}
}
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);
graph.opset_imports.insert(String::new(), 0);
assert_eq!(effective_opset(&graph, &node), 0);
}
#[test]
#[should_panic(expected = "internal invariant violated")]
fn effective_opset_requires_validated_import() {
effective_opset(
&Graph::default(),
&Node::new(NodeId(0), "Softmax", vec![], vec![]),
);
}
use onnx_runtime_ir::{static_shape, WeightRef};
use std::path::PathBuf;
fn weightstream_tmp_dir() -> PathBuf {
let dir = PathBuf::from(concat!(
env!("CARGO_MANIFEST_DIR"),
"/../../target/weightstream_test"
));
std::fs::create_dir_all(&dir).expect("create weight-streaming test dir");
dir
}
fn f32_le(data: &[f32]) -> Vec<u8> {
data.iter().flat_map(|v| v.to_le_bytes()).collect()
}
#[test]
fn aligned_external_initializer_is_borrowed_zero_copy() {
let align = TensorLayout::contiguous().alignment;
let path = weightstream_tmp_dir().join("aligned_init.bin");
let w_data = [1.0f32, 2.0, 3.0, 4.0];
std::fs::write(&path, f32_le(&w_data)).unwrap();
let mut store = WeightStore::new();
store.map_external(&path).unwrap();
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let w = g.create_named_value("W", DataType::Float32, static_shape([4]));
g.set_initializer(
w,
WeightRef::External {
path: path.clone(),
offset: 0, length: 16,
dtype: DataType::Float32,
dims: vec![4],
},
);
let y = g.create_value(DataType::Float32, static_shape([4]));
g.insert_node(Node::new(NodeId(0), "Relu", vec![Some(w)], vec![y]));
g.add_output(y);
let ep = auto_detect_cpu_ep().unwrap();
let exec = Executor::build(g, Arc::new(store), ep).unwrap();
let weight = &exec.graph.initializers[&w];
let src = exec.weights().bytes(weight).unwrap();
assert!(
(src.as_ptr() as usize).is_multiple_of(align),
"mmap window must be aligned for this test to exercise the zero-copy path"
);
let buf = &exec.buffers[&w];
assert!(buf.is_borrowed(), "aligned initializer must be borrowed, not copied");
assert_eq!(
buf.as_ptr() as *const u8,
src.as_ptr(),
"zero-copy: the buffer must alias the mmap bytes (no copy)"
);
let _ = std::fs::remove_file(&path);
}
#[test]
fn unaligned_external_initializer_falls_back_to_owned_copy() {
let align = TensorLayout::contiguous().alignment;
let path = weightstream_tmp_dir().join("unaligned_init.bin");
let offset = 8usize;
let w_data = [5.0f32, 6.0, 7.0, 8.0];
let mut file = vec![0u8; offset];
file.extend_from_slice(&f32_le(&w_data));
std::fs::write(&path, &file).unwrap();
let mut store = WeightStore::new();
store.map_external(&path).unwrap();
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let w = g.create_named_value("W", DataType::Float32, static_shape([4]));
g.set_initializer(
w,
WeightRef::External {
path: path.clone(),
offset,
length: 16,
dtype: DataType::Float32,
dims: vec![4],
},
);
let x = g.create_named_value("X", DataType::Float32, static_shape([4]));
g.add_input(x);
let y = g.create_value(DataType::Float32, static_shape([4]));
g.insert_node(Node::new(NodeId(0), "Add", vec![Some(x), Some(w)], vec![y]));
g.add_output(y);
let ep = auto_detect_cpu_ep().unwrap();
let mut exec = Executor::build(g, Arc::new(store), ep).unwrap();
let weight = &exec.graph.initializers[&w];
let src = exec.weights().bytes(weight).unwrap();
assert!(
!(src.as_ptr() as usize).is_multiple_of(align),
"window must be unaligned for this test to exercise the fallback"
);
let buf = &exec.buffers[&w];
assert!(
!buf.is_borrowed(),
"unaligned initializer must fall back to an owned copy"
);
assert_ne!(
buf.as_ptr() as *const u8,
src.as_ptr(),
"fallback: the buffer must be a fresh copy, not an alias"
);
let x_tensor = Tensor::from_f32(&[4], &[10.0, 20.0, 30.0, 40.0]).unwrap();
let out = exec.run(&[("X", &x_tensor)]).unwrap();
assert_eq!(out.len(), 1);
let got = out[0].to_vec_f32();
let want = [15.0f32, 26.0, 37.0, 48.0];
assert_eq!(got.len(), want.len());
for (g, w) in got.iter().zip(want.iter()) {
assert!((g - w).abs() < 1e-5, "got {g}, want {w}");
}
let _ = std::fs::remove_file(&path);
}
#[test]
fn producer_backed_initializer_is_not_borrowed() {
let align = TensorLayout::contiguous().alignment;
let path = weightstream_tmp_dir().join("producer_backed_init.bin");
let w_data = [1.0f32, 2.0, 3.0, 4.0];
std::fs::write(&path, f32_le(&w_data)).unwrap();
let mut store = WeightStore::new();
store.map_external(&path).unwrap();
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let x = g.create_named_value("X", DataType::Float32, static_shape([4]));
g.add_input(x);
let w = g.create_named_value("W", DataType::Float32, static_shape([4]));
g.set_initializer(
w,
WeightRef::External {
path: path.clone(),
offset: 0, length: 16,
dtype: DataType::Float32,
dims: vec![4],
},
);
g.insert_node(Node::new(NodeId(0), "Identity", vec![Some(x)], vec![w]));
let y = g.create_value(DataType::Float32, static_shape([4]));
g.insert_node(Node::new(NodeId(1), "Add", vec![Some(x), Some(w)], vec![y]));
g.add_output(y);
assert!(
g.value(w).producer.is_some(),
"test setup: initializer value must have a producer",
);
let ep = auto_detect_cpu_ep().unwrap();
let exec = Executor::build(g, Arc::new(store), ep).unwrap();
let weight = &exec.graph.initializers[&w];
let src = exec.weights().bytes(weight).unwrap();
assert!(
(src.as_ptr() as usize).is_multiple_of(align),
"mmap window must be aligned so only the producer guard prevents borrowing",
);
let buf = &exec.buffers[&w];
assert!(
!buf.is_borrowed(),
"producer-backed initializer must fall back to an owned writable copy",
);
assert_ne!(
buf.as_ptr() as *const u8,
src.as_ptr(),
"producer-backed initializer must not alias read-only mmap bytes",
);
let _ = std::fs::remove_file(&path);
}
}