Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 1 addition & 1 deletion crates/onnx-runtime-ep-cpu-plugin/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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 }
Expand Down
176 changes: 152 additions & 24 deletions crates/onnx-runtime-ep-cpu-plugin/tests/shape_tables_agree.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<Option<ValueId>> = (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));
}
Expand Down Expand Up @@ -146,17 +147,35 @@ fn assert_agree(
whichever path a user exercised would look correct.",
plugin[0]
);

let input_shapes: Vec<Vec<Option<usize>>> = 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());
let reps: [i64; 2] = [3, 2];
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],
Expand All @@ -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.
Expand All @@ -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],
Expand All @@ -188,33 +206,104 @@ 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],
vec![typed_with_values(&[3], &dims)],
);
}

#[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
Expand All @@ -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];
Expand All @@ -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,
Expand All @@ -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"
);
}
3 changes: 3 additions & 0 deletions crates/onnx-runtime-ep-plugin/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading
Loading