diff --git a/Cargo.lock b/Cargo.lock index 6f74ba344c..83f1e686f6 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2524,6 +2524,7 @@ dependencies = [ "onnx-runtime-ir", "onnx-runtime-memory-governor", "onnx-runtime-ort-testkit", + "onnx-runtime-shape-inference", "windows-sys 0.61.2", ] diff --git a/crates/onnx-runtime-ep-cpu-plugin/Cargo.toml b/crates/onnx-runtime-ep-cpu-plugin/Cargo.toml index e3f9fc5825..3f72770c1d 100644 --- a/crates/onnx-runtime-ep-cpu-plugin/Cargo.toml +++ b/crates/onnx-runtime-ep-cpu-plugin/Cargo.toml @@ -38,7 +38,7 @@ onnx-runtime-shape-inference = { workspace = true } onnx-runtime-ort-testkit = { workspace = true } libloading = { workspace = true } onnx-genai-ort-sys = { workspace = true } -onnx-runtime-ep-plugin = { path = "../onnx-runtime-ep-plugin", version = "=0.1.0-dev.5" } +onnx-runtime-ep-plugin = { path = "../onnx-runtime-ep-plugin", version = "=0.1.0-dev.5", features = ["testutil"] } onnx-runtime-ep-cpu = { workspace = true } onnx-runtime-ep-api = { workspace = true } onnx-runtime-ir = { workspace = true } diff --git a/crates/onnx-runtime-ep-cpu-plugin/tests/shape_tables_agree.rs b/crates/onnx-runtime-ep-cpu-plugin/tests/shape_tables_agree.rs index 8e46402abb..87b762f564 100644 --- a/crates/onnx-runtime-ep-cpu-plugin/tests/shape_tables_agree.rs +++ b/crates/onnx-runtime-ep-cpu-plugin/tests/shape_tables_agree.rs @@ -90,9 +90,10 @@ fn typed_with_values(dims: &[i64], values: &[i64]) -> NodeIo { io } -fn node(op: &str, n_inputs: usize, attrs: &[(&str, i64)]) -> Node { +fn node(op: &str, n_inputs: usize, attrs: &[(&str, i64)], opset: u64) -> Node { let inputs: Vec> = (0..n_inputs).map(|i| Some(ValueId(i as u32))).collect(); let mut n = Node::new(NodeId(0), op, inputs, vec![ValueId(100)]); + n.version = Some(opset as i64); for (k, v) in attrs { n.attributes.insert((*k).to_string(), Attribute::Int(*v)); } @@ -146,9 +147,27 @@ fn assert_agree( whichever path a user exercised would look correct.", plugin[0] ); + + let input_shapes: Vec>> = plugin_inputs + .iter() + .map(|input| input.shape.iter().copied().map(Some).collect()) + .collect(); + let strategy = ShapeInference::for_node(node, &input_shapes, node.outputs.len()); + let ShapeInference::SharedNative { fallback, .. } = &strategy else { + panic!( + "{what}: the production registry stopped routing this op through the shared adapter" + ); + }; + assert!( + !matches!(fallback.as_ref(), ShapeInference::SharedNative { .. }), + "{what}: a shared rule's fallback must not recurse into the shared adapter" + ); + let production = + onnx_runtime_ep_plugin::compute::infer_shapes_for_test(&strategy, plugin_inputs) + .unwrap_or_else(|e| panic!("{what}: production shared rule failed: {e}")); + assert_eq!(production, vec![native_dims]); } -#[test] fn tile_agrees() { let buf = [0u8; 24]; let data = view(DataType::Float32, &[2, 3], &[3, 1], buf.as_ptr()); @@ -156,7 +175,7 @@ fn tile_agrees() { let r = view(DataType::Int64, &[2], &[1], reps.as_ptr().cast()); assert_agree( "Tile", - &node("Tile", 2, &[]), + &node("Tile", 2, &[], 13), 13, ShapeInference::Tile, &[data, r], @@ -167,7 +186,6 @@ fn tile_agrees() { ); } -#[test] fn expand_agrees_on_bidirectional_broadcast() { // The case a "just take the target" implementation gets wrong. If the two // tables ever diverge, this is where it shows. @@ -177,7 +195,7 @@ fn expand_agrees_on_bidirectional_broadcast() { let s = view(DataType::Int64, &[2], &[1], want.as_ptr().cast()); assert_agree( "Expand", - &node("Expand", 2, &[]), + &node("Expand", 2, &[], 13), 13, ShapeInference::Expand, &[data, s], @@ -188,13 +206,12 @@ fn expand_agrees_on_bidirectional_broadcast() { ); } -#[test] fn constant_of_shape_agrees() { let dims: [i64; 3] = [2, 3, 4]; let t = view(DataType::Int64, &[3], &[1], dims.as_ptr().cast()); assert_agree( "ConstantOfShape", - &node("ConstantOfShape", 1, &[]), + &node("ConstantOfShape", 1, &[], 9), 9, ShapeInference::ConstantOfShape, &[t], @@ -202,19 +219,91 @@ fn constant_of_shape_agrees() { ); } +#[test] +fn migrated_shared_rules_agree() { + const EXPECTED_RULES: &[&str] = &["ConstantOfShape", "Expand", "Tile"]; + let rules = onnx_runtime_ep_plugin::compute::shared_native_rule_names_for_test(); + assert_eq!( + rules, EXPECTED_RULES, + "the production shared-rule census changed; update the explicit agreement fixtures in the same commit" + ); + let mut compared = 0; + for rule in rules { + match rule { + "ConstantOfShape" => constant_of_shape_agrees(), + "Expand" => expand_agrees_on_bidirectional_broadcast(), + "Tile" => tile_agrees(), + other => panic!("shared rule {other} has no agreement fixture"), + } + compared += 1; + } + assert_eq!(compared, EXPECTED_RULES.len()); +} + +#[test] +fn migrated_shared_rules_agree_on_edge_extents() { + let empty: [i64; 0] = []; + let empty_shape = view(DataType::Int64, &[0], &[1], empty.as_ptr().cast()); + assert_agree( + "ConstantOfShape empty shape", + &node("ConstantOfShape", 1, &[], 9), + 9, + ShapeInference::ConstantOfShape, + &[empty_shape], + vec![typed_with_values(&[0], &empty)], + ); + + let scalar = [0.0f32; 1]; + let data = view(DataType::Float32, &[1], &[1], scalar.as_ptr().cast()); + let zero_target = [0i64]; + let target = view(DataType::Int64, &[1], &[1], zero_target.as_ptr().cast()); + assert_agree( + "Expand zero target extent", + &node("Expand", 2, &[], 8), + 8, + ShapeInference::Expand, + &[data, target], + vec![ + typed(DataType::Float32, &[1]), + typed_with_values(&[1], &zero_target), + ], + ); + + let tile_data = [0.0f32; 6]; + let data = view( + DataType::Float32, + &[2, 3], + &[3, 1], + tile_data.as_ptr().cast(), + ); + let zero_repeat = [0i64, 2]; + let repeats = view(DataType::Int64, &[2], &[1], zero_repeat.as_ptr().cast()); + assert_agree( + "Tile zero repeat", + &node("Tile", 2, &[], 13), + 13, + ShapeInference::Tile, + &[data, repeats], + vec![ + typed(DataType::Float32, &[2, 3]), + typed_with_values(&[2], &zero_repeat), + ], + ); +} + // ── Sweep ──────────────────────────────────────────────────────────────────── -/// Every plain shape-preserving op must agree, driven from the registry rather -/// than a hand-written list. +/// Every plain shape-preserving op added by #2049 must agree. /// /// The three cases above are hand-built because they need *values*. These do /// not: one input, output shape equals it. Sweeping them costs nothing and /// catches the case where one table quietly stops preserving the shape — a /// `DequantizeLinear` that started broadcasting against its scale, say. /// -/// A hand-written list would drift from the rules themselves, which is the very -/// failure this file exists to catch, so the op names come from the same place -/// the plugin's arm does. +/// Unlike the migrated cases above, these not-yet-shared rules have no +/// production descriptor to enumerate. Their explicit list is temporary +/// duplication and should disappear as later slices move them into +/// [`SharedNativeShapeRule`]. #[test] fn shape_preserving_ops_agree() { // (op, opset, extra input count) — the companions are dtype/scale operands @@ -229,6 +318,7 @@ fn shape_preserving_ops_agree() { ]; let buf = [0u8; 64]; + let mut compared = 0; for &(op, opset, extras) in CASES { let dims: [usize; 2] = [3, 4]; let strides: [i64; 2] = [4, 1]; @@ -253,19 +343,16 @@ fn shape_preserving_ops_agree() { )); native_inputs.push(typed(DataType::Float32, &[3])); } - let n = node(op, 1 + extras, &[]); - let nat = match native_try(&n, native_inputs, opset) { - Ok(outs) => outs, - // The native table rejected the node outright. That is a validity - // judgement, not a shape answer, so there is nothing to compare — - // and it is recorded rather than swallowed so a reader knows the - // two tables differ in strictness, not in shape. - Err(_) => continue, - }; - let Some(native_dims) = native_static(&nat) else { - // The native table declined to resolve this one; nothing to compare. - continue; + let attrs: &[(&str, i64)] = if matches!(op, "DequantizeLinear" | "QuantizeLinear") { + &[("axis", 0)] + } else { + &[] }; + let n = node(op, 1 + extras, attrs, opset); + let nat = native_try(&n, native_inputs, opset) + .unwrap_or_else(|error| panic!("{op}: expected native inference to resolve: {error}")); + let native_dims = + native_static(&nat).unwrap_or_else(|| panic!("{op}: expected a concrete native shape")); let plugin = onnx_runtime_ep_plugin::compute::infer_shapes_for_test( &ShapeInference::SameAsInput(0), &plugin_inputs, @@ -277,5 +364,46 @@ fn shape_preserving_ops_agree() { {native_dims:?}.", plugin[0] ); + compared += 1; } + assert_eq!( + compared, + CASES.len(), + "every expected shape-preserving fixture must reach the comparison" + ); +} + +#[test] +fn malformed_dequantize_preserves_the_plugin_permissiveness_boundary() { + let data_buf = [0.0f32; 12]; + let scale_buf = [1.0f32; 3]; + let data = view( + DataType::Float32, + &[3, 4], + &[4, 1], + data_buf.as_ptr().cast(), + ); + let scale = view(DataType::Float32, &[3], &[1], scale_buf.as_ptr().cast()); + let n = node("DequantizeLinear", 2, &[], 13); + + let native_error = native_try( + &n, + vec![ + typed(DataType::Float32, &[3, 4]), + typed(DataType::Float32, &[3]), + ], + 13, + ) + .expect_err("native inference validates the per-axis scale length"); + assert!(native_error.to_string().contains("scale length 3")); + let strategy = ShapeInference::for_node(&n, &[vec![Some(3), Some(4)], vec![Some(3)]], 1); + assert!( + matches!(strategy, ShapeInference::SameAsInput(0)), + "DequantizeLinear is intentionally outside the first shared slice" + ); + assert_eq!( + onnx_runtime_ep_plugin::compute::infer_shapes_for_test(&strategy, &[data, scale]).unwrap(), + vec![vec![3, 4]], + "the plugin's pre-existing permissive sizing contract must remain unchanged" + ); } diff --git a/crates/onnx-runtime-ep-plugin/Cargo.toml b/crates/onnx-runtime-ep-plugin/Cargo.toml index 9ec454cc7d..0fd7a14f3b 100644 --- a/crates/onnx-runtime-ep-plugin/Cargo.toml +++ b/crates/onnx-runtime-ep-plugin/Cargo.toml @@ -14,10 +14,13 @@ categories = ["science", "hardware-support"] # calls are empty, so production builds carry none of it. Timing is gated a # second time at runtime by ONNX_GENAI_PROFILE_DISPATCH=1. dispatch_probe = [] +# Cross-crate shape-table agreement hooks. Shipped artifacts do not enable this. +testutil = [] [dependencies] onnx-runtime-ep-api = { workspace = true } onnx-runtime-ir = { workspace = true } +onnx-runtime-shape-inference = { workspace = true } onnx-genai-ort-sys = { workspace = true } # `pin.rs` takes one never-released reference to this library's own module so diff --git a/crates/onnx-runtime-ep-plugin/src/compute.rs b/crates/onnx-runtime-ep-plugin/src/compute.rs index e03578b27f..07ca61ca28 100644 --- a/crates/onnx-runtime-ep-plugin/src/compute.rs +++ b/crates/onnx-runtime-ep-plugin/src/compute.rs @@ -8,12 +8,10 @@ //! an error naming the op and domain. This surfaces at Compute time — never //! silently producing a wrong-shape tensor. //! -//! [`ShapeInference::for_node`] is the single source of truth: it takes the -//! compiled IR `Node`, its input shapes and its output count, and resolves the -//! rule — reading attributes where the shape depends on them (Reshape, Conv, -//! reductions, the LayerNorm family, …). There is deliberately no op-name-only -//! entry point; every caller has a `Node` at capability and compile time, so a -//! second, attribute-blind table would only be a place for rules to drift. +//! [`ShapeInference::for_node`] is the single dispatch point: it takes the +//! compiled IR `Node`, its input shapes and its output count, and selects either +//! the shared native registry adapter or a plugin-only rule. There is +//! deliberately no attribute-blind op-name entry point. use std::collections::HashSet; use std::ffi::c_void; @@ -29,6 +27,7 @@ use onnx_runtime_ep_api::tensor::{DevicePtr, DevicePtrMut, TensorMut, TensorView use onnx_runtime_ir::{DataType, DeviceId, Node}; use crate::kernel_ctx::{allocate_output, read_inputs_into}; +use crate::shared_shapes::{SharedNativeShapeRule, SharedShapeResult, infer_shared_node}; use crate::status::{fail_status, ok_status}; // ────────────────────────────────────────────────────────────────────────────── @@ -52,6 +51,13 @@ pub struct ConvSpatialAxis { /// How to infer output shapes at runtime from the concrete input shapes. #[derive(Clone, Debug)] pub enum ShapeInference { + /// Ask the native symbolic registry first, then preserve the existing + /// plugin rule as a fallback when values remain unavailable or the native + /// rule is stricter than the historical plugin contract. + SharedNative { + node: Box, + fallback: Box, + }, /// numpy-style broadcast of all inputs → one output. ElementwiseBroadcast, /// Output shape == input[idx].shape. @@ -276,6 +282,18 @@ impl ShapeInference { let domain = node.domain.as_str(); let opset = node.version.unwrap_or(0); + if let Some(rule) = SharedNativeShapeRule::for_node(node) { + let fallback = match rule { + SharedNativeShapeRule::ConstantOfShape => Self::ConstantOfShape, + SharedNativeShapeRule::Expand => Self::Expand, + SharedNativeShapeRule::Tile => Self::Tile, + }; + return Self::SharedNative { + node: Box::new(node.clone()), + fallback: Box::new(fallback), + }; + } + let int_attr = |name: &str| -> Option { node.attr(name)?.as_int() }; let ints_attr = |name: &str| -> Option> { Some(node.attr(name)?.as_ints()?.to_vec()) }; @@ -592,13 +610,9 @@ impl ShapeInference { | "QuantizeLinear" => Self::SameAsInput(0), // ── Shapes carried in input values ──────────────────────────── - // Each of these is data-dependent, and each is *cheap*: the extent - // is a handful of int64s, not a computation over the payload. That - // is what separates them from `Unique` / `NonMaxSuppression`, where - // the extent is the whole algorithm. - "ConstantOfShape" => Self::ConstantOfShape, - "Expand" => Self::Expand, - "Tile" => Self::Tile, + // The shared native adapter owns ConstantOfShape/Expand/Tile above; + // their local variants remain as compatibility fallbacks. Window + // sizing is the next cheap value-carried rule still local here. "HannWindow" | "HammingWindow" | "BlackmanWindow" => Self::Window, // ── DFT ─────────────────────────────────────────────────────── @@ -1947,6 +1961,7 @@ unsafe fn mem_info_is_device(api: &ort::OrtApi, mem_info: *const ort::OrtMemoryI /// operand that crossed a CPU→device boundary and must be staged (#982). fn host_operand_indices(strategy: &ShapeInference) -> &'static [usize] { match strategy { + ShapeInference::SharedNative { fallback, .. } => host_operand_indices(fallback), ShapeInference::ReshapeData { .. } => &[1], ShapeInference::SliceData => &[1, 2, 3, 4], ShapeInference::ReductionFromInput { .. } => &[1], @@ -3650,7 +3665,7 @@ fn read_i64_vec(t: &TensorView<'_>, what: &str) -> Result, String> { if len == 0 { return Ok(Vec::new()); } - let base = t.data.as_ptr() as *const u8; + let base = t.data.as_ptr::(); if base.is_null() { return Err(format!("{what} has a null data pointer")); } @@ -3688,7 +3703,7 @@ fn read_scalar_i64(t: &TensorView<'_>, what: &str) -> Result { t.device )); } - let base = t.data.as_ptr() as *const u8; + let base = t.data.as_ptr::(); if base.is_null() { return Err(format!("{what} has a null data pointer")); } @@ -3746,7 +3761,7 @@ fn count_true(condition: &TensorView<'_>) -> Result, String> { let len = condition.shape.first().copied().unwrap_or(0); let stride = condition.strides.first().copied().unwrap_or(1); - let base = condition.data.as_ptr() as *const u8; + let base = condition.data.as_ptr::(); if base.is_null() { return Err("Compress: condition has a null data pointer".into()); } @@ -3767,12 +3782,12 @@ fn count_true(condition: &TensorView<'_>) -> Result, String> { /// Test-only entry point to [`infer_shapes`]. /// -/// Exposed so `shape_tables_agree` in `onnx-runtime-ep-cpu-plugin` can drive -/// this table and the native `onnx-runtime-shape-inference` one over the same -/// node and compare them. Two independent encodings of one ONNX specification -/// drift, and the drift is silent — the two paths simply disagree about the same -/// graph. Keeping the function itself private preserves the invariant that -/// production callers reach it only through the Compute path. +/// Exposed so `shape_tables_agree` in `onnx-runtime-ep-cpu-plugin` can compare +/// the compatibility fallback with the native registry and exercise the +/// production shared route. The `testutil` feature is enabled only by that +/// crate's dev-dependency, so shipped plugin artifacts expose no test-only API. +#[cfg(feature = "testutil")] +#[doc(hidden)] pub fn infer_shapes_for_test( strategy: &ShapeInference, inputs: &[TensorView<'_>], @@ -3780,12 +3795,31 @@ pub fn infer_shapes_for_test( infer_shapes(strategy, inputs) } +/// Exact production census used by the cross-crate anti-vacuity test. +#[cfg(feature = "testutil")] +#[doc(hidden)] +pub fn shared_native_rule_names_for_test() -> Vec<&'static str> { + SharedNativeShapeRule::all() + .iter() + .map(|rule| rule.op_type()) + .collect() +} + /// Infer output shapes from the shape inference strategy and input views. fn infer_shapes( strategy: &ShapeInference, inputs: &[TensorView<'_>], ) -> Result>, String> { match strategy { + ShapeInference::SharedNative { node, fallback } => match infer_shared_node(node, inputs) { + SharedShapeResult::Resolved(shapes) => Ok(shapes), + // Preserve the plugin's established permissiveness. In particular, + // native validation can reject malformed companion operands that a + // shape-only plugin rule historically ignored. + SharedShapeResult::SymbolicOrUnknown | SharedShapeResult::Rejected(_) => { + infer_shapes(fallback, inputs) + } + }, ShapeInference::ElementwiseBroadcast => { if inputs.is_empty() { return Err("ElementwiseBroadcast: no inputs".into()); @@ -5166,6 +5200,150 @@ fn after() {} ); } + #[test] + fn shared_rule_rejection_preserves_the_existing_plugin_fallback() { + let buf = vec![0.0f32; 6]; + let data = f32_data(&[2, 3], &[3, 1], &buf); + let reps = [3i64, 2]; + // The native rule correctly rejects this malformed rank-2 `repeats` + // tensor. The historical plugin rule only reads its two values and + // sizes the output, so the first migration slice deliberately retains + // that permissive behavior. + let r = i64_scalar(&reps, &[2, 1], &[1, 1]); + let mut node = Node::new( + onnx_runtime_ir::NodeId(0), + "Tile", + vec![ + Some(onnx_runtime_ir::ValueId(0)), + Some(onnx_runtime_ir::ValueId(1)), + ], + vec![onnx_runtime_ir::ValueId(2)], + ); + node.version = Some(13); + assert!(matches!( + infer_shared_node(&node, &[data, r]), + SharedShapeResult::Rejected(reason) if reason.contains("invalid rank 2") + )); + + let strategy = + ShapeInference::for_node(&node, &[vec![Some(2), Some(3)], vec![Some(2), Some(1)]], 1); + assert_eq!(infer(&strategy, &[data, r]).unwrap(), vec![vec![6, 6]]); + } + + #[test] + fn shared_expand_opset_floor_falls_back_before_version_eight() { + let buf = vec![0.0f32; 3]; + let data = f32_data(&[3, 1], &[1, 1], &buf); + let want = [1i64, 4]; + let shape_in = i64_scalar(&want, &[2], &[1]); + let make_node = |version| { + let mut node = Node::new( + onnx_runtime_ir::NodeId(0), + "Expand", + vec![ + Some(onnx_runtime_ir::ValueId(0)), + Some(onnx_runtime_ir::ValueId(1)), + ], + vec![onnx_runtime_ir::ValueId(2)], + ); + node.version = Some(version); + node + }; + + let version_seven = make_node(7); + assert_eq!( + infer_shared_node(&version_seven, &[data, shape_in]), + SharedShapeResult::SymbolicOrUnknown + ); + let strategy = + ShapeInference::for_node(&version_seven, &[vec![Some(3), Some(1)], vec![Some(2)]], 1); + assert_eq!( + infer(&strategy, &[data, shape_in]).unwrap(), + vec![vec![3, 4]], + "Expand@7 must reach the compatibility fallback" + ); + + assert_eq!( + infer_shared_node(&make_node(8), &[data, shape_in]), + SharedShapeResult::Resolved(vec![vec![3, 4]]) + ); + } + + #[test] + fn foreign_domain_expand_is_not_a_shared_native_rule() { + let mut node = Node::new( + onnx_runtime_ir::NodeId(0), + "Expand", + vec![None, None], + vec![onnx_runtime_ir::ValueId(2)], + ); + node.domain = "example.foreign".into(); + node.version = Some(8); + assert!( + !matches!( + ShapeInference::for_node(&node, &[vec![Some(3), Some(1)], vec![Some(2)]], 1), + ShapeInference::SharedNative { .. } + ), + "operator names from foreign domains must not enter default-domain shared rules" + ); + } + + #[test] + fn unregistered_native_rule_can_use_a_synthetic_fallback() { + let node = Node::new( + onnx_runtime_ir::NodeId(0), + "NotRegisteredAnywhere", + vec![Some(onnx_runtime_ir::ValueId(0))], + vec![onnx_runtime_ir::ValueId(1)], + ); + let input = view(&[2, 3], &[3, 1]); + assert_eq!( + infer_shared_node(&node, &[input]), + SharedShapeResult::SymbolicOrUnknown + ); + let strategy = ShapeInference::SharedNative { + node: Box::new(node), + fallback: Box::new(ShapeInference::SameAsInput(0)), + }; + assert_eq!(infer(&strategy, &[input]).unwrap(), vec![vec![2, 3]]); + } + + #[test] + fn device_shape_operand_reaches_safe_plugin_error() { + let buf = vec![0.0f32; 3]; + let data = f32_data(&[3, 1], &[1, 1], &buf); + let want = [1i64, 4]; + let shape_in = TensorView::new( + DevicePtr(want.as_ptr().cast()), + DataType::Int64, + &[2], + &[1], + DeviceId::cuda(0), + ); + let mut node = Node::new( + onnx_runtime_ir::NodeId(0), + "Expand", + vec![ + Some(onnx_runtime_ir::ValueId(0)), + Some(onnx_runtime_ir::ValueId(1)), + ], + vec![onnx_runtime_ir::ValueId(2)], + ); + node.version = Some(8); + assert_eq!( + infer_shared_node(&node, &[data, shape_in]), + SharedShapeResult::SymbolicOrUnknown + ); + + let strategy = ShapeInference::for_node(&node, &[vec![Some(3), Some(1)], vec![Some(2)]], 1); + let error = infer(&strategy, &[data, shape_in]) + .expect_err("the fallback must refuse to dereference a device pointer"); + assert!( + error.contains("host cannot read") && error.contains("Cuda"), + "device-resident fallback error should explain the safe refusal: {error}" + ); + } + #[test] fn window_length_is_the_scalar_input_value() { let n = [16i64]; diff --git a/crates/onnx-runtime-ep-plugin/src/lib.rs b/crates/onnx-runtime-ep-plugin/src/lib.rs index aeb7aadf44..86fd768858 100644 --- a/crates/onnx-runtime-ep-plugin/src/lib.rs +++ b/crates/onnx-runtime-ep-plugin/src/lib.rs @@ -37,6 +37,7 @@ pub mod graph_reader; pub mod host_pool; pub mod kernel_ctx; pub mod pin; +mod shared_shapes; pub mod status; pub mod transfer; diff --git a/crates/onnx-runtime-ep-plugin/src/shared_shapes.rs b/crates/onnx-runtime-ep-plugin/src/shared_shapes.rs new file mode 100644 index 0000000000..d42ec0540b --- /dev/null +++ b/crates/onnx-runtime-ep-plugin/src/shared_shapes.rs @@ -0,0 +1,343 @@ +//! Adapter from concrete plugin inputs to the native symbolic shape registry. + +use std::collections::HashMap; +use std::sync::OnceLock; + +use onnx_runtime_ep_api::TensorView; +use onnx_runtime_ir::{Node, normalize_domain}; +use onnx_runtime_shape_inference::{ + DimExpr, InferenceRegistry, MAX_SHAPE_DATA_ELEMS, MergePolicy, NodeIo, ShapeData, + SymbolInterner, TypeInfo, +}; + +/// The deliberately small first set of plugin rules delegated to the native +/// shape registry. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum SharedNativeShapeRule { + ConstantOfShape, + Expand, + Tile, +} + +impl SharedNativeShapeRule { + const ALL: [Self; 3] = [Self::ConstantOfShape, Self::Expand, Self::Tile]; + + /// Every rule currently routed through the shared adapter. + #[cfg(feature = "testutil")] + pub(crate) const fn all() -> &'static [Self] { + &Self::ALL + } + + /// The ONNX operator name for this rule. + pub(crate) const fn op_type(self) -> &'static str { + match self { + Self::ConstantOfShape => "ConstantOfShape", + Self::Expand => "Expand", + Self::Tile => "Tile", + } + } + + pub(crate) fn for_node(node: &Node) -> Option { + if !node.is_default_domain() { + return None; + } + Self::ALL + .iter() + .copied() + .find(|rule| rule.op_type() == node.op_type) + } +} + +/// Result of asking the native shape registry to resolve one plugin node. +#[derive(Clone, Debug, PartialEq, Eq)] +pub(crate) enum SharedShapeResult { + /// Every output rank and extent was resolved concretely. + Resolved(Vec>), + /// The rule left an output unknown or symbolic. + SymbolicOrUnknown, + /// The native rule rejected malformed input. + Rejected(String), +} + +/// Infer one node through the native registry using concrete plugin metadata. +/// +/// Host-accessible rank-0/rank-1 values are supplied as [`ShapeData`], bounded +/// by [`MAX_SHAPE_DATA_ELEMS`]. Device values are deliberately not read: the +/// native rule remains symbolic and the caller can use its existing +/// Compute-time fallback. +/// +/// The plugin graph reader stores ORT's `Node_GetSinceVersion` result in +/// [`Node::version`]. That is the selected kernel schema's `since_version`, not +/// necessarily the model's graph-level opset. This adapter therefore dispatches +/// with that exact local version. If it is absent or invalid, it deliberately +/// uses opset 1 rather than guessing the latest version. Future migrations with +/// version-dependent semantics must account for this contract explicitly. +pub(crate) fn infer_shared_node(node: &Node, inputs: &[TensorView<'_>]) -> SharedShapeResult { + let input_ios = match inputs.iter().map(node_io).collect::, _>>() { + Ok(inputs) => inputs, + Err(reason) => return SharedShapeResult::Rejected(reason), + }; + let opset = node.local_opset().unwrap_or(1); + let mut imports = HashMap::new(); + imports.insert(normalize_domain(&node.domain).to_string(), opset); + let mut interner = SymbolInterner::new(0x8000_0000); + + static REGISTRY: OnceLock = OnceLock::new(); + let outputs = match REGISTRY + .get_or_init(InferenceRegistry::default_registry) + .infer_node( + node, + &imports, + input_ios, + MergePolicy::Permissive, + &mut interner, + ) { + Ok(outputs) => outputs, + Err(error) => return SharedShapeResult::Rejected(error.to_string()), + }; + + let Some(shapes) = outputs + .iter() + .map(|output| { + output + .type_info + .as_ref()? + .shape + .iter() + .map(|dim| dim.as_const().and_then(|n| usize::try_from(n).ok())) + .collect::>>() + }) + .collect::>>() + else { + return SharedShapeResult::SymbolicOrUnknown; + }; + + if shapes.is_empty() { + SharedShapeResult::SymbolicOrUnknown + } else { + SharedShapeResult::Resolved(shapes) + } +} + +fn node_io(input: &TensorView<'_>) -> Result { + if input.is_absent() { + return Ok(NodeIo::default()); + } + let shape = input + .shape + .iter() + .map(|&extent| { + i64::try_from(extent) + .map(DimExpr::constant) + .map_err(|_| format!("input extent {extent} exceeds i64::MAX")) + }) + .collect::, _>>()?; + let type_info = TypeInfo::new(input.dtype, shape); + Ok(NodeIo { + type_info: Some(type_info), + shape_data: tensor_shape_data(input), + value_type: None, + }) +} + +fn tensor_shape_data(input: &TensorView<'_>) -> Option { + if !input.device.is_host_accessible() || input.shape.len() > 1 { + return None; + } + let numel = input.shape.first().copied().unwrap_or(1); + if numel > MAX_SHAPE_DATA_ELEMS { + return None; + } + let element_size = input.dtype.byte_size(); + if element_size == 0 || (numel != 0 && input.data.is_null()) { + return None; + } + + let mut bytes = Vec::with_capacity(numel.checked_mul(element_size)?); + let stride = isize::try_from(input.strides.first().copied().unwrap_or(1)).ok()?; + let element_size_isize = isize::try_from(element_size).ok()?; + let origin = input.data.as_ptr::(); + for index in 0..numel { + let element_offset = isize::try_from(index).ok()?.checked_mul(stride)?; + let byte_offset = element_offset.checked_mul(element_size_isize)?; + // SAFETY: TensorView's contract makes each logical element address + // readable, including negative strides; only the bounded scalar/vector + // shape-data subset is copied. + let element = unsafe { origin.add(input.byte_offset).offset(byte_offset) }; + // SAFETY: `element` addresses one complete element under that contract. + bytes.extend_from_slice(unsafe { std::slice::from_raw_parts(element, element_size) }); + } + ShapeData::from_tensor(input.dtype, input.shape, &bytes) +} + +#[cfg(test)] +mod tests { + use onnx_runtime_ep_api::DevicePtr; + use onnx_runtime_ir::{DataType, DeviceId, NodeId, ValueId}; + + use super::*; + + fn node_at(op: &str, input_count: usize, version: Option) -> Node { + let mut node = Node::new( + NodeId(0), + op, + (0..input_count) + .map(|index| Some(ValueId(index as u32))) + .collect(), + vec![ValueId(100)], + ); + node.version = version; + node + } + + fn node(op: &str, input_count: usize) -> Node { + node_at(op, input_count, Some(13)) + } + + fn view<'a>( + data: *const u8, + dtype: DataType, + shape: &'a [usize], + strides: &'a [i64], + device: DeviceId, + ) -> TensorView<'a> { + TensorView::new(DevicePtr(data.cast()), dtype, shape, strides, device) + } + + #[test] + fn device_value_leaves_the_native_answer_symbolic() { + let data = [0.0f32; 3]; + let target = [1i64, 4]; + let inputs = [ + view( + data.as_ptr().cast(), + DataType::Float32, + &[3, 1], + &[1, 1], + DeviceId::cpu(), + ), + view( + target.as_ptr().cast(), + DataType::Int64, + &[2], + &[1], + DeviceId::cuda(0), + ), + ]; + + assert_eq!( + infer_shared_node(&node("Expand", 2), &inputs), + SharedShapeResult::SymbolicOrUnknown + ); + } + + #[test] + fn expand_dispatches_with_the_exact_node_version_not_latest() { + let data = [0.0f32; 3]; + let target = [1i64, 4]; + let inputs = [ + view( + data.as_ptr().cast(), + DataType::Float32, + &[3, 1], + &[1, 1], + DeviceId::cpu(), + ), + view( + target.as_ptr().cast(), + DataType::Int64, + &[2], + &[1], + DeviceId::cpu(), + ), + ]; + + assert_eq!( + infer_shared_node(&node_at("Expand", 2, Some(7)), &inputs), + SharedShapeResult::SymbolicOrUnknown, + "Expand was introduced at opset 8" + ); + assert_eq!( + infer_shared_node(&node_at("Expand", 2, Some(8)), &inputs), + SharedShapeResult::Resolved(vec![vec![3, 4]]) + ); + assert_eq!( + infer_shared_node(&node_at("Expand", 2, None), &inputs), + SharedShapeResult::SymbolicOrUnknown, + "an absent version must use opset 1, not silently select latest" + ); + } + + #[test] + fn shared_rule_selection_requires_the_default_domain() { + let mut foreign = node("Expand", 2); + foreign.domain = "example.foreign".into(); + assert_eq!(SharedNativeShapeRule::for_node(&foreign), None); + assert_eq!( + SharedNativeShapeRule::for_node(&node("Expand", 2)), + Some(SharedNativeShapeRule::Expand) + ); + } + + #[test] + fn unregistered_operator_is_symbolic_not_rejected() { + assert_eq!( + infer_shared_node(&node("NotRegisteredAnywhere", 0), &[]), + SharedShapeResult::SymbolicOrUnknown + ); + } + + #[test] + fn malformed_input_is_a_rejection_not_an_unknown_shape() { + let data = [0.0f32; 6]; + let repeats = [3i64]; + let inputs = [ + view( + data.as_ptr().cast(), + DataType::Float32, + &[2, 3], + &[3, 1], + DeviceId::cpu(), + ), + view( + repeats.as_ptr().cast(), + DataType::Int64, + &[1], + &[1], + DeviceId::cpu(), + ), + ]; + + assert!(matches!( + infer_shared_node(&node("Tile", 2), &inputs), + SharedShapeResult::Rejected(reason) if reason.contains("input rank is 2") + )); + } + + #[test] + fn strided_shape_data_is_copied_in_logical_order() { + let data = [0.0f32; 3]; + let target = [1i64, 99, 4]; + let inputs = [ + view( + data.as_ptr().cast(), + DataType::Float32, + &[3, 1], + &[1, 1], + DeviceId::cpu(), + ), + view( + target.as_ptr().cast(), + DataType::Int64, + &[2], + &[2], + DeviceId::cpu(), + ), + ]; + + assert_eq!( + infer_shared_node(&node("Expand", 2), &inputs), + SharedShapeResult::Resolved(vec![vec![3, 4]]) + ); + } +}