use prost::Message;
use onnx_runtime_einsum_conformance::{
ConformanceDType, DeclaredDType, MalformedCase, malformed_cases, named_cases,
};
use onnx_runtime_ir::{Attribute, DataType, Graph, Node, NodeId, static_shape};
use onnx_runtime_loader::{LoaderError, load_model_bytes, proto::onnx, validate_model};
fn tensor_type(elem_type: i32, shape: Option<&[i64]>) -> onnx::TypeProto {
use onnx::tensor_shape_proto::{Dimension, dimension::Value};
onnx::TypeProto {
value: Some(onnx::type_proto::Value::TensorType(
onnx::type_proto::Tensor {
elem_type,
shape: shape.map(|shape| onnx::TensorShapeProto {
dim: shape
.iter()
.map(|&dimension| Dimension {
value: Some(Value::DimValue(dimension)),
..Default::default()
})
.collect(),
}),
},
)),
..Default::default()
}
}
fn value_info(name: &str, elem_type: Option<i32>, shape: Option<&[i64]>) -> onnx::ValueInfoProto {
onnx::ValueInfoProto {
name: name.to_string(),
r#type: elem_type.map(|elem_type| tensor_type(elem_type, shape)),
..Default::default()
}
}
fn einsum_node(input: &str, output: &str, equation: &str) -> onnx::NodeProto {
onnx::NodeProto {
op_type: "Einsum".to_string(),
input: vec![input.to_string()],
output: vec![output.to_string()],
attribute: vec![onnx::AttributeProto {
name: "equation".to_string(),
r#type: onnx::attribute_proto::AttributeType::String as i32,
s: equation.as_bytes().to_vec(),
..Default::default()
}],
..Default::default()
}
}
fn einsum_node_with_io(inputs: &[&str], outputs: &[&str], equation: &str) -> onnx::NodeProto {
let mut node = einsum_node("", "", equation);
node.input = inputs.iter().map(|name| (*name).to_string()).collect();
node.output = outputs.iter().map(|name| (*name).to_string()).collect();
node
}
fn model(
opset: i64,
input: onnx::ValueInfoProto,
output: onnx::ValueInfoProto,
nodes: Vec<onnx::NodeProto>,
) -> Vec<u8> {
onnx::ModelProto {
ir_version: 8,
opset_import: vec![onnx::OperatorSetIdProto {
domain: String::new(),
version: opset,
}],
graph: Some(onnx::GraphProto {
input: vec![input],
output: vec![output],
node: nodes,
..Default::default()
}),
..Default::default()
}
.encode_to_vec()
}
fn find(graph: &Graph, name: &str) -> onnx_runtime_ir::ValueId {
graph
.values
.iter()
.find_map(|(id, value)| (value.name.as_deref() == Some(name)).then_some(id))
.unwrap_or_else(|| panic!("value {name:?} was not loaded"))
}
fn einsum_graph(opset: u64, dtype: DataType) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), opset);
let input = graph.create_named_value("input", dtype, static_shape([]));
let output = graph.create_named_value("output", dtype, static_shape([]));
graph.add_input(input);
graph.add_output(output);
let mut node = Node::new(NodeId(0), "Einsum", vec![Some(input)], vec![output]);
node.name = "einsum".to_string();
node.attributes
.insert("equation".to_string(), Attribute::String(b"->".to_vec()));
graph.insert_node(node);
graph
}
#[test]
fn loader_resolves_einsum_schema_from_imported_opset() {
let error = validate_model(&einsum_graph(11, DataType::Float32)).unwrap_err();
assert!(matches!(error, LoaderError::InvalidEinsum { .. }));
assert!(error.to_string().contains("predates Einsum-12"));
for opset in [12, 27] {
let error = validate_model(&einsum_graph(opset, DataType::BFloat16)).unwrap_err();
assert!(matches!(error, LoaderError::InvalidEinsum { .. }));
assert!(error.to_string().contains("not admitted by Einsum-12"));
}
validate_model(&einsum_graph(28, DataType::BFloat16)).unwrap();
}
#[test]
fn loader_rejects_einsum_invalid_input_and_output_arity_before_metadata() {
let zero_inputs = load_model_bytes(&model(
12,
value_info("X", None, None),
value_info("Y", None, None),
vec![einsum_node_with_io(&[], &["Y"], "i->i")],
))
.unwrap_err();
assert!(matches!(zero_inputs, LoaderError::InvalidEinsum { .. }));
let zero_inputs = zero_inputs.to_string();
assert!(zero_inputs.contains("<unnamed node #0>"), "{zero_inputs}");
assert!(zero_inputs.contains("equation `i->i`"), "{zero_inputs}");
assert!(
zero_inputs.contains("expected at least one input"),
"{zero_inputs}"
);
assert!(zero_inputs.contains("found none"), "{zero_inputs}");
for output_count in [0, 2, 3] {
let outputs = ["Y", "extra_1", "extra_2"];
let invalid = load_model_bytes(&model(
12,
value_info("X", None, None),
value_info("X", None, None),
vec![einsum_node_with_io(
&["X"],
&outputs[..output_count],
"i->i",
)],
))
.unwrap_err();
assert!(matches!(invalid, LoaderError::InvalidEinsum { .. }));
let invalid = invalid.to_string();
assert!(invalid.contains("<unnamed node #0>"), "{invalid}");
assert!(invalid.contains("equation `i->i`"), "{invalid}");
assert!(
invalid.contains(&format!("declares {output_count} outputs")),
"{invalid}"
);
assert!(invalid.contains("requires exactly 1 output"), "{invalid}");
}
}
#[test]
fn loader_rejects_einsum_empty_required_output_name() {
let invalid = load_model_bytes(&model(
12,
value_info("X", None, None),
value_info("X", None, None),
vec![einsum_node_with_io(&["X"], &[""], "i->i")],
))
.unwrap_err();
assert!(matches!(invalid, LoaderError::InvalidEinsum { .. }));
let invalid = invalid.to_string();
assert!(invalid.contains("<unnamed node #0>"), "{invalid}");
assert!(invalid.contains("equation `i->i`"), "{invalid}");
assert!(invalid.contains("declares 1 output"), "{invalid}");
assert!(invalid.contains("required output #0"), "{invalid}");
assert!(invalid.contains("empty or omitted name"), "{invalid}");
}
#[test]
fn loader_accepts_single_named_einsum_output_with_unknown_metadata() {
let graph = load_model_bytes(&model(
12,
value_info("X", Some(1), Some(&[2])),
value_info("Y", None, None),
vec![einsum_node("X", "Y", "i->i")],
))
.unwrap();
let output = find(&graph, "Y");
assert!(!graph.value_type_is_known(output));
assert!(!graph.value_shape_is_known(output));
assert_eq!(graph.value(output).dtype, DataType::Float32);
assert_eq!(graph.value(output).shape, static_shape([2]));
}
#[test]
fn omitted_interior_value_info_is_inferred_after_fail_fast_validation() {
let graph = load_model_bytes(&model(
12,
value_info("X", Some(1), Some(&[2])),
value_info("Y", Some(1), Some(&[2])),
vec![
einsum_node("X", "interior", "i->i"),
einsum_node("interior", "Y", "i->i"),
],
))
.unwrap();
let interior = find(&graph, "interior");
assert!(!graph.value_type_is_known(interior));
assert!(!graph.value_shape_is_known(interior));
assert_eq!(graph.value(interior).dtype, DataType::Float32);
assert_eq!(graph.value(interior).shape, static_shape([2]));
}
#[test]
fn unannotated_f16_and_opset28_bf16_values_do_not_claim_placeholder_f32() {
for (opset, raw_dtype, dtype) in [(12, 10, DataType::Float16), (28, 16, DataType::BFloat16)] {
let graph = load_model_bytes(&model(
opset,
value_info("X", Some(raw_dtype), Some(&[3])),
value_info("Y", None, None),
vec![
einsum_node("X", "interior", "i->i"),
einsum_node("interior", "Y", "i->i"),
],
))
.unwrap_or_else(|error| panic!("opset {opset}, dtype {dtype:?}: {error}"));
for name in ["interior", "Y"] {
let value = find(&graph, name);
assert!(!graph.value_type_is_known(value), "{name}, opset {opset}");
assert!(!graph.value_shape_is_known(value), "{name}, opset {opset}");
assert_eq!(graph.value(value).dtype, dtype, "{name}, opset {opset}");
assert_eq!(
graph.value(value).shape,
static_shape([3]),
"{name}, opset {opset}"
);
}
}
}
#[test]
fn tensor_type_without_shape_is_known_type_not_a_declared_scalar() {
let graph = load_model_bytes(&model(
12,
value_info("X", Some(10), None),
value_info("Y", Some(10), None),
vec![einsum_node("X", "Y", "i->i")],
))
.unwrap();
for name in ["X", "Y"] {
let value = find(&graph, name);
assert!(graph.value_type_is_known(value), "{name}");
assert!(!graph.value_shape_is_known(value), "{name}");
assert_eq!(graph.value(value).dtype, DataType::Float16, "{name}");
assert!(graph.value(value).shape.is_empty(), "{name}");
}
}
#[test]
fn partial_metadata_still_rejects_every_known_invalid_fact() {
let invalid_dtype = load_model_bytes(&model(
27,
value_info("X", Some(16), None),
value_info("Y", None, None),
vec![einsum_node("X", "Y", "i->i")],
))
.unwrap_err();
assert!(
invalid_dtype
.to_string()
.contains("not admitted by Einsum-12")
);
let output_mismatch = load_model_bytes(&model(
12,
value_info("X", Some(10), None),
value_info("Y", Some(1), None),
vec![einsum_node("X", "Y", "i->i")],
))
.unwrap_err();
assert!(
output_mismatch
.to_string()
.contains("does not match known homogeneous input dtype")
);
let malformed_equation = load_model_bytes(&model(
12,
value_info("X", None, None),
value_info("Y", None, None),
vec![einsum_node("X", "Y", "i$->i")],
))
.unwrap_err();
assert!(malformed_equation.to_string().contains("invalid character"));
let mut known_shape_unknown_type = einsum_graph(12, DataType::Float32);
let input = known_shape_unknown_type.inputs[0];
let output = known_shape_unknown_type.outputs[0];
known_shape_unknown_type.value_mut(input).shape = static_shape([2, 3]);
known_shape_unknown_type.mark_value_type_unknown(input);
known_shape_unknown_type.mark_value_type_unknown(output);
let invalid_rank = validate_model(&known_shape_unknown_type).unwrap_err();
assert!(invalid_rank.to_string().contains("rank 2 does not match"));
let mut one_known_bad_shape = Graph::new();
one_known_bad_shape.opset_imports.insert(String::new(), 12);
let left =
one_known_bad_shape.create_named_value("left", DataType::Float32, static_shape([2, 3]));
let right = one_known_bad_shape.create_named_value("right", DataType::Float32, Vec::new());
let result = one_known_bad_shape.create_named_value("result", DataType::Float32, Vec::new());
one_known_bad_shape.mark_value_shape_unknown(right);
one_known_bad_shape.mark_value_type_unknown(result);
one_known_bad_shape.mark_value_shape_unknown(result);
one_known_bad_shape.add_input(left);
one_known_bad_shape.add_input(right);
one_known_bad_shape.add_output(result);
let mut node = Node::new(
NodeId(0),
"Einsum",
vec![Some(left), Some(right)],
vec![result],
);
node.attributes.insert(
"equation".to_string(),
Attribute::String(b"i,j->ij".to_vec()),
);
one_known_bad_shape.insert_node(node);
let invalid_partial_rank = validate_model(&one_known_bad_shape).unwrap_err();
assert!(
invalid_partial_rank
.to_string()
.contains("input #0 rank 2 does not match")
);
}
fn malformed_dtype(dtype: DeclaredDType) -> DataType {
match dtype {
DeclaredDType::Numeric(ConformanceDType::Uint8) => DataType::Uint8,
DeclaredDType::Numeric(ConformanceDType::Uint16) => DataType::Uint16,
DeclaredDType::Numeric(ConformanceDType::Uint32) => DataType::Uint32,
DeclaredDType::Numeric(ConformanceDType::Uint64) => DataType::Uint64,
DeclaredDType::Numeric(ConformanceDType::Int8) => DataType::Int8,
DeclaredDType::Numeric(ConformanceDType::Int16) => DataType::Int16,
DeclaredDType::Numeric(ConformanceDType::Int32) => DataType::Int32,
DeclaredDType::Numeric(ConformanceDType::Int64) => DataType::Int64,
DeclaredDType::Numeric(ConformanceDType::Float16) => DataType::Float16,
DeclaredDType::Numeric(ConformanceDType::Float32) => DataType::Float32,
DeclaredDType::Numeric(ConformanceDType::Float64) => DataType::Float64,
DeclaredDType::Numeric(ConformanceDType::BFloat16) => DataType::BFloat16,
DeclaredDType::Bool => DataType::Bool,
DeclaredDType::String => DataType::String,
DeclaredDType::Complex64 => DataType::Complex64,
DeclaredDType::Complex128 => DataType::Complex128,
}
}
fn malformed_graph(case: &MalformedCase) -> Graph {
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), case.opset);
let mut inputs = Vec::new();
for (index, (&dtype, shape)) in case.input_dtypes.iter().zip(&case.input_shapes).enumerate() {
let value = graph.create_named_value(
format!("input_{index}"),
malformed_dtype(dtype),
static_shape(shape.iter().copied()),
);
graph.add_input(value);
inputs.push(Some(value));
}
let output_dtype = case
.input_dtypes
.first()
.copied()
.map(malformed_dtype)
.unwrap_or(DataType::Float32);
let mut outputs = Vec::new();
for index in 0..case.output_count {
let value =
graph.create_named_value(format!("output_{index}"), output_dtype, static_shape([]));
graph.add_output(value);
outputs.push(value);
}
let mut node = Node::new(NodeId(0), "Einsum", inputs, outputs);
node.name = case.id.clone();
node.attributes.insert(
"equation".into(),
Attribute::String(case.equation.as_bytes().to_vec()),
);
graph.insert_node(node);
graph
}
#[test]
fn loader_accepts_the_independent_64_operand_legal_case() {
let case = named_cases()
.into_iter()
.find(|case| case.id == "scalar-product-64-operands")
.expect("high-arity conformance case");
let mut graph = Graph::new();
graph.opset_imports.insert(String::new(), case.opset);
let mut inputs = Vec::new();
for (index, shape) in case.input_shapes.iter().enumerate() {
let value = graph.create_named_value(
format!("input_{index}"),
malformed_dtype(DeclaredDType::Numeric(case.dtype)),
static_shape(shape.iter().copied()),
);
graph.add_input(value);
inputs.push(Some(value));
}
let output = graph.create_named_value("output", DataType::Float32, static_shape([]));
graph.add_output(output);
let mut node = Node::new(NodeId(0), "Einsum", inputs, vec![output]);
node.name = case.id;
node.attributes.insert(
"equation".into(),
Attribute::String(case.equation.into_bytes()),
);
graph.insert_node(node);
validate_model(&graph).unwrap();
}
#[test]
fn loader_rejects_every_record_in_the_independent_malformed_corpus() {
let cases = malformed_cases();
assert!(!cases.is_empty());
for case in cases {
let error = match validate_model(&malformed_graph(&case)) {
Ok(()) => panic!("{} unexpectedly passed loader validation", case.id),
Err(error) => error,
};
assert!(
matches!(error, LoaderError::InvalidEinsum { .. }),
"{}: {error}",
case.id
);
}
}