Skip to main content

onnx_std/schema/
mod.rs

1//! Declarative operator schemas and an opset-aware registry (ONNX_RS §7).
2//!
3//! Schemas are authored as YAML and loaded into owned Rust values. The built-in
4//! registry embeds high-value standard operators and can expand the YAML
5//! catalogue without changing the registry API.
6
7use std::collections::HashMap;
8
9use onnx_runtime_ir::DataType;
10use serde::{Deserialize, Serialize};
11
12/// A complete operator definition for one opset interval.
13#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
14pub struct OpSchema {
15    /// Operator domain (`""` and `"ai.onnx"` both mean the standard domain).
16    #[serde(default)]
17    pub domain: String,
18    /// Operator type name.
19    pub name: String,
20    /// First opset version for which this schema is valid.
21    pub since_version: u64,
22    /// Last valid opset version, inclusive. `None` means no upper bound.
23    #[serde(default)]
24    pub until_version: Option<u64>,
25    /// Human-readable operator documentation.
26    #[serde(default)]
27    pub doc: String,
28    /// Positional input definitions.
29    #[serde(default)]
30    pub inputs: Vec<InputSpec>,
31    /// Positional output definitions.
32    #[serde(default)]
33    pub outputs: Vec<OutputSpec>,
34    /// Attribute definitions.
35    #[serde(default)]
36    pub attributes: Vec<AttributeSpec>,
37    /// Type variables and the element types they admit.
38    #[serde(default)]
39    pub type_constraints: Vec<TypeConstraint>,
40}
41
42impl OpSchema {
43    /// Whether this schema applies to `opset`.
44    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/// One positional operator input.
50#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
51pub struct InputSpec {
52    /// Schema-visible input name.
53    pub name: String,
54    /// Type variable or concrete type expression.
55    pub type_str: String,
56    /// Human-readable input documentation.
57    #[serde(default)]
58    pub doc: String,
59    /// Whether this position may be omitted.
60    #[serde(default)]
61    pub optional: bool,
62    /// Whether this position accepts a variable number of trailing values.
63    #[serde(default)]
64    pub variadic: bool,
65    /// Minimum number of actual values consumed by a variadic position.
66    #[serde(default = "default_min_arity")]
67    pub min_arity: usize,
68}
69
70/// One positional operator output.
71#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
72pub struct OutputSpec {
73    /// Schema-visible output name.
74    pub name: String,
75    /// Type variable or concrete type expression.
76    pub type_str: String,
77    /// Human-readable output documentation.
78    #[serde(default)]
79    pub doc: String,
80    /// Whether this output may be omitted.
81    #[serde(default)]
82    pub optional: bool,
83    /// Whether this position accepts a variable number of trailing values.
84    #[serde(default)]
85    pub variadic: bool,
86    /// Minimum number of actual values produced by a variadic position.
87    #[serde(default = "default_min_arity")]
88    pub min_arity: usize,
89}
90
91const fn default_min_arity() -> usize {
92    1
93}
94
95/// ONNX attribute kinds.
96#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
97#[serde(rename_all = "snake_case")]
98pub enum AttributeType {
99    /// Scalar 64-bit integer.
100    Int,
101    /// Scalar 32-bit float.
102    Float,
103    /// Raw byte string.
104    String,
105    /// Tensor value.
106    Tensor,
107    /// Graph value.
108    Graph,
109    /// Sparse tensor value.
110    SparseTensor,
111    /// Type-proto value.
112    TypeProto,
113    /// Integer list.
114    Ints,
115    /// Float list.
116    Floats,
117    /// String list.
118    Strings,
119    /// Graph list.
120    Graphs,
121    /// Tensor list.
122    Tensors,
123    /// Sparse tensor list.
124    SparseTensors,
125    /// Type-proto list.
126    TypeProtos,
127}
128
129/// A typed YAML-compatible attribute default.
130#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
131#[serde(untagged)]
132pub enum AttributeDefault {
133    /// Integer scalar.
134    Int(i64),
135    /// Floating-point scalar.
136    Float(f64),
137    /// String scalar.
138    String(String),
139    /// Integer list.
140    Ints(Vec<i64>),
141    /// Floating-point list.
142    Floats(Vec<f64>),
143    /// String list.
144    Strings(Vec<String>),
145}
146
147/// One operator attribute definition.
148#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
149pub struct AttributeSpec {
150    /// Attribute name.
151    pub name: String,
152    /// Required ONNX attribute kind.
153    #[serde(rename = "type")]
154    pub attr_type: AttributeType,
155    /// Whether callers must provide the attribute.
156    #[serde(default)]
157    pub required: bool,
158    /// Default used when the attribute is omitted.
159    #[serde(default)]
160    pub default: Option<AttributeDefault>,
161    /// Human-readable attribute documentation.
162    #[serde(default)]
163    pub doc: String,
164}
165
166/// Allowed element types for a type variable.
167#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
168pub struct TypeConstraint {
169    /// Type variable name, such as `T`.
170    pub type_param: String,
171    /// Allowed shared-IR data types.
172    #[serde(with = "data_types")]
173    pub allowed: Vec<DataType>,
174}
175
176/// Failure while loading or registering schemas.
177#[derive(Debug, thiserror::Error)]
178pub enum SchemaError {
179    /// YAML could not be decoded.
180    #[error("invalid op-schema YAML: {0}")]
181    Yaml(#[from] serde_yaml::Error),
182    /// The schema declares an invalid or overlapping version interval.
183    #[error("invalid schema {domain}::{name}: {message}")]
184    Invalid {
185        /// Operator domain.
186        domain: String,
187        /// Operator name.
188        name: String,
189        /// Explanation of the invalid schema.
190        message: String,
191    },
192}
193
194/// Owned registry resolving `(op_type, domain, opset)` to an operator schema.
195#[derive(Clone, Debug, Default)]
196pub struct SchemaRegistry {
197    schemas: HashMap<(String, String), Vec<OpSchema>>,
198}
199
200impl SchemaRegistry {
201    /// Create an empty registry.
202    pub fn new() -> Self {
203        Self::default()
204    }
205
206    /// Load one schema from YAML and register it.
207    pub fn load_yaml(&mut self, yaml: &str) -> Result<(), SchemaError> {
208        self.register(serde_yaml::from_str(yaml)?)
209    }
210
211    /// Register a schema. A later `since_version` supersedes earlier versions.
212    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    /// Resolve the schema whose interval contains `opset`.
236    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    /// Whether any version of this operator is registered.
245    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    /// Iterate over every registered schema.
251    pub fn iter(&self) -> impl Iterator<Item = &OpSchema> {
252        self.schemas.values().flatten()
253    }
254
255    /// Load the embedded standard-schema starter set.
256    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
440// FOLLOW-UP §7.4: complete the standard and ONNX-ML YAML catalogues.
441
442mod 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}