1use std::collections::HashMap;
8
9use onnx_runtime_ir::DataType;
10use serde::{Deserialize, Serialize};
11
12#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
14pub struct OpSchema {
15 #[serde(default)]
17 pub domain: String,
18 pub name: String,
20 pub since_version: u64,
22 #[serde(default)]
24 pub until_version: Option<u64>,
25 #[serde(default)]
27 pub doc: String,
28 #[serde(default)]
30 pub inputs: Vec<InputSpec>,
31 #[serde(default)]
33 pub outputs: Vec<OutputSpec>,
34 #[serde(default)]
36 pub attributes: Vec<AttributeSpec>,
37 #[serde(default)]
39 pub type_constraints: Vec<TypeConstraint>,
40}
41
42impl OpSchema {
43 pub fn supports_opset(&self, opset: u64) -> bool {
45 self.since_version <= opset && self.until_version.is_none_or(|until| opset <= until)
46 }
47}
48
49#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
51pub struct InputSpec {
52 pub name: String,
54 pub type_str: String,
56 #[serde(default)]
58 pub doc: String,
59 #[serde(default)]
61 pub optional: bool,
62 #[serde(default)]
64 pub variadic: bool,
65 #[serde(default = "default_min_arity")]
67 pub min_arity: usize,
68}
69
70#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
72pub struct OutputSpec {
73 pub name: String,
75 pub type_str: String,
77 #[serde(default)]
79 pub doc: String,
80 #[serde(default)]
82 pub optional: bool,
83 #[serde(default)]
85 pub variadic: bool,
86 #[serde(default = "default_min_arity")]
88 pub min_arity: usize,
89}
90
91const fn default_min_arity() -> usize {
92 1
93}
94
95#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
97#[serde(rename_all = "snake_case")]
98pub enum AttributeType {
99 Int,
101 Float,
103 String,
105 Tensor,
107 Graph,
109 SparseTensor,
111 TypeProto,
113 Ints,
115 Floats,
117 Strings,
119 Graphs,
121 Tensors,
123 SparseTensors,
125 TypeProtos,
127}
128
129#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
131#[serde(untagged)]
132pub enum AttributeDefault {
133 Int(i64),
135 Float(f64),
137 String(String),
139 Ints(Vec<i64>),
141 Floats(Vec<f64>),
143 Strings(Vec<String>),
145}
146
147#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
149pub struct AttributeSpec {
150 pub name: String,
152 #[serde(rename = "type")]
154 pub attr_type: AttributeType,
155 #[serde(default)]
157 pub required: bool,
158 #[serde(default)]
160 pub default: Option<AttributeDefault>,
161 #[serde(default)]
163 pub doc: String,
164}
165
166#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
168pub struct TypeConstraint {
169 pub type_param: String,
171 #[serde(with = "data_types")]
173 pub allowed: Vec<DataType>,
174}
175
176#[derive(Debug, thiserror::Error)]
178pub enum SchemaError {
179 #[error("invalid op-schema YAML: {0}")]
181 Yaml(#[from] serde_yaml::Error),
182 #[error("invalid schema {domain}::{name}: {message}")]
184 Invalid {
185 domain: String,
187 name: String,
189 message: String,
191 },
192}
193
194#[derive(Clone, Debug, Default)]
196pub struct SchemaRegistry {
197 schemas: HashMap<(String, String), Vec<OpSchema>>,
198}
199
200impl SchemaRegistry {
201 pub fn new() -> Self {
203 Self::default()
204 }
205
206 pub fn load_yaml(&mut self, yaml: &str) -> Result<(), SchemaError> {
208 self.register(serde_yaml::from_str(yaml)?)
209 }
210
211 pub fn register(&mut self, mut schema: OpSchema) -> Result<(), SchemaError> {
213 schema.domain = normalize_domain(&schema.domain).to_string();
214 validate_schema(&schema)?;
215 let key = (schema.domain.clone(), schema.name.clone());
216 let versions = self.schemas.entry(key).or_default();
217 if versions
218 .iter()
219 .any(|current| current.since_version == schema.since_version)
220 {
221 return Err(SchemaError::Invalid {
222 domain: schema.domain.clone(),
223 name: schema.name.clone(),
224 message: format!(
225 "a schema already exists at since_version {}",
226 schema.since_version
227 ),
228 });
229 }
230 versions.push(schema);
231 versions.sort_by_key(|schema| schema.since_version);
232 Ok(())
233 }
234
235 pub fn lookup(&self, op_type: &str, domain: &str, opset: u64) -> Option<&OpSchema> {
237 self.schemas
238 .get(&(normalize_domain(domain).to_string(), op_type.to_string()))?
239 .iter()
240 .rev()
241 .find(|schema| schema.supports_opset(opset))
242 }
243
244 pub fn contains_operator(&self, op_type: &str, domain: &str) -> bool {
246 self.schemas
247 .contains_key(&(normalize_domain(domain).to_string(), op_type.to_string()))
248 }
249
250 pub fn iter(&self) -> impl Iterator<Item = &OpSchema> {
252 self.schemas.values().flatten()
253 }
254
255 pub fn builtins() -> Self {
257 let mut registry = Self::new();
258 for yaml in BUILTIN_YAML {
259 registry
260 .load_yaml(yaml)
261 .expect("embedded ONNX op schema must be valid");
262 }
263 registry
264 }
265}
266
267fn normalize_domain(domain: &str) -> &str {
268 if domain.is_empty() || domain == "ai.onnx" {
269 "ai.onnx"
270 } else {
271 domain
272 }
273}
274
275fn validate_schema(schema: &OpSchema) -> Result<(), SchemaError> {
276 let invalid = |message: &str| SchemaError::Invalid {
277 domain: schema.domain.clone(),
278 name: schema.name.clone(),
279 message: message.into(),
280 };
281 if schema.name.is_empty() {
282 return Err(invalid("operator name must not be empty"));
283 }
284 if schema.since_version == 0 {
285 return Err(invalid("since_version must be at least 1"));
286 }
287 if schema
288 .until_version
289 .is_some_and(|until| until < schema.since_version)
290 {
291 return Err(invalid("until_version precedes since_version"));
292 }
293 if schema.inputs.iter().filter(|input| input.variadic).count() > 1
294 || schema
295 .inputs
296 .iter()
297 .position(|input| input.variadic)
298 .is_some_and(|index| index + 1 != schema.inputs.len())
299 {
300 return Err(invalid(
301 "a variadic input must be the only trailing variadic",
302 ));
303 }
304 if schema
305 .outputs
306 .iter()
307 .filter(|output| output.variadic)
308 .count()
309 > 1
310 || schema
311 .outputs
312 .iter()
313 .position(|output| output.variadic)
314 .is_some_and(|index| index + 1 != schema.outputs.len())
315 {
316 return Err(invalid(
317 "a variadic output must be the only trailing variadic",
318 ));
319 }
320 if schema.attributes.iter().any(|attribute| {
321 attribute
322 .default
323 .as_ref()
324 .is_some_and(|value| !default_matches(value, attribute.attr_type))
325 }) {
326 return Err(invalid(
327 "an attribute default does not match its declared type",
328 ));
329 }
330 Ok(())
331}
332
333fn default_matches(value: &AttributeDefault, attr_type: AttributeType) -> bool {
334 matches!(
335 (value, attr_type),
336 (AttributeDefault::Int(_), AttributeType::Int)
337 | (AttributeDefault::Float(_), AttributeType::Float)
338 | (AttributeDefault::String(_), AttributeType::String)
339 | (AttributeDefault::Ints(_), AttributeType::Ints)
340 | (AttributeDefault::Floats(_), AttributeType::Floats)
341 | (AttributeDefault::Strings(_), AttributeType::Strings)
342 )
343}
344
345const BUILTIN_YAML: &[&str] = &[
346 include_str!("../../schemas/standard/matmul.yaml"),
347 include_str!("../../schemas/standard/gemm.yaml"),
348 include_str!("../../schemas/standard/add.yaml"),
349 include_str!("../../schemas/standard/sub.yaml"),
350 include_str!("../../schemas/standard/div.yaml"),
351 include_str!("../../schemas/standard/relu.yaml"),
352 include_str!("../../schemas/standard/conv.yaml"),
353 include_str!("../../schemas/standard/mul.yaml"),
354 include_str!("../../schemas/standard/identity.yaml"),
355 include_str!("../../schemas/standard/if.yaml"),
356 include_str!("../../schemas/standard/softmax.yaml"),
357 include_str!("../../schemas/standard/layer_normalization.yaml"),
358 include_str!("../../schemas/standard/gather.yaml"),
359 include_str!("../../schemas/standard/reshape_v14.yaml"),
360 include_str!("../../schemas/standard/reshape_v19.yaml"),
361 include_str!("../../schemas/standard/reshape_v21.yaml"),
362 include_str!("../../schemas/standard/reshape_v23.yaml"),
363 include_str!("../../schemas/standard/reshape.yaml"),
364 include_str!("../../schemas/standard/transpose_v13.yaml"),
365 include_str!("../../schemas/standard/transpose_v21.yaml"),
366 include_str!("../../schemas/standard/transpose_v23.yaml"),
367 include_str!("../../schemas/standard/transpose.yaml"),
368 include_str!("../../schemas/standard/concat.yaml"),
369 include_str!("../../schemas/standard/slice.yaml"),
370 include_str!("../../schemas/standard/sigmoid.yaml"),
371 include_str!("../../schemas/standard/tanh.yaml"),
372 include_str!("../../schemas/standard/erf.yaml"),
373 include_str!("../../schemas/standard/sqrt.yaml"),
374 include_str!("../../schemas/standard/exp.yaml"),
375 include_str!("../../schemas/standard/log.yaml"),
376 include_str!("../../schemas/standard/pow.yaml"),
377 include_str!("../../schemas/standard/clip.yaml"),
378 include_str!("../../schemas/standard/expand.yaml"),
379 include_str!("../../schemas/standard/where.yaml"),
380 include_str!("../../schemas/standard/reduce_sum.yaml"),
381 include_str!("../../schemas/standard/reduce_mean.yaml"),
382 include_str!("../../schemas/standard/neg.yaml"),
383 include_str!("../../schemas/standard/abs.yaml"),
384 include_str!("../../schemas/standard/mod.yaml"),
385 include_str!("../../schemas/standard/log_softmax.yaml"),
386 include_str!("../../schemas/standard/rms_normalization.yaml"),
387 include_str!("../../schemas/standard/reduce_max.yaml"),
388 include_str!("../../schemas/standard/reduce_min.yaml"),
389 include_str!("../../schemas/standard/reduce_prod.yaml"),
390 include_str!("../../schemas/standard/reduce_l1.yaml"),
391 include_str!("../../schemas/standard/reduce_l2.yaml"),
392 include_str!("../../schemas/standard/reduce_log_sum.yaml"),
393 include_str!("../../schemas/standard/reduce_log_sum_exp.yaml"),
394 include_str!("../../schemas/standard/reduce_sum_square.yaml"),
395 include_str!("../../schemas/standard/arg_max.yaml"),
396 include_str!("../../schemas/standard/arg_min.yaml"),
397 include_str!("../../schemas/standard/gather_elements.yaml"),
398 include_str!("../../schemas/standard/gather_nd.yaml"),
399 include_str!("../../schemas/standard/equal.yaml"),
400 include_str!("../../schemas/standard/greater.yaml"),
401 include_str!("../../schemas/standard/less.yaml"),
402 include_str!("../../schemas/standard/and.yaml"),
403 include_str!("../../schemas/standard/or.yaml"),
404 include_str!("../../schemas/standard/not.yaml"),
405 include_str!("../../schemas/standard/cast.yaml"),
406 include_str!("../../schemas/standard/shape.yaml"),
407 include_str!("../../schemas/standard/size.yaml"),
408 include_str!("../../schemas/standard/non_zero.yaml"),
409 include_str!("../../schemas/standard/range.yaml"),
410 include_str!("../../schemas/standard/split.yaml"),
411 include_str!("../../schemas/standard/tile.yaml"),
412 include_str!("../../schemas/standard/pad.yaml"),
413 include_str!("../../schemas/standard/scatter_nd.yaml"),
414 include_str!("../../schemas/standard/scatter_elements.yaml"),
415 include_str!("../../schemas/standard/constant_of_shape.yaml"),
416 include_str!("../../schemas/standard/max_pool.yaml"),
417 include_str!("../../schemas/standard/average_pool.yaml"),
418 include_str!("../../schemas/standard/global_average_pool.yaml"),
419 include_str!("../../schemas/standard/global_max_pool.yaml"),
420 include_str!("../../schemas/standard/resize.yaml"),
421 include_str!("../../schemas/standard/quantize_linear.yaml"),
422 include_str!("../../schemas/standard/dequantize_linear.yaml"),
423 include_str!("../../schemas/standard/dynamic_quantize_linear.yaml"),
424 include_str!("../../schemas/standard/attention.yaml"),
425 include_str!("../../schemas/standard/cast_like.yaml"),
426 include_str!("../../schemas/standard/cum_sum.yaml"),
427 include_str!("../../schemas/standard/greater_or_equal.yaml"),
428 include_str!("../../schemas/standard/less_or_equal.yaml"),
429 include_str!("../../schemas/standard/min.yaml"),
430 include_str!("../../schemas/standard/max.yaml"),
431 include_str!("../../schemas/standard/rotary_embedding.yaml"),
432 include_str!("../../schemas/standard/softplus.yaml"),
433 include_str!("../../schemas/standard/squeeze.yaml"),
434 include_str!("../../schemas/standard/top_k.yaml"),
435 include_str!("../../schemas/standard/unsqueeze_v11.yaml"),
436 include_str!("../../schemas/standard/unsqueeze_v13.yaml"),
437 include_str!("../../schemas/standard/unsqueeze.yaml"),
438];
439
440mod data_types {
443 use onnx_runtime_ir::DataType;
444 use serde::{Deserialize, Deserializer, Serialize, Serializer};
445
446 pub fn serialize<S>(types: &[DataType], serializer: S) -> Result<S::Ok, S::Error>
447 where
448 S: Serializer,
449 {
450 types
451 .iter()
452 .map(|data_type| name(*data_type))
453 .collect::<Vec<_>>()
454 .serialize(serializer)
455 }
456
457 pub fn deserialize<'de, D>(deserializer: D) -> Result<Vec<DataType>, D::Error>
458 where
459 D: Deserializer<'de>,
460 {
461 Vec::<String>::deserialize(deserializer)?
462 .into_iter()
463 .map(|value| {
464 parse(&value)
465 .ok_or_else(|| serde::de::Error::custom(format!("unknown data type '{value}'")))
466 })
467 .collect()
468 }
469
470 fn parse(value: &str) -> Option<DataType> {
471 Some(match value {
472 "undefined" => DataType::Undefined,
473 "float32" => DataType::Float32,
474 "uint8" => DataType::Uint8,
475 "int8" => DataType::Int8,
476 "uint16" => DataType::Uint16,
477 "int16" => DataType::Int16,
478 "int32" => DataType::Int32,
479 "int64" => DataType::Int64,
480 "string" => DataType::String,
481 "bool" => DataType::Bool,
482 "float16" => DataType::Float16,
483 "float64" => DataType::Float64,
484 "uint32" => DataType::Uint32,
485 "uint64" => DataType::Uint64,
486 "complex64" => DataType::Complex64,
487 "complex128" => DataType::Complex128,
488 "bfloat16" => DataType::BFloat16,
489 "float8e4m3fn" => DataType::Float8E4M3FN,
490 "float8e4m3fnuz" => DataType::Float8E4M3FNUZ,
491 "float8e5m2" => DataType::Float8E5M2,
492 "float8e5m2fnuz" => DataType::Float8E5M2FNUZ,
493 "uint4" => DataType::Uint4,
494 "int4" => DataType::Int4,
495 "float4e2m1" => DataType::Float4E2M1,
496 "float8e8m0" => DataType::Float8E8M0,
497 "uint2" => DataType::Uint2,
498 "int2" => DataType::Int2,
499 _ => return None,
500 })
501 }
502
503 fn name(value: DataType) -> &'static str {
504 match value {
505 DataType::Undefined => "undefined",
506 DataType::Float32 => "float32",
507 DataType::Uint8 => "uint8",
508 DataType::Int8 => "int8",
509 DataType::Uint16 => "uint16",
510 DataType::Int16 => "int16",
511 DataType::Int32 => "int32",
512 DataType::Int64 => "int64",
513 DataType::String => "string",
514 DataType::Bool => "bool",
515 DataType::Float16 => "float16",
516 DataType::Float64 => "float64",
517 DataType::Uint32 => "uint32",
518 DataType::Uint64 => "uint64",
519 DataType::Complex64 => "complex64",
520 DataType::Complex128 => "complex128",
521 DataType::BFloat16 => "bfloat16",
522 DataType::Float8E4M3FN => "float8e4m3fn",
523 DataType::Float8E4M3FNUZ => "float8e4m3fnuz",
524 DataType::Float8E5M2 => "float8e5m2",
525 DataType::Float8E5M2FNUZ => "float8e5m2fnuz",
526 DataType::Uint4 => "uint4",
527 DataType::Int4 => "int4",
528 DataType::Float4E2M1 => "float4e2m1",
529 DataType::Float8E8M0 => "float8e8m0",
530 DataType::Uint2 => "uint2",
531 DataType::Int2 => "int2",
532 }
533 }
534}
535
536#[cfg(test)]
537mod tests {
538 use super::*;
539
540 const RELU_V6: &str = r#"
541domain: ""
542name: Relu
543since_version: 6
544inputs: [{ name: X, type_str: T }]
545outputs: [{ name: Y, type_str: T }]
546type_constraints:
547 - type_param: T
548 allowed: [float16, float32]
549"#;
550
551 #[test]
552 fn yaml_schema_round_trips_every_public_field() {
553 let schema: OpSchema = serde_yaml::from_str(
554 r#"
555domain: example
556name: Variadic
557since_version: 2
558until_version: 4
559doc: example op
560inputs:
561 - { name: X, type_str: T, doc: input, optional: true, variadic: true, min_arity: 2 }
562outputs:
563 - { name: Y, type_str: T, doc: output, optional: true, variadic: true }
564attributes:
565 - { name: axis, type: int, required: true, default: 1, doc: axis }
566type_constraints:
567 - { type_param: T, allowed: [float32, int64] }
568"#,
569 )
570 .unwrap();
571 assert_eq!(schema.domain, "example");
572 assert!(schema.supports_opset(3));
573 assert!(!schema.supports_opset(5));
574 assert!(schema.inputs[0].optional && schema.inputs[0].variadic);
575 assert!(schema.outputs[0].optional && schema.outputs[0].variadic);
576 assert_eq!(schema.inputs[0].min_arity, 2);
577 assert_eq!(schema.outputs[0].min_arity, 1);
578 assert_eq!(schema.attributes[0].attr_type, AttributeType::Int);
579 assert_eq!(schema.attributes[0].default, Some(AttributeDefault::Int(1)));
580 assert_eq!(
581 schema.type_constraints[0].allowed,
582 vec![DataType::Float32, DataType::Int64]
583 );
584 let encoded = serde_yaml::to_string(&schema).unwrap();
585 assert_eq!(serde_yaml::from_str::<OpSchema>(&encoded).unwrap(), schema);
586 }
587
588 #[test]
589 fn registry_resolves_domains_and_opset_ranges() {
590 let mut registry = SchemaRegistry::new();
591 registry.load_yaml(RELU_V6).unwrap();
592 let mut newer: OpSchema = serde_yaml::from_str(RELU_V6).unwrap();
593 newer.since_version = 13;
594 newer.until_version = None;
595 registry.register(newer).unwrap();
596 assert_eq!(registry.lookup("Relu", "", 10).unwrap().since_version, 6);
597 assert_eq!(
598 registry
599 .lookup("Relu", "ai.onnx", 21)
600 .unwrap()
601 .since_version,
602 13
603 );
604 assert!(registry.lookup("Relu", "", 5).is_none());
605 assert!(registry.contains_operator("Relu", ""));
606 assert_eq!(registry.iter().count(), 2);
607 }
608
609 #[test]
610 fn registry_rejects_invalid_and_duplicate_version_schemas() {
611 let mut registry = SchemaRegistry::new();
612 registry.load_yaml(RELU_V6).unwrap();
613 assert!(registry.load_yaml(RELU_V6).is_err());
614 let invalid = RELU_V6.replace("since_version: 6", "since_version: 0");
615 assert!(matches!(
616 SchemaRegistry::new().load_yaml(&invalid),
617 Err(SchemaError::Invalid { .. })
618 ));
619 assert!(matches!(
620 SchemaRegistry::new().load_yaml("not: [valid"),
621 Err(SchemaError::Yaml(_))
622 ));
623 }
624
625 #[test]
626 fn builtins_contain_expected_common_ops() {
627 let registry = SchemaRegistry::builtins();
628 for name in [
629 "MatMul",
630 "Gemm",
631 "Add",
632 "Sub",
633 "Div",
634 "Relu",
635 "Conv",
636 "Mul",
637 "Identity",
638 "If",
639 "Softmax",
640 "LayerNormalization",
641 "Gather",
642 "Reshape",
643 "Transpose",
644 "Concat",
645 "Slice",
646 "Sigmoid",
647 "Tanh",
648 "Erf",
649 "Sqrt",
650 "Exp",
651 "Log",
652 "Pow",
653 "Clip",
654 "Expand",
655 "Where",
656 "ReduceSum",
657 "ReduceMean",
658 "Neg",
659 "Abs",
660 "Mod",
661 "LogSoftmax",
662 "RMSNormalization",
663 "ReduceMax",
664 "ReduceMin",
665 "ReduceProd",
666 "ReduceL1",
667 "ReduceL2",
668 "ReduceLogSum",
669 "ReduceLogSumExp",
670 "ReduceSumSquare",
671 "ArgMax",
672 "ArgMin",
673 "GatherElements",
674 "GatherND",
675 "Equal",
676 "Greater",
677 "Less",
678 "And",
679 "Or",
680 "Not",
681 "Cast",
682 "Shape",
683 "Size",
684 "NonZero",
685 "Range",
686 "Split",
687 "Tile",
688 "Pad",
689 "ScatterND",
690 "ScatterElements",
691 "ConstantOfShape",
692 "MaxPool",
693 "AveragePool",
694 "GlobalAveragePool",
695 "GlobalMaxPool",
696 "Resize",
697 "QuantizeLinear",
698 "DequantizeLinear",
699 "DynamicQuantizeLinear",
700 ] {
701 assert!(registry.lookup(name, "", 25).is_some(), "{name}");
702 }
703 }
704
705 #[test]
706 fn round_five_schemas_match_official_signatures() {
707 let registry = SchemaRegistry::builtins();
708
709 for (name, since_version) in [
710 ("GatherElements", 13),
711 ("GatherND", 13),
712 ("Equal", 19),
713 ("Greater", 13),
714 ("Less", 13),
715 ("And", 7),
716 ("Or", 7),
717 ("Not", 1),
718 ("Cast", 25),
719 ("Shape", 25),
720 ("Size", 25),
721 ("NonZero", 13),
722 ("Range", 11),
723 ("Split", 18),
724 ] {
725 assert_eq!(
726 registry.lookup(name, "", 25).unwrap().since_version,
727 since_version,
728 "{name}"
729 );
730 }
731
732 let gather_elements = registry.lookup("GatherElements", "", 25).unwrap();
733 assert_eq!(
734 gather_elements.attributes[0].default,
735 Some(AttributeDefault::Int(0))
736 );
737 assert_eq!(
738 gather_elements.type_constraints[1].allowed,
739 [DataType::Int32, DataType::Int64]
740 );
741
742 let gather_nd = registry.lookup("GatherND", "", 25).unwrap();
743 assert_eq!(gather_nd.inputs[1].type_str, "tensor(int64)");
744 assert_eq!(
745 gather_nd.attributes[0].default,
746 Some(AttributeDefault::Int(0))
747 );
748
749 for name in ["Equal", "Greater", "Less", "And", "Or"] {
750 let schema = registry.lookup(name, "", 25).unwrap();
751 assert_eq!(schema.outputs[0].type_str, "T1");
752 assert_eq!(
753 schema.type_constraints.last().unwrap().allowed,
754 [DataType::Bool],
755 "{name}"
756 );
757 }
758 assert_eq!(
759 registry.lookup("Not", "", 25).unwrap().type_constraints[0].allowed,
760 [DataType::Bool]
761 );
762
763 let cast = registry.lookup("Cast", "", 25).unwrap();
764 assert_eq!(
765 cast.attributes
766 .iter()
767 .find(|attribute| attribute.name == "round_mode")
768 .unwrap()
769 .default,
770 Some(AttributeDefault::String("up".into()))
771 );
772 assert!(
773 cast.attributes
774 .iter()
775 .find(|attribute| attribute.name == "to")
776 .unwrap()
777 .required
778 );
779 assert_eq!(cast.type_constraints[0].allowed.len(), 24);
780
781 for name in ["Shape", "Size"] {
782 let schema = registry.lookup(name, "", 25).unwrap();
783 assert_eq!(schema.type_constraints[0].allowed.len(), 26);
784 assert_eq!(schema.type_constraints[1].allowed, [DataType::Int64]);
785 }
786
787 let split = registry.lookup("Split", "", 25).unwrap();
788 assert!(split.inputs[1].optional);
789 assert!(split.outputs[0].variadic);
790 assert_eq!(split.outputs[0].min_arity, 1);
791 assert_eq!(
792 registry.lookup("Range", "", 25).unwrap().type_constraints[0].allowed,
793 [
794 DataType::Float32,
795 DataType::Float64,
796 DataType::Int16,
797 DataType::Int32,
798 DataType::Int64
799 ]
800 );
801 }
802
803 #[test]
804 fn round_six_schemas_match_official_signatures() {
805 let registry = SchemaRegistry::builtins();
806
807 for (name, since_version, inputs) in [
808 ("Slice", 13, 5),
809 ("Concat", 13, 1),
810 ("Tile", 13, 2),
811 ("Expand", 13, 2),
812 ("Pad", 25, 4),
813 ("ScatterND", 18, 3),
814 ("ScatterElements", 18, 3),
815 ("ConstantOfShape", 25, 1),
816 ] {
817 let schema = registry.lookup(name, "", 25).unwrap();
818 assert_eq!(schema.since_version, since_version, "{name}");
819 assert_eq!(schema.inputs.len(), inputs, "{name}");
820 assert_eq!(schema.outputs.len(), 1, "{name}");
821 }
822
823 let tile = registry.lookup("Tile", "", 25).unwrap();
824 assert_eq!(tile.type_constraints[1].allowed, [DataType::Int64]);
825
826 let pad = registry.lookup("Pad", "", 25).unwrap();
827 assert_eq!(
828 pad.attributes[0].default,
829 Some(AttributeDefault::String("constant".into()))
830 );
831 assert!(pad.inputs[2].optional && pad.inputs[3].optional);
832 assert_eq!(pad.type_constraints[0].allowed.len(), 26);
833 assert_eq!(
834 pad.type_constraints[1].allowed,
835 [DataType::Int32, DataType::Int64]
836 );
837
838 for name in ["ScatterND", "ScatterElements"] {
839 let scatter = registry.lookup(name, "", 25).unwrap();
840 assert_eq!(
841 scatter
842 .attributes
843 .iter()
844 .find(|attribute| attribute.name == "reduction")
845 .and_then(|attribute| attribute.default.clone()),
846 Some(AttributeDefault::String("none".into())),
847 "{name}"
848 );
849 }
850 let scatter_elements = registry.lookup("ScatterElements", "", 25).unwrap();
851 assert_eq!(
852 scatter_elements.attributes[0].default,
853 Some(AttributeDefault::Int(0))
854 );
855 assert_eq!(
856 scatter_elements.type_constraints[1].allowed,
857 [DataType::Int32, DataType::Int64]
858 );
859
860 let constant = registry.lookup("ConstantOfShape", "", 25).unwrap();
861 assert!(!constant.attributes[0].required);
862 assert_eq!(constant.type_constraints[0].allowed, [DataType::Int64]);
863 assert_eq!(constant.type_constraints[1].allowed.len(), 23);
864 }
865
866 #[test]
867 fn round_seven_schemas_match_official_signatures() {
868 let registry = SchemaRegistry::builtins();
869 for (name, since_version, inputs, outputs) in [
870 ("MaxPool", 22, 1, 2),
871 ("AveragePool", 22, 1, 1),
872 ("GlobalAveragePool", 22, 1, 1),
873 ("GlobalMaxPool", 22, 1, 1),
874 ("Resize", 19, 4, 1),
875 ("QuantizeLinear", 21, 3, 1),
876 ("DequantizeLinear", 21, 3, 1),
877 ("DynamicQuantizeLinear", 11, 1, 3),
878 ] {
879 let schema = registry.lookup(name, "", 25).unwrap();
880 assert_eq!(schema.since_version, since_version, "{name}");
881 assert_eq!(schema.inputs.len(), inputs, "{name}");
882 assert_eq!(schema.outputs.len(), outputs, "{name}");
883 }
884
885 let max_pool = registry.lookup("MaxPool", "", 25).unwrap();
886 assert!(max_pool.outputs[1].optional);
887 assert_eq!(max_pool.type_constraints[1].allowed, [DataType::Int64]);
888 assert_eq!(max_pool.type_constraints[0].allowed.len(), 6);
889 assert!(
890 max_pool
891 .attributes
892 .iter()
893 .find(|attribute| attribute.name == "kernel_shape")
894 .unwrap()
895 .required
896 );
897
898 let average_pool = registry.lookup("AveragePool", "", 25).unwrap();
899 assert_eq!(average_pool.type_constraints[0].allowed.len(), 4);
900 assert_eq!(
901 average_pool
902 .attributes
903 .iter()
904 .find(|attribute| attribute.name == "count_include_pad")
905 .unwrap()
906 .default,
907 Some(AttributeDefault::Int(0))
908 );
909
910 let resize = registry.lookup("Resize", "", 25).unwrap();
911 assert!(resize.inputs[1..].iter().all(|input| input.optional));
912 assert_eq!(resize.attributes.len(), 9);
913 assert_eq!(resize.type_constraints[0].allowed.len(), 16);
914
915 let quantize = registry.lookup("QuantizeLinear", "", 25).unwrap();
916 assert!(quantize.inputs[2].optional);
917 assert_eq!(quantize.type_constraints[1].allowed.len(), 10);
918 for dtype in [
919 DataType::Uint4,
920 DataType::Int4,
921 DataType::Float8E4M3FN,
922 DataType::Float8E4M3FNUZ,
923 DataType::Float8E5M2,
924 DataType::Float8E5M2FNUZ,
925 ] {
926 assert!(quantize.type_constraints[1].allowed.contains(&dtype));
927 }
928
929 let dequantize = registry.lookup("DequantizeLinear", "", 25).unwrap();
930 assert!(dequantize.inputs[2].optional);
931 assert_eq!(dequantize.type_constraints[0].allowed.len(), 11);
932
933 let dynamic = registry.lookup("DynamicQuantizeLinear", "", 25).unwrap();
934 assert_eq!(dynamic.outputs[1].type_str, "tensor(float)");
935 assert_eq!(dynamic.type_constraints[1].allowed, [DataType::Uint8]);
936 }
937
938 #[test]
939 fn round_eight_schemas_match_official_signatures() {
940 let registry = SchemaRegistry::builtins();
941 for (name, since_version, inputs, outputs, attributes) in [
942 ("Attention", 24, 7, 4, 7),
943 ("CastLike", 24, 2, 1, 2),
944 ("CumSum", 14, 2, 1, 2),
945 ("GreaterOrEqual", 16, 2, 1, 0),
946 ("LessOrEqual", 16, 2, 1, 0),
947 ("Min", 13, 1, 1, 0),
948 ("Max", 13, 1, 1, 0),
949 ("RotaryEmbedding", 23, 4, 1, 3),
950 ("Softplus", 22, 1, 1, 0),
951 ("Squeeze", 24, 2, 1, 0),
952 ("TopK", 24, 2, 2, 3),
953 ("Unsqueeze", 24, 2, 1, 0),
954 ] {
955 let schema = registry.lookup(name, "", 24).unwrap();
956 assert_eq!(schema.since_version, since_version, "{name}");
957 assert_eq!(schema.inputs.len(), inputs, "{name}");
958 assert_eq!(schema.outputs.len(), outputs, "{name}");
959 assert_eq!(schema.attributes.len(), attributes, "{name}");
960 }
961
962 let attention = registry.lookup("Attention", "", 24).unwrap();
963 assert!(attention.inputs[3..].iter().all(|input| input.optional));
964 assert!(attention.outputs[1..].iter().all(|output| output.optional));
965 assert_eq!(
966 attention.type_constraints[2].allowed.last(),
967 Some(&DataType::Bool)
968 );
969 for name in ["Min", "Max"] {
970 assert!(registry.lookup(name, "", 24).unwrap().inputs[0].variadic);
971 }
972 assert!(registry.lookup("Squeeze", "", 24).unwrap().inputs[1].optional);
973 let unsqueeze_v11 = registry.lookup("Unsqueeze", "", 11).unwrap();
974 assert_eq!(unsqueeze_v11.inputs.len(), 1);
975 assert!(unsqueeze_v11.attributes[0].required);
976 assert_eq!(
977 registry.lookup("Unsqueeze", "", 13).unwrap().inputs.len(),
978 2
979 );
980 assert_eq!(
981 registry.lookup("TopK", "", 24).unwrap().outputs[1].type_str,
982 "I"
983 );
984 assert_eq!(
985 registry
986 .lookup("CastLike", "", 24)
987 .unwrap()
988 .attributes
989 .iter()
990 .find(|attribute| attribute.name == "round_mode")
991 .unwrap()
992 .default,
993 Some(AttributeDefault::String("up".into()))
994 );
995 }
996
997 #[test]
998 fn round_four_schemas_match_official_signatures() {
999 let registry = SchemaRegistry::builtins();
1000 let reductions = [
1001 ("ReduceMax", 20),
1002 ("ReduceMin", 20),
1003 ("ReduceProd", 18),
1004 ("ReduceL1", 18),
1005 ("ReduceL2", 18),
1006 ("ReduceLogSum", 18),
1007 ("ReduceLogSumExp", 18),
1008 ("ReduceSumSquare", 18),
1009 ];
1010 for (name, since_version) in reductions {
1011 let schema = registry.lookup(name, "", 24).unwrap();
1012 assert_eq!(schema.since_version, since_version);
1013 assert_eq!(schema.inputs.len(), 2);
1014 assert!(schema.inputs[1].optional);
1015 assert_eq!(schema.inputs[1].type_str, "tensor(int64)");
1016 assert_eq!(schema.outputs.len(), 1);
1017 assert_eq!(
1018 schema
1019 .attributes
1020 .iter()
1021 .find(|attribute| attribute.name == "keepdims")
1022 .unwrap()
1023 .default,
1024 Some(AttributeDefault::Int(1))
1025 );
1026 assert_eq!(
1027 schema
1028 .attributes
1029 .iter()
1030 .find(|attribute| attribute.name == "noop_with_empty_axes")
1031 .unwrap()
1032 .default,
1033 Some(AttributeDefault::Int(0))
1034 );
1035 }
1036
1037 let rms = registry.lookup("RMSNormalization", "", 24).unwrap();
1038 assert_eq!(rms.since_version, 23);
1039 assert_eq!(
1040 rms.inputs
1041 .iter()
1042 .map(|input| input.type_str.as_str())
1043 .collect::<Vec<_>>(),
1044 ["T", "V"]
1045 );
1046 assert_eq!(rms.outputs[0].type_str, "V");
1047 assert_eq!(rms.type_constraints.len(), 2);
1048
1049 for name in ["ArgMax", "ArgMin"] {
1050 let schema = registry.lookup(name, "", 24).unwrap();
1051 assert_eq!(schema.since_version, 13);
1052 assert_eq!(schema.outputs[0].type_str, "tensor(int64)");
1053 assert_eq!(schema.attributes.len(), 3);
1054 }
1055 let log_softmax = registry.lookup("LogSoftmax", "", 24).unwrap();
1056 assert_eq!(log_softmax.since_version, 13);
1057 assert_eq!(
1058 log_softmax.attributes[0].default,
1059 Some(AttributeDefault::Int(-1))
1060 );
1061 }
1062
1063 #[test]
1064 fn softmax_schema_matches_opset_13() {
1065 let schema = SchemaRegistry::builtins()
1066 .lookup("Softmax", "", 24)
1067 .unwrap()
1068 .clone();
1069 assert_eq!(
1070 (
1071 schema.since_version,
1072 schema.inputs.len(),
1073 schema.outputs.len()
1074 ),
1075 (13, 1, 1)
1076 );
1077 assert_eq!(
1078 schema.attributes[0].default,
1079 Some(AttributeDefault::Int(-1))
1080 );
1081 }
1082
1083 #[test]
1084 fn layer_normalization_schema_matches_opset_17() {
1085 let registry = SchemaRegistry::builtins();
1086 let schema = registry.lookup("LayerNormalization", "", 24).unwrap();
1087 assert_eq!(
1088 (
1089 schema.since_version,
1090 schema.inputs.len(),
1091 schema.outputs.len()
1092 ),
1093 (17, 3, 3)
1094 );
1095 assert!(schema.inputs[2].optional);
1096 assert!(schema.outputs[1].optional && schema.outputs[2].optional);
1097 assert_eq!(schema.type_constraints.len(), 2);
1098 }
1099
1100 #[test]
1101 fn gather_schema_matches_opset_13() {
1102 let registry = SchemaRegistry::builtins();
1103 let schema = registry.lookup("Gather", "", 24).unwrap();
1104 assert_eq!(
1105 (
1106 schema.since_version,
1107 schema.inputs.len(),
1108 schema.outputs.len()
1109 ),
1110 (13, 2, 1)
1111 );
1112 assert_eq!(schema.inputs[1].type_str, "Tind");
1113 assert_eq!(
1114 schema.type_constraints[1].allowed,
1115 [DataType::Int32, DataType::Int64]
1116 );
1117 }
1118
1119 #[test]
1120 fn reshape_schema_matches_opset_24() {
1121 let registry = SchemaRegistry::builtins();
1122 assert_eq!(
1123 registry.lookup("Reshape", "", 18).unwrap().since_version,
1124 14
1125 );
1126 assert_eq!(
1127 registry.lookup("Reshape", "", 19).unwrap().since_version,
1128 19
1129 );
1130 assert_eq!(
1131 registry.lookup("Reshape", "", 21).unwrap().since_version,
1132 21
1133 );
1134 assert_eq!(
1135 registry.lookup("Reshape", "", 23).unwrap().since_version,
1136 23
1137 );
1138 let schema = registry.lookup("Reshape", "", 24).unwrap();
1139 assert_eq!(
1140 (
1141 schema.since_version,
1142 schema.inputs.len(),
1143 schema.outputs.len()
1144 ),
1145 (24, 2, 1)
1146 );
1147 assert_eq!(schema.inputs[1].type_str, "tensor(int64)");
1148 assert_eq!(schema.attributes[0].name, "allowzero");
1149 }
1150
1151 #[test]
1152 fn transpose_schema_matches_opset_24() {
1153 let registry = SchemaRegistry::builtins();
1154 assert_eq!(
1155 registry.lookup("Transpose", "", 20).unwrap().since_version,
1156 13
1157 );
1158 assert_eq!(
1159 registry.lookup("Transpose", "", 21).unwrap().since_version,
1160 21
1161 );
1162 assert_eq!(
1163 registry.lookup("Transpose", "", 23).unwrap().since_version,
1164 23
1165 );
1166 let schema = registry.lookup("Transpose", "", 24).unwrap();
1167 assert_eq!(
1168 (
1169 schema.since_version,
1170 schema.inputs.len(),
1171 schema.outputs.len()
1172 ),
1173 (24, 1, 1)
1174 );
1175 assert_eq!(schema.attributes[0].attr_type, AttributeType::Ints);
1176 }
1177
1178 #[test]
1179 fn concat_schema_matches_opset_13() {
1180 let registry = SchemaRegistry::builtins();
1181 let schema = registry.lookup("Concat", "", 24).unwrap();
1182 assert_eq!(
1183 (
1184 schema.since_version,
1185 schema.inputs.len(),
1186 schema.outputs.len()
1187 ),
1188 (13, 1, 1)
1189 );
1190 assert!(schema.inputs[0].variadic);
1191 assert_eq!(schema.inputs[0].min_arity, 1);
1192 assert!(schema.attributes[0].required);
1193 }
1194
1195 #[test]
1196 fn slice_schema_matches_opset_13() {
1197 let registry = SchemaRegistry::builtins();
1198 let schema = registry.lookup("Slice", "", 24).unwrap();
1199 assert_eq!(
1200 (
1201 schema.since_version,
1202 schema.inputs.len(),
1203 schema.outputs.len()
1204 ),
1205 (13, 5, 1)
1206 );
1207 assert!(!schema.inputs[2].optional);
1208 assert!(schema.inputs[3].optional && schema.inputs[4].optional);
1209 }
1210
1211 macro_rules! common_schema_test {
1212 ($test:ident, $name:literal, $since:literal, $inputs:literal, $outputs:literal) => {
1213 #[test]
1214 fn $test() {
1215 let registry = SchemaRegistry::builtins();
1216 let schema = registry.lookup($name, "", 25).unwrap();
1217 assert_eq!(schema.since_version, $since);
1218 assert_eq!(schema.inputs.len(), $inputs);
1219 assert_eq!(schema.outputs.len(), $outputs);
1220 }
1221 };
1222 }
1223
1224 common_schema_test!(sigmoid_schema_matches_opset_13, "Sigmoid", 13, 1, 1);
1225 common_schema_test!(tanh_schema_matches_opset_13, "Tanh", 13, 1, 1);
1226 common_schema_test!(erf_schema_matches_opset_13, "Erf", 13, 1, 1);
1227 common_schema_test!(sqrt_schema_matches_opset_13, "Sqrt", 13, 1, 1);
1228 common_schema_test!(exp_schema_matches_opset_13, "Exp", 13, 1, 1);
1229 common_schema_test!(log_schema_matches_opset_13, "Log", 13, 1, 1);
1230 common_schema_test!(pow_schema_matches_opset_15, "Pow", 15, 2, 1);
1231 common_schema_test!(clip_schema_matches_opset_13, "Clip", 13, 3, 1);
1232 common_schema_test!(expand_schema_matches_opset_13, "Expand", 13, 2, 1);
1233 common_schema_test!(where_schema_matches_opset_16, "Where", 16, 3, 1);
1234 common_schema_test!(reduce_sum_schema_matches_opset_13, "ReduceSum", 13, 2, 1);
1235 common_schema_test!(reduce_mean_schema_matches_opset_18, "ReduceMean", 18, 2, 1);
1236 common_schema_test!(sub_schema_matches_opset_14, "Sub", 14, 2, 1);
1237 common_schema_test!(div_schema_matches_opset_14, "Div", 14, 2, 1);
1238 common_schema_test!(neg_schema_matches_opset_13, "Neg", 13, 1, 1);
1239 common_schema_test!(abs_schema_matches_opset_13, "Abs", 13, 1, 1);
1240 common_schema_test!(mod_schema_matches_opset_13, "Mod", 13, 2, 1);
1241
1242 #[test]
1243 fn added_schema_details_match_onnx_v1_20() {
1244 let registry = SchemaRegistry::builtins();
1245
1246 let pow = registry.lookup("Pow", "", 25).unwrap();
1247 assert_eq!(pow.type_constraints[0].type_param, "T");
1248 assert_eq!(pow.type_constraints[1].type_param, "T1");
1249 assert_eq!(pow.type_constraints[0].allowed.len(), 6);
1250 assert_eq!(pow.type_constraints[1].allowed.len(), 12);
1251
1252 let clip = registry.lookup("Clip", "", 25).unwrap();
1253 assert!(clip.inputs[1].optional && clip.inputs[2].optional);
1254 assert_eq!(clip.type_constraints[0].allowed.len(), 12);
1255
1256 let expand = registry.lookup("Expand", "", 25).unwrap();
1257 assert_eq!(expand.inputs[1].type_str, "tensor(int64)");
1258 assert_eq!(expand.type_constraints[0].allowed.len(), 16);
1259
1260 let where_op = registry.lookup("Where", "", 25).unwrap();
1261 assert_eq!(where_op.type_constraints[0].allowed, [DataType::Bool]);
1262 assert_eq!(where_op.type_constraints[1].allowed.len(), 16);
1263
1264 for name in ["ReduceSum", "ReduceMean"] {
1265 let reduce = registry.lookup(name, "", 25).unwrap();
1266 assert!(reduce.inputs[1].optional);
1267 assert_eq!(reduce.attributes[0].default, Some(AttributeDefault::Int(1)));
1268 assert_eq!(reduce.attributes[1].default, Some(AttributeDefault::Int(0)));
1269 assert_eq!(reduce.type_constraints[0].allowed.len(), 8);
1270 }
1271
1272 for name in ["Sub", "Div", "Abs", "Mod"] {
1273 assert_eq!(
1274 registry.lookup(name, "", 25).unwrap().type_constraints[0]
1275 .allowed
1276 .len(),
1277 12,
1278 "{name}"
1279 );
1280 }
1281 assert_eq!(
1282 registry.lookup("Neg", "", 25).unwrap().type_constraints[0]
1283 .allowed
1284 .len(),
1285 8
1286 );
1287 assert_eq!(
1288 registry.lookup("Mod", "", 25).unwrap().attributes[0].default,
1289 Some(AttributeDefault::Int(0))
1290 );
1291 }
1292}