1#![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
24mod 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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
89pub struct Shape<'a> {
90 pub name: &'a str,
92 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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
104pub struct Tensor<'a> {
105 pub name: &'a str,
107 pub shape: &'a [i32],
109 pub dtype: DType,
111 pub data: Option<&'a [u8]>,
113}
114
115impl<'a> Tensor<'a> {
116 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 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
137pub const MAX_CLAMP_BOUND_BYTES: usize = 8;
145
146pub 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 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 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 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#[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#[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#[derive(Clone, Debug, PartialEq, Eq)]
336pub struct OwnedTensor<'a> {
337 pub name: String,
338 pub shape: Vec<i32>,
339 pub dtype: DType,
340 pub data: Option<Cow<'a, [u8]>>,
343}
344
345impl<'a> OwnedTensor<'a> {
346 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 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#[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 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#[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 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#[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 pub const fn with_shapes(mut self, shapes: &'a [Shape<'a>]) -> Self {
575 self.shapes = shapes;
576 self
577 }
578
579 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 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#[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 #[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 _ => (®ular_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}