Skip to main content

virtio_accel_tosa_build/
lib.rs

1//! Safe authoring of static, single-block TOSA 1.0 artifacts.
2//!
3//! The raw FlatBuffers layout is private to this crate. Borrowed [`Graph`] definitions suit
4//! statically declared artifacts, while [`OwnedGraph`] supports frontends that discover names,
5//! shapes, constants, and operators incrementally. Both surfaces round-trip every artifact through
6//! `virtio-accel-tosa`'s bounded parser and complete target validator before returning owned bytes.
7
8#![no_std]
9#![forbid(unsafe_code)]
10
11extern crate alloc;
12
13use alloc::borrow::Cow;
14use alloc::collections::BTreeSet;
15use alloc::string::{String, ToString};
16use alloc::vec::Vec;
17use core::fmt;
18
19use flatbuffers::{FlatBufferBuilder, TableFinishedWIPOffset, WIPOffset};
20use virtio_accel_tosa::{DType, NanPropagationMode, Op, RoundingMode, SemanticError, Target};
21
22type Table = WIPOffset<TableFinishedWIPOffset>;
23
24// These are the six stable container-table vtable positions used by the
25// single-block surface. Keeping them here prevents consumers from duplicating
26// schema layout knowledge; the completed artifact is always verified by the
27// ingestion crate before it crosses this API boundary.
28mod slot {
29    pub const TENSOR_NAME: u16 = 4;
30    pub const TENSOR_SHAPE: u16 = 6;
31    pub const TENSOR_TYPE: u16 = 8;
32    pub const TENSOR_DATA: u16 = 10;
33
34    pub const SHAPE_NAME: u16 = 4;
35    pub const SHAPE_RANK: u16 = 6;
36    pub const SHAPE_DATA: u16 = 8;
37
38    pub const OPERATOR_OP: u16 = 4;
39    pub const OPERATOR_ATTRIBUTE_TYPE: u16 = 6;
40    pub const OPERATOR_ATTRIBUTE: u16 = 8;
41    pub const OPERATOR_INPUTS: u16 = 10;
42    pub const OPERATOR_OUTPUTS: u16 = 12;
43
44    pub const BLOCK_NAME: u16 = 4;
45    pub const BLOCK_OPERATORS: u16 = 6;
46    pub const BLOCK_TENSORS: u16 = 8;
47    pub const BLOCK_INPUTS: u16 = 10;
48    pub const BLOCK_OUTPUTS: u16 = 12;
49    pub const BLOCK_SHAPES: u16 = 14;
50
51    pub const REGION_NAME: u16 = 4;
52    pub const REGION_BLOCKS: u16 = 6;
53
54    pub const GRAPH_VERSION: u16 = 4;
55    pub const GRAPH_REGIONS: u16 = 6;
56
57    pub const VERSION_MAJOR: u16 = 4;
58    pub const VERSION_MINOR: u16 = 6;
59    pub const VERSION_PATCH: u16 = 8;
60    pub const VERSION_DRAFT: u16 = 10;
61
62    pub const AXIS: u16 = 4;
63    pub const AXIS_NAN_MODE: u16 = 6;
64
65    pub const CLAMP_MIN_VAL: u16 = 4;
66    pub const CLAMP_MAX_VAL: u16 = 6;
67    pub const CLAMP_NAN_MODE: u16 = 8;
68
69    pub const TRANSPOSE_PERMS: u16 = 4;
70
71    pub const NAN_MODE: u16 = 4;
72    pub const MAX_POOL_KERNEL: u16 = 4;
73    pub const MAX_POOL_STRIDE: u16 = 6;
74    pub const MAX_POOL_PAD: u16 = 8;
75    pub const MAX_POOL_NAN_MODE: u16 = 10;
76
77    pub const RESCALE_SCALE32: u16 = 4;
78    pub const RESCALE_ROUNDING_MODE: u16 = 6;
79    pub const RESCALE_PER_CHANNEL: u16 = 8;
80    pub const RESCALE_INPUT_UNSIGNED: u16 = 10;
81    pub const RESCALE_OUTPUT_UNSIGNED: u16 = 12;
82}
83
84/// A serialized compile-time shape value.
85///
86/// Shape values occupy the TOSA shape namespace rather than the tensor namespace. They must be
87/// produced by [`OperatorKind::ConstShape`] before an operator such as `RESHAPE` consumes them.
88#[derive(Clone, Copy, Debug, PartialEq, Eq)]
89pub struct Shape<'a> {
90    /// Unique block-local shape name.
91    pub name: &'a str,
92    /// Signed 64-bit TOSA shape components. A scalar target shape uses an empty slice.
93    pub values: &'a [i64],
94}
95
96impl<'a> Shape<'a> {
97    pub const fn new(name: &'a str, values: &'a [i64]) -> Self {
98        Self { name, values }
99    }
100}
101
102/// A statically shaped tensor definition.
103#[derive(Clone, Copy, Debug, PartialEq, Eq)]
104pub struct Tensor<'a> {
105    /// Unique block-local tensor name.
106    pub name: &'a str,
107    /// Concrete dimensions. Scalars use an empty slice.
108    pub shape: &'a [i32],
109    /// Stable TOSA 1.0 element type.
110    pub dtype: DType,
111    /// Inline constant bytes, when this tensor is produced by `CONST`.
112    pub data: Option<&'a [u8]>,
113}
114
115impl<'a> Tensor<'a> {
116    /// Define a non-constant tensor.
117    pub const fn new(name: &'a str, shape: &'a [i32], dtype: DType) -> Self {
118        Self {
119            name,
120            shape,
121            dtype,
122            data: None,
123        }
124    }
125
126    /// Define a tensor with inline constant storage.
127    pub const fn constant(name: &'a str, shape: &'a [i32], dtype: DType, data: &'a [u8]) -> Self {
128        Self {
129            name,
130            shape,
131            dtype,
132            data: Some(data),
133        }
134    }
135}
136
137/// Typed operator kinds supported by the initial authoring surface.
138///
139/// Variants with no fields still name their exact TOSA attribute table. This
140/// prevents a caller from pairing an opcode with the wrong union member.
141/// Bytes of inline storage for one `CLAMP` bound: enough for every TOSA scalar dtype. The
142/// bounds are one scalar each in the operand's own dtype, so they are carried as serialized
143/// bytes; a wider numeric type could not represent an integer or FP8 bound faithfully.
144pub const MAX_CLAMP_BOUND_BYTES: usize = 8;
145
146/// Permutation entries a `TRANSPOSE` carries inline, one per dimension.
147pub const MAX_TRANSPOSE_RANK: usize = 6;
148
149#[derive(Clone, Copy, Debug, PartialEq, Eq)]
150#[non_exhaustive]
151pub enum OperatorKind {
152    MatMul,
153    ArgMax {
154        axis: i32,
155        nan_mode: NanPropagationMode,
156    },
157    Ceil,
158    Floor,
159    /// The bounds are one scalar each in the operand's own dtype, so they are carried as their
160    /// serialized bytes (`bound_bytes` of each array is significant) rather than a wider type
161    /// that would not represent an integer or FP8 bound faithfully.
162    Clamp {
163        min_val: [u8; MAX_CLAMP_BOUND_BYTES],
164        max_val: [u8; MAX_CLAMP_BOUND_BYTES],
165        bound_bytes: u8,
166        nan_mode: NanPropagationMode,
167    },
168    Concat {
169        axis: i32,
170    },
171    Reverse {
172        axis: i32,
173    },
174    /// One permutation entry per dimension; `rank` of `perms` is significant.
175    Transpose {
176        perms: [i32; MAX_TRANSPOSE_RANK],
177        rank: u8,
178    },
179    ReduceMax {
180        axis: i32,
181        nan_mode: NanPropagationMode,
182    },
183    ReduceMin {
184        axis: i32,
185        nan_mode: NanPropagationMode,
186    },
187    ReduceProduct {
188        axis: i32,
189    },
190    ReduceSum {
191        axis: i32,
192    },
193    MaxPool2d {
194        kernel: [i32; 2],
195        stride: [i32; 2],
196        pad: [i32; 4],
197        nan_mode: NanPropagationMode,
198    },
199    Sigmoid,
200    Tanh,
201    Add,
202    LogicalAnd,
203    LogicalOr,
204    LogicalXor,
205    Maximum {
206        nan_mode: NanPropagationMode,
207    },
208    Minimum {
209        nan_mode: NanPropagationMode,
210    },
211    Mul,
212    Pow,
213    Sub,
214    Abs,
215    Cos,
216    Exp,
217    Log,
218    LogicalNot,
219    Negate,
220    Reciprocal,
221    Rsqrt,
222    Sin,
223    Select,
224    Equal,
225    Greater,
226    GreaterEqual,
227    Erf,
228    Reshape,
229    Cast,
230    Rescale {
231        scale32: bool,
232        rounding_mode: RoundingMode,
233        per_channel: bool,
234        input_unsigned: bool,
235        output_unsigned: bool,
236    },
237    Const,
238    Identity,
239    ConstShape,
240}
241
242impl OperatorKind {
243    /// Stable TOSA opcode selected by this typed variant.
244    pub const fn op(self) -> Op {
245        match self {
246            Self::MatMul => Op::MATMUL,
247            Self::ArgMax { .. } => Op::ARGMAX,
248            Self::Ceil => Op::CEIL,
249            Self::Floor => Op::FLOOR,
250            Self::Clamp { .. } => Op::CLAMP,
251            Self::Concat { .. } => Op::CONCAT,
252            Self::Reverse { .. } => Op::REVERSE,
253            Self::Transpose { .. } => Op::TRANSPOSE,
254            Self::ReduceMax { .. } => Op::REDUCE_MAX,
255            Self::ReduceMin { .. } => Op::REDUCE_MIN,
256            Self::ReduceProduct { .. } => Op::REDUCE_PRODUCT,
257            Self::ReduceSum { .. } => Op::REDUCE_SUM,
258            Self::MaxPool2d { .. } => Op::MAX_POOL2D,
259            Self::Sigmoid => Op::SIGMOID,
260            Self::Tanh => Op::TANH,
261            Self::Add => Op::ADD,
262            Self::LogicalAnd => Op::LOGICAL_AND,
263            Self::LogicalOr => Op::LOGICAL_OR,
264            Self::LogicalXor => Op::LOGICAL_XOR,
265            Self::Maximum { .. } => Op::MAXIMUM,
266            Self::Minimum { .. } => Op::MINIMUM,
267            Self::Mul => Op::MUL,
268            Self::Pow => Op::POW,
269            Self::Sub => Op::SUB,
270            Self::Abs => Op::ABS,
271            Self::Cos => Op::COS,
272            Self::Exp => Op::EXP,
273            Self::Log => Op::LOG,
274            Self::LogicalNot => Op::LOGICAL_NOT,
275            Self::Negate => Op::NEGATE,
276            Self::Reciprocal => Op::RECIPROCAL,
277            Self::Rsqrt => Op::RSQRT,
278            Self::Sin => Op::SIN,
279            Self::Select => Op::SELECT,
280            Self::Equal => Op::EQUAL,
281            Self::Greater => Op::GREATER,
282            Self::GreaterEqual => Op::GREATER_EQUAL,
283            Self::Erf => Op::ERF,
284            Self::Reshape => Op::RESHAPE,
285            Self::Cast => Op::CAST,
286            Self::Rescale { .. } => Op::RESCALE,
287            Self::Const => Op::CONST,
288            Self::Identity => Op::IDENTITY,
289            Self::ConstShape => Op::CONST_SHAPE,
290        }
291    }
292}
293
294/// One operator and its block-local tensor references.
295#[derive(Clone, Copy, Debug, PartialEq, Eq)]
296pub struct Operator<'a> {
297    pub kind: OperatorKind,
298    pub inputs: &'a [&'a str],
299    pub outputs: &'a [&'a str],
300}
301
302impl<'a> Operator<'a> {
303    pub const fn new(kind: OperatorKind, inputs: &'a [&'a str], outputs: &'a [&'a str]) -> Self {
304        Self {
305            kind,
306            inputs,
307            outputs,
308        }
309    }
310}
311
312/// An owned compile-time shape for incrementally assembled graphs.
313#[derive(Clone, Debug, PartialEq, Eq)]
314pub struct OwnedShape {
315    pub name: String,
316    pub values: Vec<i64>,
317}
318
319impl OwnedShape {
320    pub fn new(name: impl Into<String>, values: Vec<i64>) -> Self {
321        Self {
322            name: name.into(),
323            values,
324        }
325    }
326}
327
328impl From<Shape<'_>> for OwnedShape {
329    fn from(shape: Shape<'_>) -> Self {
330        Self::new(shape.name, shape.values.to_vec())
331    }
332}
333
334/// An owned tensor definition for incrementally assembled graphs.
335#[derive(Clone, Debug, PartialEq, Eq)]
336pub struct OwnedTensor<'a> {
337    pub name: String,
338    pub shape: Vec<i32>,
339    pub dtype: DType,
340    /// Inline constant bytes. Existing frontend storage may be borrowed to avoid a second copy;
341    /// newly transformed data can be moved into this field.
342    pub data: Option<Cow<'a, [u8]>>,
343}
344
345impl<'a> OwnedTensor<'a> {
346    /// Define a non-constant owned tensor.
347    pub fn new(name: impl Into<String>, shape: Vec<i32>, dtype: DType) -> Self {
348        Self {
349            name: name.into(),
350            shape,
351            dtype,
352            data: None,
353        }
354    }
355
356    /// Define an owned tensor with inline constant storage.
357    pub fn constant(
358        name: impl Into<String>,
359        shape: Vec<i32>,
360        dtype: DType,
361        data: impl Into<Cow<'a, [u8]>>,
362    ) -> Self {
363        Self {
364            name: name.into(),
365            shape,
366            dtype,
367            data: Some(data.into()),
368        }
369    }
370}
371
372impl<'a> From<Tensor<'a>> for OwnedTensor<'a> {
373    fn from(tensor: Tensor<'a>) -> Self {
374        match tensor.data {
375            Some(data) => Self::constant(tensor.name, tensor.shape.to_vec(), tensor.dtype, data),
376            None => Self::new(tensor.name, tensor.shape.to_vec(), tensor.dtype),
377        }
378    }
379}
380
381/// An owned typed operator for incrementally assembled graphs.
382#[derive(Clone, Debug, PartialEq, Eq)]
383pub struct OwnedOperator {
384    pub kind: OperatorKind,
385    pub inputs: Vec<String>,
386    pub outputs: Vec<String>,
387}
388
389impl OwnedOperator {
390    /// Define an operator from owned operand names.
391    pub const fn new(kind: OperatorKind, inputs: Vec<String>, outputs: Vec<String>) -> Self {
392        Self {
393            kind,
394            inputs,
395            outputs,
396        }
397    }
398}
399
400impl From<Operator<'_>> for OwnedOperator {
401    fn from(operator: Operator<'_>) -> Self {
402        Self::new(
403            operator.kind,
404            operator
405                .inputs
406                .iter()
407                .map(|name| (*name).to_string())
408                .collect(),
409            operator
410                .outputs
411                .iter()
412                .map(|name| (*name).to_string())
413                .collect(),
414        )
415    }
416}
417
418/// An owned, incrementally assembled static TOSA graph.
419///
420/// This surface is intended for compiler frontends whose names, shapes, constants, and operator
421/// lists are discovered at runtime. It owns names and metadata while permitting explicit borrowing
422/// of existing constant storage. It creates borrowed [`Graph`] views only during [`Self::build`],
423/// so callers do not need parallel adapter structures. The same structural and target validation
424/// as [`Graph::build`] remains authoritative.
425#[derive(Clone, Debug, PartialEq, Eq)]
426pub struct OwnedGraph<'a> {
427    pub name: String,
428    pub tensors: Vec<OwnedTensor<'a>>,
429    pub shapes: Vec<OwnedShape>,
430    pub operators: Vec<OwnedOperator>,
431    pub inputs: Vec<String>,
432    pub outputs: Vec<String>,
433}
434
435impl<'a> OwnedGraph<'a> {
436    pub fn new(name: impl Into<String>) -> Self {
437        Self {
438            name: name.into(),
439            tensors: Vec::new(),
440            shapes: Vec::new(),
441            operators: Vec::new(),
442            inputs: Vec::new(),
443            outputs: Vec::new(),
444        }
445    }
446
447    pub fn push_tensor(&mut self, tensor: OwnedTensor<'a>) -> &mut Self {
448        self.tensors.push(tensor);
449        self
450    }
451
452    pub fn push_shape(&mut self, shape: OwnedShape) -> &mut Self {
453        self.shapes.push(shape);
454        self
455    }
456
457    pub fn push_operator(&mut self, operator: OwnedOperator) -> &mut Self {
458        self.operators.push(operator);
459        self
460    }
461
462    pub fn push_input(&mut self, name: impl Into<String>) -> &mut Self {
463        self.inputs.push(name.into());
464        self
465    }
466
467    pub fn push_output(&mut self, name: impl Into<String>) -> &mut Self {
468        self.outputs.push(name.into());
469        self
470    }
471
472    /// Serialize and semantically validate this graph for `target`.
473    ///
474    /// The graph remains reusable after a build. All metadata views allocated here are transient;
475    /// the returned artifact bytes are the only allocation retained by the caller.
476    pub fn build(&self, target: Target) -> Result<Vec<u8>, BuildError> {
477        let tensors: Vec<_> = self
478            .tensors
479            .iter()
480            .map(|tensor| match &tensor.data {
481                Some(data) => {
482                    Tensor::constant(&tensor.name, &tensor.shape, tensor.dtype, data.as_ref())
483                }
484                None => Tensor::new(&tensor.name, &tensor.shape, tensor.dtype),
485            })
486            .collect();
487        let shapes: Vec<_> = self
488            .shapes
489            .iter()
490            .map(|shape| Shape::new(&shape.name, &shape.values))
491            .collect();
492        let operator_inputs: Vec<Vec<&str>> = self
493            .operators
494            .iter()
495            .map(|operator| operator.inputs.iter().map(String::as_str).collect())
496            .collect();
497        let operator_outputs: Vec<Vec<&str>> = self
498            .operators
499            .iter()
500            .map(|operator| operator.outputs.iter().map(String::as_str).collect())
501            .collect();
502        let operators: Vec<_> = self
503            .operators
504            .iter()
505            .enumerate()
506            .map(|(index, operator)| {
507                Operator::new(
508                    operator.kind,
509                    &operator_inputs[index],
510                    &operator_outputs[index],
511                )
512            })
513            .collect();
514        let inputs: Vec<_> = self.inputs.iter().map(String::as_str).collect();
515        let outputs: Vec<_> = self.outputs.iter().map(String::as_str).collect();
516
517        Graph::new(&self.name, &tensors, &operators, &inputs, &outputs)
518            .with_shapes(&shapes)
519            .build(target)
520    }
521}
522
523impl<'a> From<Graph<'a>> for OwnedGraph<'a> {
524    fn from(graph: Graph<'a>) -> Self {
525        Self {
526            name: graph.name.to_string(),
527            tensors: graph.tensors.iter().copied().map(Into::into).collect(),
528            shapes: graph.shapes.iter().copied().map(Into::into).collect(),
529            operators: graph.operators.iter().copied().map(Into::into).collect(),
530            inputs: graph
531                .inputs
532                .iter()
533                .map(|name| (*name).to_string())
534                .collect(),
535            outputs: graph
536                .outputs
537                .iter()
538                .map(|name| (*name).to_string())
539                .collect(),
540        }
541    }
542}
543
544/// A static TOSA graph containing one region and one basic block.
545#[derive(Clone, Copy, Debug, PartialEq, Eq)]
546pub struct Graph<'a> {
547    pub name: &'a str,
548    pub tensors: &'a [Tensor<'a>],
549    pub shapes: &'a [Shape<'a>],
550    pub operators: &'a [Operator<'a>],
551    pub inputs: &'a [&'a str],
552    pub outputs: &'a [&'a str],
553}
554
555impl<'a> Graph<'a> {
556    pub const fn new(
557        name: &'a str,
558        tensors: &'a [Tensor<'a>],
559        operators: &'a [Operator<'a>],
560        inputs: &'a [&'a str],
561        outputs: &'a [&'a str],
562    ) -> Self {
563        Self {
564            name,
565            tensors,
566            shapes: &[],
567            operators,
568            inputs,
569            outputs,
570        }
571    }
572
573    /// Add compile-time shape values to this graph.
574    pub const fn with_shapes(mut self, shapes: &'a [Shape<'a>]) -> Self {
575        self.shapes = shapes;
576        self
577    }
578
579    /// Serialize and semantically validate this graph for `target`.
580    ///
581    /// The returned vector is the only allocation retained by the caller.
582    /// Construction is linear in graph metadata and constant bytes. Validation
583    /// runs once on this cold authoring path; providers still perform their own
584    /// authoritative admission during `load_program`.
585    pub fn build(self, target: Target) -> Result<Vec<u8>, BuildError> {
586        self.check_static_surface()?;
587        let bytes = self.serialize();
588        let model = virtio_accel_tosa::parse(&bytes).map_err(BuildError::Parse)?;
589        model.validate_for(target).map_err(BuildError::Semantic)?;
590        Ok(bytes)
591    }
592
593    fn check_static_surface(self) -> Result<(), BuildError> {
594        if self.name.is_empty() {
595            return Err(BuildError::EmptyGraphName);
596        }
597        for tensor in self.tensors {
598            if tensor.shape.iter().any(|dimension| *dimension < 0) {
599                return Err(BuildError::DynamicShape);
600            }
601            if !tensor.dtype.is_tosa_1_0() {
602                return Err(BuildError::UnsupportedDType(tensor.dtype));
603            }
604        }
605
606        for shape in self.shapes {
607            if u32::try_from(shape.values.len()).is_err() {
608                return Err(BuildError::ShapeRankOverflow);
609            }
610        }
611
612        let mut constant_tensors = BTreeSet::new();
613        let mut constant_shapes = BTreeSet::new();
614        for operator in self.operators {
615            let outputs = match operator.kind {
616                OperatorKind::Const => &mut constant_tensors,
617                OperatorKind::ConstShape => &mut constant_shapes,
618                _ => continue,
619            };
620            outputs.extend(operator.outputs.iter().copied());
621        }
622
623        for tensor in self.tensors {
624            match tensor.data {
625                Some([]) => return Err(BuildError::EmptyConstantData),
626                Some(_) if !constant_tensors.contains(tensor.name) => {
627                    return Err(BuildError::TensorDataWithoutConst);
628                }
629                Some(_) => {}
630                None if constant_tensors.contains(tensor.name) => {
631                    return Err(BuildError::ConstWithoutTensorData);
632                }
633                None => {}
634            }
635        }
636        if self
637            .shapes
638            .iter()
639            .any(|shape| !constant_shapes.contains(shape.name))
640        {
641            return Err(BuildError::ShapeWithoutConstShape);
642        }
643        Ok(())
644    }
645
646    fn serialize(self) -> Vec<u8> {
647        let mut builder = FlatBufferBuilder::with_capacity(4096);
648
649        let tensors = {
650            let tables: Vec<_> = self
651                .tensors
652                .iter()
653                .map(|tensor| tensor_table(&mut builder, *tensor))
654                .collect();
655            builder.create_vector(&tables)
656        };
657        let shapes = {
658            let tables: Vec<_> = self
659                .shapes
660                .iter()
661                .map(|shape| shape_table(&mut builder, *shape))
662                .collect();
663            builder.create_vector(&tables)
664        };
665        let operators = {
666            let tables: Vec<_> = self
667                .operators
668                .iter()
669                .map(|operator| operator_table(&mut builder, *operator))
670                .collect();
671            builder.create_vector(&tables)
672        };
673        let block = {
674            let name = builder.create_string(self.name);
675            let inputs = string_vector(&mut builder, self.inputs);
676            let outputs = string_vector(&mut builder, self.outputs);
677            let table = builder.start_table();
678            builder.push_slot_always(slot::BLOCK_NAME, name);
679            builder.push_slot_always(slot::BLOCK_OPERATORS, operators);
680            builder.push_slot_always(slot::BLOCK_TENSORS, tensors);
681            builder.push_slot_always(slot::BLOCK_INPUTS, inputs);
682            builder.push_slot_always(slot::BLOCK_OUTPUTS, outputs);
683            builder.push_slot_always(slot::BLOCK_SHAPES, shapes);
684            builder.end_table(table)
685        };
686        let region = {
687            let name = builder.create_string(self.name);
688            let blocks = builder.create_vector(&[block]);
689            let table = builder.start_table();
690            builder.push_slot_always(slot::REGION_NAME, name);
691            builder.push_slot_always(slot::REGION_BLOCKS, blocks);
692            builder.end_table(table)
693        };
694        let graph = {
695            let version = {
696                let table = builder.start_table();
697                builder.push_slot::<i32>(slot::VERSION_MAJOR, 1, -1);
698                builder.push_slot::<i32>(slot::VERSION_MINOR, 0, -1);
699                builder.push_slot::<i32>(slot::VERSION_PATCH, 0, -1);
700                builder.push_slot::<bool>(slot::VERSION_DRAFT, false, true);
701                builder.end_table(table)
702            };
703            let regions = builder.create_vector(&[region]);
704            let table = builder.start_table();
705            builder.push_slot_always(slot::GRAPH_VERSION, version);
706            builder.push_slot_always(slot::GRAPH_REGIONS, regions);
707            builder.end_table(table)
708        };
709
710        builder.finish(graph, Some("TOSA"));
711        builder.finished_data().to_vec()
712    }
713}
714
715fn shape_table(builder: &mut FlatBufferBuilder<'_>, shape: Shape<'_>) -> Table {
716    let name = builder.create_string(shape.name);
717    let bytes: Vec<_> = shape
718        .values
719        .iter()
720        .flat_map(|value| value.to_le_bytes())
721        .collect();
722    let data = builder.create_vector(&bytes);
723    let table = builder.start_table();
724    builder.push_slot_always(slot::SHAPE_NAME, name);
725    builder.push_slot::<u32>(
726        slot::SHAPE_RANK,
727        u32::try_from(shape.values.len()).expect("shape rank checked before serialization"),
728        0,
729    );
730    builder.push_slot_always(slot::SHAPE_DATA, data);
731    builder.end_table(table)
732}
733
734fn tensor_table(builder: &mut FlatBufferBuilder<'_>, tensor: Tensor<'_>) -> Table {
735    let name = builder.create_string(tensor.name);
736    let shape = builder.create_vector(tensor.shape);
737    let data = tensor.data.map(|bytes| builder.create_vector(bytes));
738    let table = builder.start_table();
739    builder.push_slot_always(slot::TENSOR_NAME, name);
740    builder.push_slot_always(slot::TENSOR_SHAPE, shape);
741    builder.push_slot::<u32>(slot::TENSOR_TYPE, tensor.dtype.get(), 0);
742    if let Some(data) = data {
743        builder.push_slot_always(slot::TENSOR_DATA, data);
744    }
745    builder.end_table(table)
746}
747
748fn operator_table(builder: &mut FlatBufferBuilder<'_>, operator: Operator<'_>) -> Table {
749    let inputs = string_vector(builder, operator.inputs);
750    let outputs = string_vector(builder, operator.outputs);
751    let max_pool = match operator.kind {
752        OperatorKind::MaxPool2d {
753            kernel,
754            stride,
755            pad,
756            nan_mode,
757        } => Some((
758            builder.create_vector(&kernel),
759            builder.create_vector(&stride),
760            builder.create_vector(&pad),
761            nan_mode,
762        )),
763        _ => None,
764    };
765    let clamp = match operator.kind {
766        OperatorKind::Clamp {
767            min_val,
768            max_val,
769            bound_bytes,
770            nan_mode,
771        } => {
772            let bytes = bound_bytes as usize;
773            Some((
774                builder.create_vector(&min_val[..bytes]),
775                builder.create_vector(&max_val[..bytes]),
776                nan_mode,
777            ))
778        }
779        _ => None,
780    };
781    let transpose = match operator.kind {
782        OperatorKind::Transpose { perms, rank } => {
783            Some(builder.create_vector(&perms[..rank as usize]))
784        }
785        _ => None,
786    };
787    let attribute = {
788        let table = builder.start_table();
789        match operator.kind {
790            OperatorKind::Maximum { nan_mode } | OperatorKind::Minimum { nan_mode } => {
791                builder.push_slot::<u32>(slot::NAN_MODE, nan_mode.get(), 0);
792            }
793            OperatorKind::ArgMax { axis, nan_mode }
794            | OperatorKind::ReduceMax { axis, nan_mode }
795            | OperatorKind::ReduceMin { axis, nan_mode } => {
796                builder.push_slot::<i32>(slot::AXIS, axis, 0);
797                builder.push_slot::<u32>(slot::AXIS_NAN_MODE, nan_mode.get(), 0);
798            }
799            OperatorKind::Concat { axis }
800            | OperatorKind::Reverse { axis }
801            | OperatorKind::ReduceProduct { axis }
802            | OperatorKind::ReduceSum { axis } => {
803                builder.push_slot::<i32>(slot::AXIS, axis, 0);
804            }
805            OperatorKind::Clamp { .. } => {
806                let (min_val, max_val, nan_mode) = clamp.expect("CLAMP vectors were constructed");
807                builder.push_slot_always(slot::CLAMP_MIN_VAL, min_val);
808                builder.push_slot_always(slot::CLAMP_MAX_VAL, max_val);
809                builder.push_slot::<u32>(slot::CLAMP_NAN_MODE, nan_mode.get(), 0);
810            }
811            OperatorKind::Transpose { .. } => {
812                let perms = transpose.expect("TRANSPOSE perms were constructed");
813                builder.push_slot_always(slot::TRANSPOSE_PERMS, perms);
814            }
815            OperatorKind::MaxPool2d { .. } => {
816                let (kernel, stride, pad, nan_mode) =
817                    max_pool.expect("MAX_POOL2D vectors were constructed");
818                builder.push_slot_always(slot::MAX_POOL_KERNEL, kernel);
819                builder.push_slot_always(slot::MAX_POOL_STRIDE, stride);
820                builder.push_slot_always(slot::MAX_POOL_PAD, pad);
821                builder.push_slot::<u32>(slot::MAX_POOL_NAN_MODE, nan_mode.get(), 0);
822            }
823            OperatorKind::Rescale {
824                scale32,
825                rounding_mode,
826                per_channel,
827                input_unsigned,
828                output_unsigned,
829            } => {
830                builder.push_slot::<bool>(slot::RESCALE_SCALE32, scale32, false);
831                builder.push_slot::<u32>(slot::RESCALE_ROUNDING_MODE, rounding_mode.get(), 0);
832                builder.push_slot::<bool>(slot::RESCALE_PER_CHANNEL, per_channel, false);
833                builder.push_slot::<bool>(slot::RESCALE_INPUT_UNSIGNED, input_unsigned, false);
834                builder.push_slot::<bool>(slot::RESCALE_OUTPUT_UNSIGNED, output_unsigned, false);
835            }
836            _ => {}
837        }
838        builder.end_table(table)
839    };
840    let op = operator.kind.op();
841    let table = builder.start_table();
842    builder.push_slot::<u32>(slot::OPERATOR_OP, op.get(), 0);
843    // In the pinned TOSA 1.0 schema, stable Attribute union members are in
844    // the same order as their corresponding stable Op values.
845    builder.push_slot::<u8>(slot::OPERATOR_ATTRIBUTE_TYPE, op.get() as u8, 0);
846    builder.push_slot_always(slot::OPERATOR_ATTRIBUTE, attribute);
847    builder.push_slot_always(slot::OPERATOR_INPUTS, inputs);
848    builder.push_slot_always(slot::OPERATOR_OUTPUTS, outputs);
849    builder.end_table(table)
850}
851
852fn string_vector<'a>(
853    builder: &mut FlatBufferBuilder<'a>,
854    values: &[&str],
855) -> WIPOffset<flatbuffers::Vector<'a, flatbuffers::ForwardsUOffset<&'a str>>> {
856    let offsets: Vec<_> = values
857        .iter()
858        .map(|value| builder.create_string(value))
859        .collect();
860    builder.create_vector(&offsets)
861}
862
863/// Failure to construct a static graph accepted by the shared TOSA target.
864#[derive(Debug)]
865#[non_exhaustive]
866pub enum BuildError {
867    EmptyGraphName,
868    DynamicShape,
869    ShapeRankOverflow,
870    EmptyConstantData,
871    TensorDataWithoutConst,
872    ConstWithoutTensorData,
873    ShapeWithoutConstShape,
874    UnsupportedDType(DType),
875    Parse(virtio_accel_tosa::Error),
876    Semantic(SemanticError),
877}
878
879impl fmt::Display for BuildError {
880    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
881        match self {
882            Self::EmptyGraphName => formatter.write_str("graph name must not be empty"),
883            Self::DynamicShape => {
884                formatter.write_str("the single-block authoring surface requires static shapes")
885            }
886            Self::ShapeRankOverflow => formatter.write_str("shape rank does not fit TOSA u32"),
887            Self::EmptyConstantData => {
888                formatter.write_str("tensor constants require nonempty inline data")
889            }
890            Self::TensorDataWithoutConst => {
891                formatter.write_str("inline tensor data requires a CONST producer")
892            }
893            Self::ConstWithoutTensorData => {
894                formatter.write_str("CONST outputs require nonempty inline tensor data")
895            }
896            Self::ShapeWithoutConstShape => {
897                formatter.write_str("shape values require a CONST_SHAPE producer")
898            }
899            Self::UnsupportedDType(dtype) => {
900                write!(formatter, "unsupported TOSA 1.0 dtype {dtype:?}")
901            }
902            Self::Parse(error) => write!(
903                formatter,
904                "constructed artifact failed structural validation: {error}"
905            ),
906            Self::Semantic(error) => write!(
907                formatter,
908                "constructed artifact failed semantic validation: {error}"
909            ),
910        }
911    }
912}
913
914#[cfg(test)]
915mod tests {
916    use super::*;
917    use alloc::vec;
918    use virtio_accel_tosa::{ExtensionSet, Level, ProfileSet, Version};
919
920    const FLOAT_TARGET: Target = Target::new(
921        Version::TOSA_1_0,
922        ProfileSet::FLOATING_POINT,
923        Level::Level8K,
924        ExtensionSet::NONE,
925    );
926
927    #[test]
928    fn identity_round_trips_through_production_validation() {
929        let tensors = [
930            Tensor::new("input", &[1, 4], DType::FP32),
931            Tensor::new("output", &[1, 4], DType::FP32),
932        ];
933        let operators = [Operator::new(
934            OperatorKind::Identity,
935            &["input"],
936            &["output"],
937        )];
938        let graph = Graph::new("main", &tensors, &operators, &["input"], &["output"]);
939
940        let first = graph.build(FLOAT_TARGET).unwrap();
941        let second = graph.build(FLOAT_TARGET).unwrap();
942        assert_eq!(first, second);
943
944        let model = virtio_accel_tosa::parse(&first).unwrap();
945        model.validate_for(FLOAT_TARGET).unwrap();
946        let block = model.regions().next().unwrap().blocks().next().unwrap();
947        assert_eq!(block.operators().next().unwrap().op(), Op::IDENTITY);
948    }
949
950    #[test]
951    fn constants_and_typed_nan_attributes_round_trip() {
952        let one = 1.0_f32.to_le_bytes();
953        let tensors = [
954            Tensor::new("lhs", &[1], DType::FP32),
955            Tensor::constant("rhs", &[1], DType::FP32, &one),
956            Tensor::new("output", &[1], DType::FP32),
957        ];
958        let operators = [
959            Operator::new(OperatorKind::Const, &[], &["rhs"]),
960            Operator::new(
961                OperatorKind::Maximum {
962                    nan_mode: NanPropagationMode::PROPAGATE,
963                },
964                &["lhs", "rhs"],
965                &["output"],
966            ),
967        ];
968        let bytes = Graph::new("main", &tensors, &operators, &["lhs"], &["output"])
969            .build(FLOAT_TARGET)
970            .unwrap();
971
972        let model = virtio_accel_tosa::parse(&bytes).unwrap();
973        let block = model.regions().next().unwrap().blocks().next().unwrap();
974        assert_eq!(block.operators().len(), 2);
975    }
976
977    /// The slot offsets of every newly authorable attribute, checked by parsing the artifact
978    /// back. The opcode/union-tag test above cannot catch a wrong offset: a misplaced axis
979    /// serializes cleanly and reads back as a different number.
980    #[test]
981    fn newly_authorable_attributes_survive_a_parse() {
982        use virtio_accel_tosa::OpAttributes;
983
984        fn attributes_of(
985            tensors: &[Tensor<'_>],
986            operators: &[Operator<'_>],
987            inputs: &[&str],
988            outputs: &[&str],
989        ) -> Vec<u8> {
990            Graph::new("main", tensors, operators, inputs, outputs)
991                .build(FLOAT_TARGET)
992                .expect("graph builds")
993        }
994
995        let f32_shape = &[2, 3][..];
996        let bytes = attributes_of(
997            &[
998                Tensor::new("x", f32_shape, DType::FP32),
999                Tensor::new("y", &[2, 1], DType::FP32),
1000            ],
1001            &[Operator::new(
1002                OperatorKind::ReduceSum { axis: 1 },
1003                &["x"],
1004                &["y"],
1005            )],
1006            &["x"],
1007            &["y"],
1008        );
1009        let model = virtio_accel_tosa::parse(&bytes).unwrap();
1010        let analysis = model.analyze_for(FLOAT_TARGET).unwrap();
1011        let operator = analysis.operators().last().unwrap();
1012        assert!(matches!(
1013            operator.source().attributes(),
1014            OpAttributes::ReduceSum { axis: 1 }
1015        ));
1016
1017        let bytes = attributes_of(
1018            &[
1019                Tensor::new("x", f32_shape, DType::FP32),
1020                Tensor::new("y", &[3, 2], DType::FP32),
1021            ],
1022            &[Operator::new(
1023                OperatorKind::Transpose {
1024                    perms: [1, 0, 0, 0, 0, 0],
1025                    rank: 2,
1026                },
1027                &["x"],
1028                &["y"],
1029            )],
1030            &["x"],
1031            &["y"],
1032        );
1033        let model = virtio_accel_tosa::parse(&bytes).unwrap();
1034        let analysis = model.analyze_for(FLOAT_TARGET).unwrap();
1035        let OpAttributes::Transpose { perms } =
1036            analysis.operators().last().unwrap().source().attributes()
1037        else {
1038            panic!("TRANSPOSE attribute");
1039        };
1040        assert_eq!(perms.iter().collect::<Vec<_>>(), vec![1, 0]);
1041
1042        let lo = (-1.0_f32).to_le_bytes();
1043        let hi = 1.0_f32.to_le_bytes();
1044        let mut min_val = [0; MAX_CLAMP_BOUND_BYTES];
1045        let mut max_val = [0; MAX_CLAMP_BOUND_BYTES];
1046        min_val[..4].copy_from_slice(&lo);
1047        max_val[..4].copy_from_slice(&hi);
1048        let bytes = attributes_of(
1049            &[
1050                Tensor::new("x", f32_shape, DType::FP32),
1051                Tensor::new("y", f32_shape, DType::FP32),
1052            ],
1053            &[Operator::new(
1054                OperatorKind::Clamp {
1055                    min_val,
1056                    max_val,
1057                    bound_bytes: 4,
1058                    nan_mode: NanPropagationMode::PROPAGATE,
1059                },
1060                &["x"],
1061                &["y"],
1062            )],
1063            &["x"],
1064            &["y"],
1065        );
1066        let model = virtio_accel_tosa::parse(&bytes).unwrap();
1067        let analysis = model.analyze_for(FLOAT_TARGET).unwrap();
1068        let OpAttributes::Clamp {
1069            min_val, max_val, ..
1070        } = analysis.operators().last().unwrap().source().attributes()
1071        else {
1072            panic!("CLAMP attribute");
1073        };
1074        assert_eq!(min_val, &lo[..]);
1075        assert_eq!(max_val, &hi[..]);
1076
1077        let bytes = attributes_of(
1078            &[
1079                Tensor::new("x", f32_shape, DType::FP32),
1080                Tensor::new("y", &[2], DType::INT32),
1081            ],
1082            &[Operator::new(
1083                OperatorKind::ArgMax {
1084                    axis: 1,
1085                    nan_mode: NanPropagationMode::IGNORE,
1086                },
1087                &["x"],
1088                &["y"],
1089            )],
1090            &["x"],
1091            &["y"],
1092        );
1093        let model = virtio_accel_tosa::parse(&bytes).unwrap();
1094        let analysis = model.analyze_for(FLOAT_TARGET).unwrap();
1095        let OpAttributes::ArgMax { axis, nan_mode } =
1096            analysis.operators().last().unwrap().source().attributes()
1097        else {
1098            panic!("ARGMAX attribute");
1099        };
1100        assert_eq!(axis, 1);
1101        assert_eq!(nan_mode, NanPropagationMode::IGNORE);
1102    }
1103
1104    #[test]
1105    fn incrementally_owned_graph_matches_borrowed_artifact_bytes() {
1106        let one = 1.0_f32.to_le_bytes();
1107        let tensors = [
1108            Tensor::new("lhs", &[1], DType::FP32),
1109            Tensor::constant("rhs", &[1], DType::FP32, &one),
1110            Tensor::new("output", &[1], DType::FP32),
1111        ];
1112        let operators = [
1113            Operator::new(OperatorKind::Const, &[], &["rhs"]),
1114            Operator::new(OperatorKind::Add, &["lhs", "rhs"], &["output"]),
1115        ];
1116        let expected = Graph::new("main", &tensors, &operators, &["lhs"], &["output"])
1117            .build(FLOAT_TARGET)
1118            .unwrap();
1119
1120        let mut graph = OwnedGraph::new(String::from("main"));
1121        graph
1122            .push_tensor(OwnedTensor::new(String::from("lhs"), vec![1], DType::FP32))
1123            .push_tensor(OwnedTensor::constant(
1124                String::from("rhs"),
1125                vec![1],
1126                DType::FP32,
1127                one.to_vec(),
1128            ))
1129            .push_tensor(OwnedTensor::new("output", vec![1], DType::FP32))
1130            .push_operator(OwnedOperator::new(
1131                OperatorKind::Const,
1132                vec![],
1133                vec![String::from("rhs")],
1134            ))
1135            .push_operator(OwnedOperator::new(
1136                OperatorKind::Add,
1137                vec![String::from("lhs"), String::from("rhs")],
1138                vec![String::from("output")],
1139            ))
1140            .push_input(String::from("lhs"))
1141            .push_output(String::from("output"));
1142
1143        let first = graph.build(FLOAT_TARGET).unwrap();
1144        let second = graph.build(FLOAT_TARGET).unwrap();
1145        assert_eq!(first, expected);
1146        assert_eq!(second, expected);
1147    }
1148
1149    #[test]
1150    fn borrowed_graph_converts_to_an_incremental_graph() {
1151        let tensors = [
1152            Tensor::new("input", &[1, 4], DType::FP32),
1153            Tensor::new("output", &[2, 2], DType::FP32),
1154        ];
1155        let shapes = [Shape::new("target", &[2, 2])];
1156        let operators = [
1157            Operator::new(OperatorKind::ConstShape, &[], &["target"]),
1158            Operator::new(OperatorKind::Reshape, &["input", "target"], &["output"]),
1159        ];
1160        let borrowed =
1161            Graph::new("main", &tensors, &operators, &["input"], &["output"]).with_shapes(&shapes);
1162        let expected = borrowed.build(FLOAT_TARGET).unwrap();
1163
1164        let owned = OwnedGraph::from(borrowed);
1165        assert_eq!(owned.build(FLOAT_TARGET).unwrap(), expected);
1166        assert_eq!(owned.shapes, vec![OwnedShape::new("target", vec![2, 2])]);
1167    }
1168
1169    #[test]
1170    fn owned_tensor_can_reuse_existing_constant_storage() {
1171        let data = 1.0_f32.to_le_bytes();
1172        let tensor = OwnedTensor::constant("value", vec![1], DType::FP32, data.as_slice());
1173
1174        let Some(Cow::Borrowed(stored)) = tensor.data else {
1175            panic!("borrowed constant data was copied");
1176        };
1177        assert!(core::ptr::eq(stored, data.as_slice()));
1178    }
1179
1180    #[test]
1181    fn owned_graph_uses_the_borrowed_validation_contract() {
1182        let mut graph = OwnedGraph::new("main");
1183        graph
1184            .push_tensor(OwnedTensor::constant(
1185                "value",
1186                vec![1],
1187                DType::FP32,
1188                1.0_f32.to_le_bytes().to_vec(),
1189            ))
1190            .push_input("value")
1191            .push_output("value");
1192
1193        assert!(matches!(
1194            graph.build(FLOAT_TARGET),
1195            Err(BuildError::TensorDataWithoutConst)
1196        ));
1197    }
1198
1199    #[test]
1200    fn reshape_uses_a_typed_const_shape_operand() {
1201        let tensors = [
1202            Tensor::new("input", &[1, 4], DType::FP32),
1203            Tensor::new("output", &[2, 2], DType::FP32),
1204        ];
1205        let shapes = [Shape::new("target", &[2, 2])];
1206        let operators = [
1207            Operator::new(OperatorKind::ConstShape, &[], &["target"]),
1208            Operator::new(OperatorKind::Reshape, &["input", "target"], &["output"]),
1209        ];
1210        let bytes = Graph::new("main", &tensors, &operators, &["input"], &["output"])
1211            .with_shapes(&shapes)
1212            .build(FLOAT_TARGET)
1213            .unwrap();
1214
1215        let model = virtio_accel_tosa::parse(&bytes).unwrap();
1216        let block = model.regions().next().unwrap().blocks().next().unwrap();
1217        assert_eq!(block.shapes().len(), 1);
1218        assert_eq!(block.shapes().next().unwrap().values().unwrap().len(), 2);
1219        assert_eq!(block.operators().len(), 2);
1220    }
1221
1222    #[test]
1223    fn inline_tensor_data_requires_a_nonempty_const_producer() {
1224        let one = 1.0_f32.to_le_bytes();
1225        let data_without_const = [Tensor::constant("value", &[1], DType::FP32, &one)];
1226        assert!(matches!(
1227            Graph::new("main", &data_without_const, &[], &["value"], &["value"])
1228                .build(FLOAT_TARGET),
1229            Err(BuildError::TensorDataWithoutConst)
1230        ));
1231
1232        let missing_data = [Tensor::new("value", &[1], DType::FP32)];
1233        let constant = [Operator::new(OperatorKind::Const, &[], &["value"])];
1234        assert!(matches!(
1235            Graph::new("main", &missing_data, &constant, &[], &["value"]).build(FLOAT_TARGET),
1236            Err(BuildError::ConstWithoutTensorData)
1237        ));
1238
1239        let empty_data = [Tensor::constant("value", &[0], DType::FP32, &[])];
1240        assert!(matches!(
1241            Graph::new("main", &empty_data, &constant, &[], &["value"]).build(FLOAT_TARGET),
1242            Err(BuildError::EmptyConstantData)
1243        ));
1244    }
1245
1246    #[test]
1247    fn every_operator_kind_serializes_its_pinned_opcode_and_union_tag() {
1248        let cases = [
1249            (OperatorKind::MatMul, Op::MATMUL, 7),
1250            (
1251                OperatorKind::ArgMax {
1252                    axis: 1,
1253                    nan_mode: NanPropagationMode::PROPAGATE,
1254                },
1255                Op::ARGMAX,
1256                1,
1257            ),
1258            (OperatorKind::Ceil, Op::CEIL, 34),
1259            (OperatorKind::Floor, Op::FLOOR, 38),
1260            (
1261                OperatorKind::Clamp {
1262                    min_val: [0; MAX_CLAMP_BOUND_BYTES],
1263                    max_val: [0, 0, 0x80, 0x3f, 0, 0, 0, 0],
1264                    bound_bytes: 4,
1265                    nan_mode: NanPropagationMode::PROPAGATE,
1266                },
1267                Op::CLAMP,
1268                11,
1269            ),
1270            (OperatorKind::Concat { axis: 0 }, Op::CONCAT, 55),
1271            (OperatorKind::Reverse { axis: 1 }, Op::REVERSE, 58),
1272            (
1273                OperatorKind::Transpose {
1274                    perms: [1, 0, 2, 0, 0, 0],
1275                    rank: 3,
1276                },
1277                Op::TRANSPOSE,
1278                61,
1279            ),
1280            (
1281                OperatorKind::ReduceMax {
1282                    axis: 1,
1283                    nan_mode: NanPropagationMode::PROPAGATE,
1284                },
1285                Op::REDUCE_MAX,
1286                51,
1287            ),
1288            (
1289                OperatorKind::ReduceMin {
1290                    axis: 1,
1291                    nan_mode: NanPropagationMode::IGNORE,
1292                },
1293                Op::REDUCE_MIN,
1294                52,
1295            ),
1296            (
1297                OperatorKind::ReduceProduct { axis: 0 },
1298                Op::REDUCE_PRODUCT,
1299                53,
1300            ),
1301            (OperatorKind::ReduceSum { axis: 0 }, Op::REDUCE_SUM, 54),
1302            (OperatorKind::Erf, Op::ERF, 12),
1303            (
1304                OperatorKind::MaxPool2d {
1305                    kernel: [2, 2],
1306                    stride: [2, 2],
1307                    pad: [0; 4],
1308                    nan_mode: NanPropagationMode::PROPAGATE,
1309                },
1310                Op::MAX_POOL2D,
1311                8,
1312            ),
1313            (OperatorKind::Sigmoid, Op::SIGMOID, 13),
1314            (OperatorKind::Tanh, Op::TANH, 14),
1315            (OperatorKind::Add, Op::ADD, 15),
1316            (OperatorKind::LogicalAnd, Op::LOGICAL_AND, 21),
1317            (OperatorKind::LogicalOr, Op::LOGICAL_OR, 24),
1318            (OperatorKind::LogicalXor, Op::LOGICAL_XOR, 25),
1319            (
1320                OperatorKind::Maximum {
1321                    nan_mode: NanPropagationMode::PROPAGATE,
1322                },
1323                Op::MAXIMUM,
1324                26,
1325            ),
1326            (
1327                OperatorKind::Minimum {
1328                    nan_mode: NanPropagationMode::PROPAGATE,
1329                },
1330                Op::MINIMUM,
1331                27,
1332            ),
1333            (OperatorKind::Mul, Op::MUL, 28),
1334            (OperatorKind::Pow, Op::POW, 29),
1335            (OperatorKind::Sub, Op::SUB, 30),
1336            (OperatorKind::Abs, Op::ABS, 32),
1337            (OperatorKind::Cos, Op::COS, 36),
1338            (OperatorKind::Exp, Op::EXP, 37),
1339            (OperatorKind::Log, Op::LOG, 39),
1340            (OperatorKind::LogicalNot, Op::LOGICAL_NOT, 40),
1341            (OperatorKind::Negate, Op::NEGATE, 41),
1342            (OperatorKind::Reciprocal, Op::RECIPROCAL, 42),
1343            (OperatorKind::Rsqrt, Op::RSQRT, 43),
1344            (OperatorKind::Sin, Op::SIN, 44),
1345            (OperatorKind::Select, Op::SELECT, 45),
1346            (OperatorKind::Equal, Op::EQUAL, 46),
1347            (OperatorKind::Greater, Op::GREATER, 47),
1348            (OperatorKind::GreaterEqual, Op::GREATER_EQUAL, 48),
1349            (OperatorKind::Reshape, Op::RESHAPE, 57),
1350            (OperatorKind::Cast, Op::CAST, 65),
1351            (
1352                OperatorKind::Rescale {
1353                    scale32: true,
1354                    rounding_mode: RoundingMode::SINGLE_ROUND,
1355                    per_channel: false,
1356                    input_unsigned: false,
1357                    output_unsigned: false,
1358                },
1359                Op::RESCALE,
1360                66,
1361            ),
1362            (OperatorKind::Const, Op::CONST, 67),
1363            (OperatorKind::Identity, Op::IDENTITY, 68),
1364            (OperatorKind::ConstShape, Op::CONST_SHAPE, 75),
1365        ];
1366
1367        for (kind, expected_op, expected_attribute) in cases {
1368            let tensors = [
1369                Tensor::new("input", &[1], DType::FP32),
1370                Tensor::new("output", &[1], DType::FP32),
1371            ];
1372            let shapes = [Shape::new("target", &[1])];
1373            let regular_inputs = ["input"];
1374            let reshape_inputs = ["input", "target"];
1375            let tensor_output = ["output"];
1376            let shape_output = ["target"];
1377            let (inputs, outputs): (&[&str], &[&str]) = match kind {
1378                OperatorKind::Const => (&[], &tensor_output),
1379                OperatorKind::ConstShape => (&[], &shape_output),
1380                OperatorKind::Reshape => (&reshape_inputs, &tensor_output),
1381                _ => (&regular_inputs, &tensor_output),
1382            };
1383            let operator = [Operator::new(kind, inputs, outputs)];
1384            let graph = Graph::new("main", &tensors, &operator, &[], &[]).with_shapes(&shapes);
1385            let bytes = graph.serialize();
1386            let model = virtio_accel_tosa::parse(&bytes).unwrap();
1387            let parsed = model
1388                .regions()
1389                .next()
1390                .unwrap()
1391                .blocks()
1392                .next()
1393                .unwrap()
1394                .operators()
1395                .next()
1396                .unwrap();
1397            assert_eq!(parsed.op(), expected_op);
1398            assert_eq!(parsed.attribute_kind().get(), expected_attribute);
1399        }
1400    }
1401
1402    #[test]
1403    fn rejects_dynamic_shapes_before_serialization() {
1404        let tensors = [Tensor::new("value", &[-1], DType::FP32)];
1405        let graph = Graph::new("main", &tensors, &[], &["value"], &["value"]);
1406        assert!(matches!(
1407            graph.build(FLOAT_TARGET),
1408            Err(BuildError::DynamicShape)
1409        ));
1410    }
1411}