1use alloc::vec::Vec;
2use core::fmt;
3
4use crate::{
5 BasicBlock, DType, ExtensionSet, I32List, LevelLimits, Model, ModelValidator,
6 NanPropagationMode, Op, OpAttributes, Operator, ProfileSet, ResizeMode, RoundingMode, Shape,
7 Target, TargetError, Tensor,
8};
9
10#[derive(Clone, Copy, Debug, PartialEq, Eq)]
12pub enum OperandRole {
13 Input,
14 Output,
15}
16
17#[derive(Clone, Copy, Debug, PartialEq, Eq)]
19pub enum SemanticErrorKind {
20 GraphIoMustBeTensor,
21 InvalidArity {
22 op: Op,
23 inputs: usize,
24 outputs: usize,
25 },
26 TensorListLimit {
27 op: Op,
28 actual: usize,
29 limit: usize,
30 },
31 ExpectedTensor {
32 op: Op,
33 role: OperandRole,
34 index: usize,
35 },
36 ExpectedShape {
37 op: Op,
38 role: OperandRole,
39 index: usize,
40 },
41 InvalidRank {
42 op: Op,
43 role: OperandRole,
44 index: usize,
45 rank: Option<usize>,
46 minimum: usize,
47 maximum: usize,
48 },
49 UnsupportedTypeProfile(Op),
50 InvalidAttribute(Op),
51 InvalidShape(Op),
52 ConstantRequired {
53 op: Op,
54 input: usize,
55 },
56 InvalidConstantData {
57 op: Op,
58 operand: usize,
59 },
60 InvalidTensorData,
61 TensorSizeLimit,
62 ShapeValueLimit,
63 InvalidVariable,
64 DisconnectedSymbol,
65 GraphInputProduced,
66 DataflowCycle,
67 UnknownControlFlowRegion(Op),
68 ControlFlowSignature(Op),
69 ControlFlowCycle,
70 ControlFlowNestingLimit {
71 actual: usize,
72 limit: usize,
73 },
74}
75
76#[derive(Clone, Copy, Debug, PartialEq, Eq)]
78pub enum SemanticError {
79 InvalidTarget(TargetError),
80 AllocationFailed,
81 Graph {
82 region: usize,
83 block: usize,
84 operator: Option<usize>,
85 kind: SemanticErrorKind,
86 },
87}
88
89impl fmt::Display for SemanticError {
90 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
91 write!(formatter, "{self:?}")
92 }
93}
94
95#[derive(Clone, Copy, Debug)]
97pub struct SemanticValidator {
98 target: Target,
99}
100
101impl SemanticValidator {
102 pub const fn new(target: Target) -> Self {
103 Self { target }
104 }
105
106 pub const fn target(&self) -> Target {
107 self.target
108 }
109}
110
111impl ModelValidator for SemanticValidator {
112 type Error = SemanticError;
113
114 fn validate(&mut self, model: &Model<'_>) -> Result<(), Self::Error> {
115 validate_semantics(model, self.target)
116 }
117}
118
119pub fn validate_semantics(model: &Model<'_>, target: Target) -> Result<(), SemanticError> {
121 let target = target.validate().map_err(SemanticError::InvalidTarget)?;
122 if target.version != model.version() {
123 return Err(SemanticError::InvalidTarget(TargetError::VersionMismatch {
124 target: target.version,
125 model: model.version(),
126 }));
127 }
128
129 let mut region_names = Vec::new();
130 region_names
131 .try_reserve_exact(model.regions().len())
132 .map_err(|_| SemanticError::AllocationFailed)?;
133 for region in model.regions() {
134 region_names.push(region.name());
135 }
136 region_names.sort_unstable();
137 validate_control_flow_nesting(model, ®ion_names, target.level.limits().max_nesting)?;
138
139 for (region_index, region) in model.regions().enumerate() {
140 for (block_index, block) in region.blocks().enumerate() {
141 validate_block(
142 model,
143 block,
144 target,
145 ®ion_names,
146 region_index,
147 block_index,
148 )?;
149 }
150 }
151 Ok(())
152}
153
154fn validate_control_flow_nesting(
155 model: &Model<'_>,
156 region_names: &[&str],
157 limit: usize,
158) -> Result<(), SemanticError> {
159 let mut states = Vec::new();
160 let mut depths = Vec::new();
161 states
162 .try_reserve_exact(region_names.len())
163 .map_err(|_| SemanticError::AllocationFailed)?;
164 depths
165 .try_reserve_exact(region_names.len())
166 .map_err(|_| SemanticError::AllocationFailed)?;
167 for _ in region_names {
168 states.push(0_u8);
169 depths.push(0_usize);
170 }
171 for index in 0..region_names.len() {
172 let depth = control_flow_depth(
173 model,
174 region_names,
175 &mut states,
176 &mut depths,
177 index,
178 0,
179 limit,
180 )
181 .map_err(|kind| SemanticError::Graph {
182 region: index,
183 block: 0,
184 operator: None,
185 kind,
186 })?;
187 if depth > limit {
188 return Err(SemanticError::Graph {
189 region: index,
190 block: 0,
191 operator: None,
192 kind: SemanticErrorKind::ControlFlowNestingLimit {
193 actual: depth,
194 limit,
195 },
196 });
197 }
198 }
199 Ok(())
200}
201
202fn control_flow_depth(
203 model: &Model<'_>,
204 region_names: &[&str],
205 states: &mut [u8],
206 depths: &mut [usize],
207 index: usize,
208 path_depth: usize,
209 limit: usize,
210) -> Result<usize, SemanticErrorKind> {
211 if path_depth > limit {
212 return Err(SemanticErrorKind::ControlFlowNestingLimit {
213 actual: path_depth,
214 limit,
215 });
216 }
217 match states[index] {
218 1 => return Err(SemanticErrorKind::ControlFlowCycle),
219 2 => return Ok(depths[index]),
220 _ => {}
221 }
222 states[index] = 1;
223 let mut maximum = 0_usize;
224 if let Some(region) = model
225 .regions()
226 .find(|region| region.name() == region_names[index])
227 {
228 for block in region.blocks() {
229 for operator in block.operators() {
230 let attributes = operator.attributes();
231 let children: &[Option<&str>] = match attributes {
232 OpAttributes::CondIf {
233 then_graph,
234 else_graph,
235 } => &[then_graph, else_graph],
236 OpAttributes::WhileLoop {
237 cond_graph,
238 body_graph,
239 } => &[cond_graph, body_graph],
240 _ => &[],
241 };
242 for child in children.iter().flatten() {
243 let Ok(child_index) = region_names.binary_search(child) else {
244 continue;
245 };
246 let child_depth = control_flow_depth(
247 model,
248 region_names,
249 states,
250 depths,
251 child_index,
252 path_depth + 1,
253 limit,
254 )?;
255 maximum = maximum.max(child_depth.saturating_add(1));
256 }
257 }
258 }
259 }
260 states[index] = 2;
261 depths[index] = maximum;
262 Ok(maximum)
263}
264
265#[derive(Clone, Copy)]
266enum Symbol<'a> {
267 Tensor(Tensor<'a>),
268 Shape(Shape<'a>),
269}
270
271fn validate_block<'a>(
272 model: &Model<'a>,
273 block: BasicBlock<'a>,
274 target: Target,
275 region_names: &[&str],
276 region_index: usize,
277 block_index: usize,
278) -> Result<(), SemanticError> {
279 let symbol_count = block.tensors().len() + block.shapes().len();
280 let mut symbols = Vec::new();
281 symbols
282 .try_reserve_exact(symbol_count)
283 .map_err(|_| SemanticError::AllocationFailed)?;
284 for tensor in block.tensors() {
285 validate_tensor_encoding(model, tensor).map_err(|kind| SemanticError::Graph {
286 region: region_index,
287 block: block_index,
288 operator: None,
289 kind,
290 })?;
291 if !tensor_within_level(tensor, target.level.limits()) {
292 return Err(SemanticError::Graph {
293 region: region_index,
294 block: block_index,
295 operator: None,
296 kind: SemanticErrorKind::TensorSizeLimit,
297 });
298 }
299 if (tensor.is_variable()
300 && (!target.extensions.contains(ExtensionSet::VARIABLE)
301 || tensor.variable_name().is_none_or(str::is_empty)))
302 || (!tensor.is_variable()
303 && tensor.variable_name().is_some_and(|name| !name.is_empty()))
304 {
305 return Err(SemanticError::Graph {
306 region: region_index,
307 block: block_index,
308 operator: None,
309 kind: SemanticErrorKind::InvalidVariable,
310 });
311 }
312 symbols.push((tensor.name(), Symbol::Tensor(tensor)));
313 }
314 for shape in block.shapes() {
315 if !shape_within_level(shape, target.level.limits()) {
316 return Err(SemanticError::Graph {
317 region: region_index,
318 block: block_index,
319 operator: None,
320 kind: SemanticErrorKind::ShapeValueLimit,
321 });
322 }
323 symbols.push((shape.name(), Symbol::Shape(shape)));
324 }
325 symbols.sort_unstable_by_key(|(name, _)| *name);
326
327 let constant_count = block
328 .operators()
329 .filter(|operator| matches!(operator.op(), Op::CONST | Op::CONST_SHAPE))
330 .try_fold(0_usize, |count, operator| {
331 count.checked_add(operator.outputs().len())
332 })
333 .ok_or(SemanticError::AllocationFailed)?;
334 let mut constants = Vec::new();
335 constants
336 .try_reserve_exact(constant_count)
337 .map_err(|_| SemanticError::AllocationFailed)?;
338 for operator in block.operators() {
339 if matches!(operator.op(), Op::CONST | Op::CONST_SHAPE) {
340 constants.extend(operator.outputs());
341 }
342 }
343 constants.sort_unstable();
344
345 for name in block.inputs().chain(block.outputs()) {
346 if !matches!(resolve(&symbols, name), Symbol::Tensor(_)) {
347 return Err(SemanticError::Graph {
348 region: region_index,
349 block: block_index,
350 operator: None,
351 kind: SemanticErrorKind::GraphIoMustBeTensor,
352 });
353 }
354 }
355 validate_dataflow(block, &symbols, region_index, block_index)?;
356
357 let mut inputs = Vec::new();
358 let mut outputs = Vec::new();
359 for (operator_index, operator) in block.operators().enumerate() {
360 inputs.clear();
361 outputs.clear();
362 inputs
363 .try_reserve(operator.inputs().len())
364 .map_err(|_| SemanticError::AllocationFailed)?;
365 outputs
366 .try_reserve(operator.outputs().len())
367 .map_err(|_| SemanticError::AllocationFailed)?;
368 for name in operator.inputs() {
369 inputs.push(resolve(&symbols, name));
370 }
371 for name in operator.outputs() {
372 outputs.push(resolve(&symbols, name));
373 }
374
375 validate_operator(
376 model,
377 operator,
378 &inputs,
379 &outputs,
380 target,
381 region_names,
382 &constants,
383 )
384 .map_err(|kind| SemanticError::Graph {
385 region: region_index,
386 block: block_index,
387 operator: Some(operator_index),
388 kind,
389 })?;
390 }
391 Ok(())
392}
393
394fn resolve<'a>(symbols: &[(&'a str, Symbol<'a>)], name: &str) -> Symbol<'a> {
395 let index = symbols
396 .binary_search_by_key(&name, |(candidate, _)| *candidate)
397 .expect("structurally validated symbol reference");
398 symbols[index].1
399}
400
401fn validate_dataflow<'a>(
402 block: BasicBlock<'a>,
403 symbols: &[(&'a str, Symbol<'a>)],
404 region_index: usize,
405 block_index: usize,
406) -> Result<(), SemanticError> {
407 let allocation = |_| SemanticError::AllocationFailed;
408 let located = |kind| SemanticError::Graph {
409 region: region_index,
410 block: block_index,
411 operator: None,
412 kind,
413 };
414 let operator_count = block.operators().len();
415
416 let mut sources = Vec::new();
417 let source_capacity = block
418 .inputs()
419 .len()
420 .checked_add(block.tensors().len())
421 .ok_or(SemanticError::AllocationFailed)?;
422 sources
423 .try_reserve_exact(source_capacity)
424 .map_err(allocation)?;
425 sources.extend(block.inputs());
426 sources.extend(
427 block
428 .tensors()
429 .filter(Tensor::is_variable)
430 .map(|tensor| tensor.name()),
431 );
432 sources.sort_unstable();
433
434 let mut producers = Vec::new();
435 producers
436 .try_reserve_exact(symbols.len())
437 .map_err(allocation)?;
438 let mut edge_capacity = 0_usize;
439 for (operator_index, operator) in block.operators().enumerate() {
440 edge_capacity = edge_capacity
441 .checked_add(operator.inputs().len())
442 .ok_or(SemanticError::AllocationFailed)?;
443 for name in operator.outputs() {
444 if !matches!(resolve(symbols, name), Symbol::Tensor(value) if value.is_variable()) {
445 if sources.binary_search(&name).is_ok() {
446 return Err(located(SemanticErrorKind::GraphInputProduced));
447 }
448 producers.push((name, operator_index));
449 }
450 }
451 }
452 producers.sort_unstable_by_key(|(name, _)| *name);
453
454 let mut indegrees = Vec::new();
455 indegrees
456 .try_reserve_exact(operator_count)
457 .map_err(allocation)?;
458 indegrees.resize(operator_count, 0_usize);
459 let mut edges = Vec::new();
460 edges.try_reserve_exact(edge_capacity).map_err(allocation)?;
461 for (consumer, operator) in block.operators().enumerate() {
462 for name in operator.inputs() {
463 if sources.binary_search(&name).is_ok() {
464 continue;
465 }
466 let producer = producers
467 .binary_search_by_key(&name, |(candidate, _)| *candidate)
468 .ok()
469 .map(|index| producers[index].1)
470 .ok_or_else(|| located(SemanticErrorKind::DisconnectedSymbol))?;
471 edges.push((producer, consumer));
472 indegrees[consumer] = indegrees[consumer]
473 .checked_add(1)
474 .ok_or(SemanticError::AllocationFailed)?;
475 }
476 }
477 for name in block.outputs() {
478 if sources.binary_search(&name).is_err()
479 && producers
480 .binary_search_by_key(&name, |(candidate, _)| *candidate)
481 .is_err()
482 {
483 return Err(located(SemanticErrorKind::DisconnectedSymbol));
484 }
485 }
486
487 edges.sort_unstable();
488 let mut ready = Vec::new();
489 ready
490 .try_reserve_exact(operator_count)
491 .map_err(allocation)?;
492 ready.extend(
493 indegrees
494 .iter()
495 .enumerate()
496 .filter_map(|(index, indegree)| (*indegree == 0).then_some(index)),
497 );
498 let mut cursor = 0_usize;
499 while cursor < ready.len() {
500 let producer = ready[cursor];
501 cursor += 1;
502 let start = edges.partition_point(|(candidate, _)| *candidate < producer);
503 let end = edges.partition_point(|(candidate, _)| *candidate <= producer);
504 for &(_, consumer) in &edges[start..end] {
505 indegrees[consumer] -= 1;
506 if indegrees[consumer] == 0 {
507 ready.push(consumer);
508 }
509 }
510 }
511 if cursor == operator_count {
512 Ok(())
513 } else {
514 Err(located(SemanticErrorKind::DataflowCycle))
515 }
516}
517
518fn validate_operator(
519 model: &Model<'_>,
520 operator: Operator<'_>,
521 inputs: &[Symbol<'_>],
522 outputs: &[Symbol<'_>],
523 target: Target,
524 region_names: &[&str],
525 constants: &[&str],
526) -> Result<(), SemanticErrorKind> {
527 let op = operator.op();
528 let arity = op.arity().expect("stable operator has an arity");
529 if !arity.accepts(inputs.len(), outputs.len()) {
530 return Err(SemanticErrorKind::InvalidArity {
531 op,
532 inputs: inputs.len(),
533 outputs: outputs.len(),
534 });
535 }
536 validate_tensor_list_limit(op, inputs.len(), outputs.len(), target.level.limits())?;
537 validate_operand_kinds(op, inputs, outputs)?;
538 validate_ranks(op, inputs, outputs, target.level.limits())?;
539 validate_type_profile(op, operator.attributes(), inputs, outputs, target)?;
540 validate_attributes(op, operator.attributes(), inputs, outputs, target)?;
541 validate_ctc_inputs(
542 model,
543 op,
544 operator,
545 operator.attributes(),
546 inputs,
547 constants,
548 target,
549 )?;
550 validate_shapes(
551 model,
552 op,
553 operator.attributes(),
554 inputs,
555 outputs,
556 target,
557 region_names,
558 )
559}
560
561fn validate_tensor_list_limit(
562 op: Op,
563 inputs: usize,
564 outputs: usize,
565 limits: LevelLimits,
566) -> Result<(), SemanticErrorKind> {
567 let actual = match op.get() {
568 55 => inputs,
569 69 => inputs.max(outputs),
570 70 => inputs.saturating_sub(1).max(outputs),
571 71 => inputs.max(outputs),
572 _ => return Ok(()),
573 };
574 if actual > limits.max_tensor_list_size {
575 Err(SemanticErrorKind::TensorListLimit {
576 op,
577 actual,
578 limit: limits.max_tensor_list_size,
579 })
580 } else {
581 Ok(())
582 }
583}
584
585fn validate_operand_kinds(
586 op: Op,
587 inputs: &[Symbol<'_>],
588 outputs: &[Symbol<'_>],
589) -> Result<(), SemanticErrorKind> {
590 let shape_inputs: &[usize] = match op.get() {
591 56 | 57 | 60 => &[1],
592 59 => &[1, 2],
593 64 => &[1, 2, 3],
594 _ => &[],
595 };
596 for (index, input) in inputs.iter().enumerate() {
597 let wants_shape = shape_inputs.contains(&index);
598 match (wants_shape, input) {
599 (true, Symbol::Shape(_)) | (false, Symbol::Tensor(_)) => {}
600 (true, _) => {
601 return Err(SemanticErrorKind::ExpectedShape {
602 op,
603 role: OperandRole::Input,
604 index,
605 });
606 }
607 (false, _) => {
608 return Err(SemanticErrorKind::ExpectedTensor {
609 op,
610 role: OperandRole::Input,
611 index,
612 });
613 }
614 }
615 }
616 for (index, output) in outputs.iter().enumerate() {
617 let wants_shape = op == Op::CONST_SHAPE;
618 match (wants_shape, output) {
619 (true, Symbol::Shape(_)) | (false, Symbol::Tensor(_)) => {}
620 (true, _) => {
621 return Err(SemanticErrorKind::ExpectedShape {
622 op,
623 role: OperandRole::Output,
624 index,
625 });
626 }
627 (false, _) => {
628 return Err(SemanticErrorKind::ExpectedTensor {
629 op,
630 role: OperandRole::Output,
631 index,
632 });
633 }
634 }
635 }
636 Ok(())
637}
638
639fn tensor(symbol: Symbol<'_>) -> Tensor<'_> {
640 match symbol {
641 Symbol::Tensor(tensor) => tensor,
642 Symbol::Shape(_) => unreachable!("operand kind validated"),
643 }
644}
645
646fn shape(symbol: Symbol<'_>) -> Shape<'_> {
647 match symbol {
648 Symbol::Shape(shape) => shape,
649 Symbol::Tensor(_) => unreachable!("operand kind validated"),
650 }
651}
652
653fn validate_tensor_encoding(
654 model: &Model<'_>,
655 tensor: Tensor<'_>,
656) -> Result<(), SemanticErrorKind> {
657 if tensor.rank().is_none() {
658 return Ok(());
659 }
660 let mut elements = 1_usize;
661 for dimension in tensor.dimensions() {
662 elements = elements
663 .checked_mul(dimension as usize)
664 .ok_or(SemanticErrorKind::InvalidTensorData)?;
665 }
666 let data = tensor_data(model, tensor);
667 if !data.is_empty() {
668 let required = if tensor.dtype() == DType::INT4 {
669 elements
670 .checked_add(1)
671 .ok_or(SemanticErrorKind::InvalidTensorData)?
672 / 2
673 } else {
674 elements
675 .checked_mul(
676 dtype_width(tensor.dtype()).ok_or(SemanticErrorKind::InvalidTensorData)?,
677 )
678 .ok_or(SemanticErrorKind::InvalidTensorData)?
679 };
680 if data.len() != required {
681 return Err(SemanticErrorKind::InvalidTensorData);
682 }
683 let valid_values = match tensor.dtype() {
684 DType::BOOL => data.iter().all(|value| *value <= 1),
685 DType::INT4 => (0..elements).all(|index| crate::unpack_int4(data, index) != Some(-8)),
686 DType::INT48 => data.chunks_exact(8).all(|bytes| {
687 let value = i64::from_le_bytes(bytes.try_into().expect("chunk size is exact"));
688 (-(1_i64 << 47)..(1_i64 << 47)).contains(&value)
689 }),
690 _ => true,
691 };
692 if !valid_values {
693 return Err(SemanticErrorKind::InvalidTensorData);
694 }
695 }
696 Ok(())
697}
698
699fn tensor_data<'a>(model: &Model<'a>, tensor: Tensor<'a>) -> &'a [u8] {
700 if let Some((offset, size)) = tensor.external_data_range() {
701 let start = usize::try_from(offset).expect("validated external offset");
702 let len = usize::try_from(size).expect("validated external size");
703 &model.as_bytes()[start..start + len]
704 } else {
705 tensor.data()
706 }
707}
708
709fn dtype_width(dtype: DType) -> Option<usize> {
710 match dtype.get() {
711 1..=3 | 11..=12 => Some(1),
712 4 | 8 | 9 => Some(2),
713 5 | 7 => Some(4),
714 6 | 10 => Some(8),
715 _ => None,
716 }
717}
718
719fn dtype_element_bytes(dtype: DType) -> Option<usize> {
720 match dtype {
721 DType::INT48 => Some(6),
722 _ => dtype_width(dtype),
723 }
724}
725
726fn tensor_within_level(tensor: Tensor<'_>, limits: LevelLimits) -> bool {
727 let Some(_) = tensor.rank() else {
728 return true;
729 };
730 let Some(elements) = tensor.dimensions().try_fold(1_u128, |count, dimension| {
731 count.checked_mul(dimension as u128)
732 }) else {
733 return false;
734 };
735 let Some(bytes) =
736 dtype_element_bytes(tensor.dtype()).and_then(|width| elements.checked_mul(width as u128))
737 else {
738 return false;
739 };
740 let maximum = (1_u128 << limits.max_log2_size) - 1;
741 bytes <= maximum
742}
743
744fn shape_within_level(shape: Shape<'_>, limits: LevelLimits) -> bool {
745 let Some(mut values) = shape.values() else {
746 return true;
747 };
748 let magnitude = 1_i128 << limits.max_log2_size;
749 values.all(|value| i128::from(value) >= -magnitude && i128::from(value) < magnitude)
750}
751
752fn validate_ranks(
754 op: Op,
755 inputs: &[Symbol<'_>],
756 outputs: &[Symbol<'_>],
757 limits: LevelLimits,
758) -> Result<(), SemanticErrorKind> {
759 for (role, operands) in [(OperandRole::Input, inputs), (OperandRole::Output, outputs)] {
760 for (index, operand) in operands.iter().copied().enumerate() {
761 let Symbol::Tensor(value) = operand else {
762 continue;
763 };
764 let rank = value.rank();
765 if rank.is_none_or(|rank| rank > limits.max_rank) {
766 return Err(SemanticErrorKind::InvalidRank {
767 op,
768 role,
769 index,
770 rank,
771 minimum: 0,
772 maximum: limits.max_rank,
773 });
774 }
775 }
776 }
777
778 let (input_ranks, output_ranks): (&[Option<usize>], &[Option<usize>]) = match op.get() {
779 1 => (&[None], &[None]),
780 2 => (&[Some(4), Some(1), Some(1)], &[Some(4)]),
781 3 | 5 | 10 => (&[Some(4), Some(4), Some(1), Some(1), Some(1)], &[Some(4)]),
782 4 => (&[Some(5), Some(5), Some(1), Some(1), Some(1)], &[Some(5)]),
783 6 => (&[Some(3), Some(3)], &[Some(3), Some(3)]),
784 7 => (&[Some(3), Some(3), Some(1), Some(1)], &[Some(3)]),
785 8 => (&[Some(4)], &[Some(4)]),
786 9 => (&[Some(3)], &[Some(3), Some(3)]),
787 28 => (&[None, None, Some(1)], &[None]),
788 31 => (&[None, Some(1)], &[None]),
789 41 => (&[None, Some(1), Some(1)], &[None]),
790 49..=54 => (&[None], &[None]),
791 55 => (&[], &[None]),
792 56 => (&[None, None, Some(1)], &[None]),
793 57 => (&[None, None], &[None]),
794 58 => (&[None], &[None]),
795 59 => (&[None, None, None], &[None]),
796 60 => (&[None, None], &[None]),
797 61 => (&[None], &[None]),
798 62 => (&[Some(3), Some(2)], &[Some(3)]),
799 63 => (&[Some(3), Some(2), Some(3)], &[Some(3)]),
800 64 => (&[Some(4), None, None, None], &[Some(4)]),
801 66 => (&[None, Some(1), Some(1), Some(1), Some(1)], &[None]),
802 _ => (&[], &[]),
803 };
804 for (index, expected) in input_ranks.iter().copied().enumerate() {
805 if let Some(expected) = expected {
806 require_rank(
807 op,
808 OperandRole::Input,
809 index,
810 inputs[index],
811 expected,
812 expected,
813 )?;
814 }
815 }
816 for (index, expected) in output_ranks.iter().copied().enumerate() {
817 if let Some(expected) = expected {
818 require_rank(
819 op,
820 OperandRole::Output,
821 index,
822 outputs[index],
823 expected,
824 expected,
825 )?;
826 }
827 }
828
829 let minimum_rank = match op.get() {
830 1 | 49..=56 | 58..=61 => 1,
831 _ => 0,
832 };
833 if minimum_rank != 0 {
834 for (role, operands) in [(OperandRole::Input, inputs), (OperandRole::Output, outputs)] {
835 for (index, operand) in operands.iter().copied().enumerate() {
836 if matches!(operand, Symbol::Tensor(_)) {
837 require_rank(op, role, index, operand, minimum_rank, limits.max_rank)?;
838 }
839 }
840 }
841 }
842 Ok(())
843}
844
845fn validate_type_profile(
846 op: Op,
847 attributes: OpAttributes<'_>,
848 inputs: &[Symbol<'_>],
849 outputs: &[Symbol<'_>],
850 target: Target,
851) -> Result<(), SemanticErrorKind> {
852 let i = |index| tensor(inputs[index]).dtype();
853 let o = |index| tensor(outputs[index]).dtype();
854 let same = || tensor_types_equal(inputs, outputs);
855 let ok = match op.get() {
856 1 => o(0) == DType::INT32 && supports_argmax(i(0), target),
857 2 => {
858 let acc = match attributes {
859 OpAttributes::AvgPool2d { acc_type, .. } => acc_type,
860 _ => unreachable!(),
861 };
862 i(0) == i(1) && i(0) == i(2) && i(0) == o(0) && supports_pool(i(0), acc, target)
863 }
864 3..=5 | 10 => {
865 let acc = conv_acc_type(attributes);
866 i(3) == i(0)
867 && i(4) == i(1)
868 && i(2) == o(0)
869 && supports_conv(i(0), i(1), o(0), acc, target)
870 }
871 6 => same() && i(0) == DType::FP32 && has_ext(target, ExtensionSet::FFT),
872 7 => i(0) == i(1) && i(2) == i(0) && i(3) == i(1) && supports_matmul(i(0), o(0), target),
873 8 => same() && supports_pool_value(i(0), target),
874 9 => same() && i(0) == DType::FP32 && has_ext(target, ExtensionSet::FFT),
875 11 => same() && supports_clamp(i(0), target),
876 12..=14 | 34 | 36..=39 | 42..=44 => same() && supports_float(i(0), target),
877 15 | 30 => same() && supports_add_sub(i(0), target),
878 16..=19 | 33 => same() && supports_integer_bits(i(0), target),
879 20 => same() && i(0) == DType::INT32 && has_any_profile(target),
880 21 | 24 | 25 | 40 => same() && i(0) == DType::BOOL && has_any_profile(target),
881 22 | 23 => same() && supports_base_integer(i(0), target, true),
882 26 | 27 | 32 => same() && supports_ordered(i(0), target),
883 28 => i(0) == i(1) && i(2) == DType::INT8 && supports_mul(i(0), o(0), target),
884 29 => same() && supports_float(i(0), target),
885 31 => {
886 (i(0) == DType::INT8
887 && i(1) == DType::INT8
888 && o(0) == DType::INT8
889 && has_profile(target, ProfileSet::INTEGER))
890 || (i(0) == DType::INT16
891 && i(1) == DType::INT16
892 && o(0) == DType::INT32
893 && has_ext(target, ExtensionSet::INT16))
894 }
895 35 => same() && i(0) == DType::INT32 && has_profile(target, ProfileSet::INTEGER),
896 41 => i(0) == i(1) && i(0) == i(2) && i(0) == o(0) && supports_negate(i(0), target),
897 45 => i(0) == DType::BOOL && i(1) == i(2) && i(1) == o(0) && supports_select(i(1), target),
898 46..=48 => i(0) == i(1) && o(0) == DType::BOOL && supports_ordered(i(0), target),
899 49 | 50 => same() && i(0) == DType::BOOL && has_any_profile(target),
900 51 | 52 => same() && supports_reduce_minmax(i(0), target),
901 53 => same() && supports_float(i(0), target),
902 54 => same() && supports_reduce_sum(i(0), target),
903 55 => same() && supports_concat(i(0), target),
904 56..=61 => same() && supports_data_movement(i(0), target),
905 62 => i(1) == DType::INT32 && i(0) == o(0) && supports_gather(i(0), target),
906 63 => i(1) == DType::INT32 && i(0) == i(2) && i(0) == o(0) && supports_gather(i(0), target),
907 64 => supports_resize(i(0), o(0), target),
908 65 => supports_cast(i(0), o(0), target),
909 66 => {
910 let scale32 = match attributes {
911 OpAttributes::Rescale { scale32, .. } => scale32,
912 _ => unreachable!(),
913 };
914 i(1) == if scale32 { DType::INT32 } else { DType::INT16 }
915 && i(2) == DType::INT8
916 && i(3) == i(0)
917 && i(4) == o(0)
918 && supports_rescale(i(0), o(0), target)
919 }
920 67 => supports_constant(o(0), target),
921 68 => same() && supports_constant(i(0), target),
922 69 => inputs
923 .iter()
924 .chain(outputs)
925 .all(|operand| supports_constant(tensor(*operand).dtype(), target)),
926 70 => {
927 i(0) == DType::BOOL
928 && has_ext(target, ExtensionSet::CONTROL_FLOW)
929 && inputs[1..]
930 .iter()
931 .chain(outputs)
932 .all(|operand| supports_constant(tensor(*operand).dtype(), target))
933 }
934 71 => {
935 has_ext(target, ExtensionSet::CONTROL_FLOW)
936 && inputs
937 .iter()
938 .chain(outputs)
939 .all(|operand| supports_constant(tensor(*operand).dtype(), target))
940 }
941 72 => has_ext(target, ExtensionSet::VARIABLE),
942 73 => has_ext(target, ExtensionSet::VARIABLE) && supports_variable(i(0), target),
943 74 => has_ext(target, ExtensionSet::VARIABLE) && supports_variable(o(0), target),
944 75 => has_any_profile(target),
945 _ => false,
946 };
947 if ok {
948 Ok(())
949 } else {
950 Err(SemanticErrorKind::UnsupportedTypeProfile(op))
951 }
952}
953
954fn require_rank(
955 op: Op,
956 role: OperandRole,
957 index: usize,
958 operand: Symbol<'_>,
959 minimum: usize,
960 maximum: usize,
961) -> Result<(), SemanticErrorKind> {
962 let rank = match operand {
963 Symbol::Tensor(value) => value.rank(),
964 Symbol::Shape(value) => Some(value.rank() as usize),
965 };
966 if rank.is_some_and(|rank| rank >= minimum && rank <= maximum) {
967 Ok(())
968 } else {
969 Err(SemanticErrorKind::InvalidRank {
970 op,
971 role,
972 index,
973 rank,
974 minimum,
975 maximum,
976 })
977 }
978}
979
980fn tensor_types_equal(inputs: &[Symbol<'_>], outputs: &[Symbol<'_>]) -> bool {
981 let mut operands = inputs
982 .iter()
983 .chain(outputs)
984 .filter_map(|operand| match operand {
985 Symbol::Tensor(value) => Some(value.dtype()),
986 Symbol::Shape(_) => None,
987 });
988 let Some(first) = operands.next() else {
989 return true;
990 };
991 operands.all(|dtype| dtype == first)
992}
993
994fn has_profile(target: Target, profile: ProfileSet) -> bool {
995 target.profiles.intersects(profile)
996}
997
998fn has_any_profile(target: Target) -> bool {
999 target.profiles.intersects(ProfileSet::ALL)
1000}
1001
1002fn has_ext(target: Target, extension: ExtensionSet) -> bool {
1003 target.extensions.contains(extension)
1004}
1005
1006fn supports_argmax(dtype: DType, target: Target) -> bool {
1007 match dtype {
1008 DType::INT8 => has_profile(target, ProfileSet::INTEGER),
1009 DType::INT16 => has_ext(target, ExtensionSet::INT16),
1010 DType::FP8E4M3 => has_ext(target, ExtensionSet::FP8E4M3),
1011 DType::FP8E5M2 => has_ext(target, ExtensionSet::FP8E5M2),
1012 DType::FP16 | DType::FP32 => has_profile(target, ProfileSet::FLOATING_POINT),
1013 DType::BF16 => has_ext(target, ExtensionSet::BF16),
1014 _ => false,
1015 }
1016}
1017
1018fn supports_pool(dtype: DType, acc: DType, target: Target) -> bool {
1019 match (dtype, acc) {
1020 (DType::INT8, DType::INT32) => has_profile(target, ProfileSet::INTEGER),
1021 (DType::INT16, DType::INT32) => has_ext(target, ExtensionSet::INT16),
1022 (DType::FP8E4M3, DType::FP16) => has_ext(target, ExtensionSet::FP8E4M3),
1023 (DType::FP8E5M2, DType::FP16) => has_ext(target, ExtensionSet::FP8E5M2),
1024 (DType::FP16, DType::FP16 | DType::FP32) => has_profile(target, ProfileSet::FLOATING_POINT),
1025 (DType::BF16, DType::FP32) => has_ext(target, ExtensionSet::BF16),
1026 (DType::FP32, DType::FP32) => has_profile(target, ProfileSet::FLOATING_POINT),
1027 _ => false,
1028 }
1029}
1030
1031fn supports_pool_value(dtype: DType, target: Target) -> bool {
1032 match dtype {
1033 DType::INT8 => has_profile(target, ProfileSet::INTEGER),
1034 DType::INT16 => has_ext(target, ExtensionSet::INT16),
1035 DType::FP8E4M3 => has_ext(target, ExtensionSet::FP8E4M3),
1036 DType::FP8E5M2 => has_ext(target, ExtensionSet::FP8E5M2),
1037 DType::FP16 | DType::FP32 => has_profile(target, ProfileSet::FLOATING_POINT),
1038 DType::BF16 => has_ext(target, ExtensionSet::BF16),
1039 _ => false,
1040 }
1041}
1042
1043fn conv_acc_type(attributes: OpAttributes<'_>) -> DType {
1044 match attributes {
1045 OpAttributes::Conv2d { acc_type, .. }
1046 | OpAttributes::Conv3d { acc_type, .. }
1047 | OpAttributes::DepthwiseConv2d { acc_type, .. }
1048 | OpAttributes::TransposeConv2d { acc_type, .. } => acc_type,
1049 _ => unreachable!(),
1050 }
1051}
1052
1053fn supports_conv(input: DType, weight: DType, output: DType, acc: DType, target: Target) -> bool {
1054 match (input, weight, output, acc) {
1055 (DType::INT8, DType::INT8, DType::INT32, DType::INT32) => {
1056 has_profile(target, ProfileSet::INTEGER)
1057 }
1058 (DType::INT8, DType::INT4, DType::INT32, DType::INT32) => {
1059 has_ext(target, ExtensionSet::INT4)
1060 }
1061 (DType::INT16, DType::INT8, DType::INT48, DType::INT48) => {
1062 has_ext(target, ExtensionSet::INT16)
1063 }
1064 (DType::FP8E4M3, DType::FP8E4M3, DType::FP16, DType::FP16) => {
1065 has_ext(target, ExtensionSet::FP8E4M3)
1066 }
1067 (DType::FP8E5M2, DType::FP8E5M2, DType::FP16, DType::FP16) => {
1068 has_ext(target, ExtensionSet::FP8E5M2)
1069 }
1070 (DType::FP16, DType::FP16, DType::FP16, DType::FP16 | DType::FP32) => {
1071 has_profile(target, ProfileSet::FLOATING_POINT)
1072 }
1073 (DType::BF16, DType::BF16, DType::BF16, DType::FP32) => has_ext(target, ExtensionSet::BF16),
1074 (DType::FP32, DType::FP32, DType::FP32, DType::FP32) => {
1075 has_profile(target, ProfileSet::FLOATING_POINT)
1076 }
1077 _ => false,
1078 }
1079}
1080
1081fn supports_matmul(input: DType, output: DType, target: Target) -> bool {
1082 match (input, output) {
1083 (DType::INT8, DType::INT32) => has_profile(target, ProfileSet::INTEGER),
1084 (DType::INT16, DType::INT48) => has_ext(target, ExtensionSet::INT16),
1085 (DType::FP8E4M3, DType::FP16) => has_ext(target, ExtensionSet::FP8E4M3),
1086 (DType::FP8E5M2, DType::FP16) => has_ext(target, ExtensionSet::FP8E5M2),
1087 (DType::FP16, DType::FP16 | DType::FP32) => has_profile(target, ProfileSet::FLOATING_POINT),
1088 (DType::BF16, DType::FP32) => has_ext(target, ExtensionSet::BF16),
1089 (DType::FP32, DType::FP32) => has_profile(target, ProfileSet::FLOATING_POINT),
1090 _ => false,
1091 }
1092}
1093
1094fn supports_float(dtype: DType, target: Target) -> bool {
1095 match dtype {
1096 DType::FP16 | DType::FP32 => has_profile(target, ProfileSet::FLOATING_POINT),
1097 DType::BF16 => has_ext(target, ExtensionSet::BF16),
1098 _ => false,
1099 }
1100}
1101
1102fn supports_clamp(dtype: DType, target: Target) -> bool {
1103 match dtype {
1104 DType::INT8 => has_profile(target, ProfileSet::INTEGER),
1105 DType::INT16 => has_ext(target, ExtensionSet::INT16),
1106 _ => supports_float(dtype, target),
1107 }
1108}
1109
1110fn supports_add_sub(dtype: DType, target: Target) -> bool {
1111 (dtype == DType::INT32 && has_any_profile(target)) || supports_float(dtype, target)
1112}
1113
1114fn supports_integer_bits(dtype: DType, target: Target) -> bool {
1115 matches!(dtype, DType::INT8 | DType::INT16 | DType::INT32)
1116 && has_profile(target, ProfileSet::INTEGER)
1117}
1118
1119fn supports_base_integer(dtype: DType, target: Target, either_profile: bool) -> bool {
1120 matches!(dtype, DType::INT8 | DType::INT16 | DType::INT32)
1121 && if either_profile {
1122 has_any_profile(target)
1123 } else {
1124 has_profile(target, ProfileSet::INTEGER)
1125 }
1126}
1127
1128fn supports_ordered(dtype: DType, target: Target) -> bool {
1129 (dtype == DType::INT32 && has_profile(target, ProfileSet::INTEGER))
1130 || supports_float(dtype, target)
1131}
1132
1133fn supports_mul(input: DType, output: DType, target: Target) -> bool {
1134 match (input, output) {
1135 (DType::INT8 | DType::INT16, DType::INT32) => has_profile(target, ProfileSet::INTEGER),
1136 (DType::INT32, DType::INT32) => has_any_profile(target),
1137 (DType::FP16, DType::FP16) | (DType::FP32, DType::FP32) => {
1138 has_profile(target, ProfileSet::FLOATING_POINT)
1139 }
1140 (DType::BF16, DType::BF16) => has_ext(target, ExtensionSet::BF16),
1141 _ => false,
1142 }
1143}
1144
1145fn supports_negate(dtype: DType, target: Target) -> bool {
1146 matches!(dtype, DType::INT8 | DType::INT16 | DType::INT32)
1147 && has_profile(target, ProfileSet::INTEGER)
1148 || supports_float(dtype, target)
1149}
1150
1151fn supports_select(dtype: DType, target: Target) -> bool {
1152 match dtype {
1153 DType::BOOL => has_any_profile(target),
1154 DType::INT8 | DType::INT16 | DType::INT32 => has_profile(target, ProfileSet::INTEGER),
1155 _ => supports_float(dtype, target),
1156 }
1157}
1158
1159fn supports_reduce_minmax(dtype: DType, target: Target) -> bool {
1160 matches!(dtype, DType::INT8 | DType::INT16 | DType::INT32)
1161 && has_profile(target, ProfileSet::INTEGER)
1162 || supports_float(dtype, target)
1163}
1164
1165fn supports_reduce_sum(dtype: DType, target: Target) -> bool {
1166 dtype == DType::INT32 && has_profile(target, ProfileSet::INTEGER)
1167 || supports_float(dtype, target)
1168}
1169
1170fn supports_concat(dtype: DType, target: Target) -> bool {
1171 match dtype {
1172 DType::BOOL => has_any_profile(target),
1173 DType::INT8 | DType::INT32 => has_profile(target, ProfileSet::INTEGER),
1174 DType::INT16 => has_ext(target, ExtensionSet::INT16),
1175 DType::FP8E4M3 => has_ext(target, ExtensionSet::FP8E4M3),
1176 DType::FP8E5M2 => has_ext(target, ExtensionSet::FP8E5M2),
1177 _ => supports_float(dtype, target),
1178 }
1179}
1180
1181fn supports_data_movement(dtype: DType, target: Target) -> bool {
1182 match dtype {
1183 DType::BOOL => has_any_profile(target),
1184 DType::INT8 | DType::INT16 | DType::INT32 => has_profile(target, ProfileSet::INTEGER),
1185 DType::FP8E4M3 => has_ext(target, ExtensionSet::FP8E4M3),
1186 DType::FP8E5M2 => has_ext(target, ExtensionSet::FP8E5M2),
1187 _ => supports_float(dtype, target),
1188 }
1189}
1190
1191fn supports_gather(dtype: DType, target: Target) -> bool {
1192 match dtype {
1193 DType::INT8 | DType::INT16 | DType::INT32 => has_profile(target, ProfileSet::INTEGER),
1194 DType::FP8E4M3 => has_ext(target, ExtensionSet::FP8E4M3),
1195 DType::FP8E5M2 => has_ext(target, ExtensionSet::FP8E5M2),
1196 _ => supports_float(dtype, target),
1197 }
1198}
1199
1200fn supports_resize(input: DType, output: DType, target: Target) -> bool {
1201 match (input, output) {
1202 (DType::INT8, DType::INT8 | DType::INT32) => has_profile(target, ProfileSet::INTEGER),
1203 (DType::INT16, DType::INT16 | DType::INT48) => has_ext(target, ExtensionSet::INT16),
1204 (DType::FP16, DType::FP16) | (DType::FP32, DType::FP32) => {
1205 has_profile(target, ProfileSet::FLOATING_POINT)
1206 }
1207 (DType::BF16, DType::BF16) => has_ext(target, ExtensionSet::BF16),
1208 _ => false,
1209 }
1210}
1211
1212fn supports_cast(input: DType, output: DType, target: Target) -> bool {
1213 if input == output {
1214 return false;
1215 }
1216 let integer = |dtype| {
1217 matches!(
1218 dtype,
1219 DType::BOOL | DType::INT8 | DType::INT16 | DType::INT32
1220 )
1221 };
1222 if integer(input) && integer(output) {
1223 return has_profile(target, ProfileSet::INTEGER);
1224 }
1225 match (input, output) {
1226 (DType::INT8 | DType::INT16 | DType::INT32, DType::FP16 | DType::FP32)
1227 | (DType::FP16 | DType::FP32, DType::INT8 | DType::INT16 | DType::INT32)
1228 | (DType::FP16, DType::FP32)
1229 | (DType::FP32, DType::FP16) => has_profile(target, ProfileSet::FLOATING_POINT),
1230 (DType::INT8 | DType::INT16 | DType::INT32, DType::BF16)
1231 | (DType::BF16, DType::INT8 | DType::INT16 | DType::INT32)
1232 | (DType::BF16, DType::FP8E4M3 | DType::FP8E5M2 | DType::FP32)
1233 | (DType::FP32, DType::BF16) => has_ext(target, ExtensionSet::BF16),
1234 (DType::FP8E4M3, DType::FP16 | DType::BF16 | DType::FP32)
1235 | (DType::FP16 | DType::FP32, DType::FP8E4M3) => has_ext(target, ExtensionSet::FP8E4M3),
1236 (DType::FP8E5M2, DType::FP16 | DType::BF16 | DType::FP32)
1237 | (DType::FP16 | DType::FP32, DType::FP8E5M2) => has_ext(target, ExtensionSet::FP8E5M2),
1238 _ => false,
1239 }
1240}
1241
1242fn supports_rescale(input: DType, output: DType, target: Target) -> bool {
1243 let output_ok = matches!(output, DType::INT8 | DType::INT16 | DType::INT32);
1244 if !output_ok {
1245 return false;
1246 }
1247 match input {
1248 DType::INT8 | DType::INT16 | DType::INT32 => has_profile(target, ProfileSet::INTEGER),
1249 DType::INT48 => has_ext(target, ExtensionSet::INT16),
1250 _ => false,
1251 }
1252}
1253
1254fn supports_constant(dtype: DType, target: Target) -> bool {
1255 match dtype {
1256 DType::BOOL | DType::INT8 | DType::INT16 | DType::INT32 => has_any_profile(target),
1257 DType::INT4 => has_ext(target, ExtensionSet::INT4),
1258 DType::INT48 => has_ext(target, ExtensionSet::INT16),
1259 DType::FP8E4M3 => has_ext(target, ExtensionSet::FP8E4M3),
1260 DType::FP8E5M2 => has_ext(target, ExtensionSet::FP8E5M2),
1261 DType::FP16 | DType::FP32 => has_profile(target, ProfileSet::FLOATING_POINT),
1262 DType::BF16 => has_ext(target, ExtensionSet::BF16),
1263 _ => false,
1264 }
1265}
1266
1267fn supports_variable(dtype: DType, target: Target) -> bool {
1268 match dtype {
1269 DType::INT8 => has_profile(target, ProfileSet::INTEGER),
1270 DType::FP16 | DType::FP32 => has_profile(target, ProfileSet::FLOATING_POINT),
1271 _ => false,
1272 }
1273}
1274
1275fn validate_attributes(
1276 op: Op,
1277 attributes: OpAttributes<'_>,
1278 inputs: &[Symbol<'_>],
1279 _outputs: &[Symbol<'_>],
1280 target: Target,
1281) -> Result<(), SemanticErrorKind> {
1282 let limits = target.level.limits();
1283 let input_rank = |index| tensor(inputs[index]).rank().expect("rank validated");
1284 let valid = match attributes {
1285 OpAttributes::Empty { op: attribute_op } => attribute_op == op,
1286 OpAttributes::ArgMax { axis, nan_mode } => {
1287 valid_axis(axis, input_rank(0)) && valid_nan_mode(nan_mode)
1288 }
1289 OpAttributes::AvgPool2d {
1290 kernel,
1291 stride,
1292 pad,
1293 ..
1294 }
1295 | OpAttributes::MaxPool2d {
1296 kernel,
1297 stride,
1298 pad,
1299 ..
1300 } => {
1301 list_positive_bounded(kernel, 2, limits.max_kernel)
1302 && list_positive_bounded(stride, 2, limits.max_stride)
1303 && list_nonnegative_bounded(pad, 4, limits.max_kernel)
1304 && pool_padding_valid(kernel, pad)
1305 && match attributes {
1306 OpAttributes::MaxPool2d { nan_mode, .. } => valid_nan_mode(nan_mode),
1307 _ => true,
1308 }
1309 }
1310 OpAttributes::Conv2d {
1311 pad,
1312 stride,
1313 dilation,
1314 ..
1315 }
1316 | OpAttributes::DepthwiseConv2d {
1317 pad,
1318 stride,
1319 dilation,
1320 ..
1321 } => {
1322 list_nonnegative_bounded(pad, 4, limits.max_kernel)
1323 && list_positive_bounded(stride, 2, limits.max_stride)
1324 && list_positive_bounded(dilation, 2, limits.max_stride)
1325 }
1326 OpAttributes::Conv3d {
1327 pad,
1328 stride,
1329 dilation,
1330 ..
1331 } => {
1332 list_nonnegative_bounded(pad, 6, limits.max_kernel)
1333 && list_positive_bounded(stride, 3, limits.max_stride)
1334 && list_positive_bounded(dilation, 3, limits.max_stride)
1335 }
1336 OpAttributes::Fft2d { .. } | OpAttributes::Rfft2d { .. } => true,
1337 OpAttributes::TransposeConv2d {
1338 out_pad, stride, ..
1339 } => {
1340 list_exact(out_pad, 4)
1341 && out_pad.iter().all(|value| value <= limits.max_kernel)
1342 && list_positive_bounded(stride, 2, limits.max_stride)
1343 }
1344 OpAttributes::Clamp {
1345 min_val,
1346 max_val,
1347 nan_mode,
1348 } => {
1349 valid_nan_mode(nan_mode)
1350 && scalar_bytes_valid(tensor(inputs[0]).dtype(), min_val, max_val)
1351 }
1352 OpAttributes::ArithmeticRightShift { .. } => true,
1353 OpAttributes::Maximum { nan_mode } | OpAttributes::Minimum { nan_mode } => {
1354 valid_nan_mode(nan_mode)
1355 }
1356 OpAttributes::ReduceAll { axis }
1357 | OpAttributes::ReduceAny { axis }
1358 | OpAttributes::ReduceProduct { axis }
1359 | OpAttributes::ReduceSum { axis }
1360 | OpAttributes::Concat { axis }
1361 | OpAttributes::Reverse { axis } => valid_axis(axis, input_rank(0)),
1362 OpAttributes::ReduceMax { axis, nan_mode } | OpAttributes::ReduceMin { axis, nan_mode } => {
1363 valid_axis(axis, input_rank(0)) && valid_nan_mode(nan_mode)
1364 }
1365 OpAttributes::Transpose { perms } => valid_permutation(perms, input_rank(0)),
1366 OpAttributes::Resize { mode } => valid_resize_mode(mode),
1367 OpAttributes::Rescale {
1368 scale32,
1369 rounding_mode,
1370 per_channel,
1371 input_unsigned,
1372 output_unsigned,
1373 } => {
1374 let input = tensor(inputs[0]).dtype();
1375 let output = match _outputs[0] {
1376 Symbol::Tensor(value) => value.dtype(),
1377 Symbol::Shape(_) => unreachable!(),
1378 };
1379 valid_rounding_mode(rounding_mode)
1380 && (rounding_mode != RoundingMode::DOUBLE_ROUND
1381 || target.extensions.contains(ExtensionSet::DOUBLE_ROUND))
1382 && (rounding_mode != RoundingMode::INEXACT_ROUND
1383 || target.extensions.contains(ExtensionSet::INEXACT_ROUND))
1384 && !(scale32 && input == DType::INT48)
1385 && (scale32 || rounding_mode != RoundingMode::DOUBLE_ROUND)
1386 && !(input_unsigned && output_unsigned)
1387 && !(output == DType::INT32 && input_unsigned)
1388 && !(matches!(input, DType::INT32 | DType::INT48) && input_unsigned)
1389 && !(matches!(input, DType::INT32 | DType::INT48) && output_unsigned)
1390 && !(output == DType::INT32 && output_unsigned)
1391 && (!per_channel || input_rank(0) >= 1)
1392 }
1393 OpAttributes::Custom {
1394 operator_name,
1395 domain_name,
1396 ..
1397 } => {
1398 operator_name.is_some_and(|name| !name.is_empty())
1399 && domain_name.is_some_and(|name| !name.is_empty())
1400 }
1401 OpAttributes::CondIf {
1402 then_graph,
1403 else_graph,
1404 } => {
1405 then_graph.is_some_and(|name| !name.is_empty())
1406 && else_graph.is_some_and(|name| !name.is_empty())
1407 }
1408 OpAttributes::WhileLoop {
1409 cond_graph,
1410 body_graph,
1411 } => {
1412 cond_graph.is_some_and(|name| !name.is_empty())
1413 && body_graph.is_some_and(|name| !name.is_empty())
1414 }
1415 };
1416 if valid {
1417 Ok(())
1418 } else {
1419 Err(SemanticErrorKind::InvalidAttribute(op))
1420 }
1421}
1422
1423fn validate_ctc_inputs(
1424 model: &Model<'_>,
1425 op: Op,
1426 operator: Operator<'_>,
1427 attributes: OpAttributes<'_>,
1428 inputs: &[Symbol<'_>],
1429 constants: &[&str],
1430 target: Target,
1431) -> Result<(), SemanticErrorKind> {
1432 let required: &[usize] = match op.get() {
1433 2 => &[1, 2],
1434 3..=5 | 10 => &[3, 4],
1435 7 => &[2, 3],
1436 28 => &[2],
1437 31 => &[1],
1438 41 => &[1, 2],
1439 56 => &[1, 2],
1440 57 => &[1],
1441 59 => &[1, 2],
1442 60 => &[1],
1443 64 => &[1, 2, 3],
1444 66 => &[1, 2, 3, 4],
1445 _ => &[],
1446 };
1447 let dynamic = target.extensions.contains(ExtensionSet::DYNAMIC);
1448 for &index in required {
1449 let connected_to_constant = operator
1450 .inputs()
1451 .nth(index)
1452 .is_some_and(|name| constants.binary_search(&name).is_ok());
1453 let present = match inputs[index] {
1454 Symbol::Tensor(value) => !tensor_data(model, value).is_empty(),
1455 Symbol::Shape(value) => value.values().is_some(),
1456 };
1457 if (!connected_to_constant || !present) && !dynamic {
1458 return Err(SemanticErrorKind::ConstantRequired { op, input: index });
1459 }
1460 }
1461
1462 let check_zero_point = |index: usize, unsigned: bool| -> bool {
1463 let value = tensor(inputs[index]);
1464 if tensor_data(model, value).is_empty() {
1465 return dynamic;
1466 }
1467 value.dtype() == DType::INT8
1468 || (value.dtype() == DType::INT16
1469 && unsigned
1470 && tensor_data(model, value)
1471 .get(..2)
1472 .and_then(|bytes| <[u8; 2]>::try_from(bytes).ok())
1473 .is_some_and(|bytes| matches!(u16::from_le_bytes(bytes), 0 | 32_768)))
1474 || constant_is_zero(model, value, 0)
1475 };
1476 let zero_points: &[usize] = match op.get() {
1477 2 => &[1, 2],
1478 3..=5 | 10 => &[3, 4],
1479 7 => &[2, 3],
1480 41 => &[1, 2],
1481 56 => &[2],
1482 _ => &[],
1483 };
1484 if zero_points
1485 .iter()
1486 .any(|index| !check_zero_point(*index, false))
1487 {
1488 return Err(SemanticErrorKind::InvalidConstantData {
1489 op,
1490 operand: *zero_points.first().unwrap_or(&0),
1491 });
1492 }
1493
1494 if op == Op::MUL {
1495 let Some(shift) = constant_integer(model, tensor(inputs[2]), 0) else {
1496 return if dynamic {
1497 Ok(())
1498 } else {
1499 Err(SemanticErrorKind::ConstantRequired { op, input: 2 })
1500 };
1501 };
1502 if !(0..=63).contains(&shift) || (tensor(inputs[0]).dtype() != DType::INT32 && shift != 0) {
1503 return Err(SemanticErrorKind::InvalidConstantData { op, operand: 2 });
1504 }
1505 }
1506
1507 if op == Op::RESCALE {
1508 let OpAttributes::Rescale {
1509 per_channel,
1510 input_unsigned,
1511 output_unsigned,
1512 ..
1513 } = attributes
1514 else {
1515 unreachable!()
1516 };
1517 if !check_zero_point(3, input_unsigned) || !check_zero_point(4, output_unsigned) {
1518 return Err(SemanticErrorKind::InvalidConstantData { op, operand: 3 });
1519 }
1520 let input = tensor(inputs[0]);
1521 let channels = if per_channel {
1522 dimension(input, input.rank().expect("rank validated") - 1) as usize
1523 } else {
1524 1
1525 };
1526 for index in [1, 2] {
1527 let value = tensor(inputs[index]);
1528 if dimension(value, 0) as usize != channels {
1529 return Err(SemanticErrorKind::InvalidConstantData { op, operand: index });
1530 }
1531 }
1532 if !tensor_data(model, tensor(inputs[1])).is_empty() {
1533 for index in 0..channels {
1534 if constant_integer(model, tensor(inputs[1]), index).is_none_or(|value| value < 0) {
1535 return Err(SemanticErrorKind::InvalidConstantData { op, operand: 1 });
1536 }
1537 }
1538 }
1539 if !tensor_data(model, tensor(inputs[2])).is_empty() {
1540 for index in 0..channels {
1541 if constant_integer(model, tensor(inputs[2]), index)
1542 .is_none_or(|value| !(2..=62).contains(&value))
1543 {
1544 return Err(SemanticErrorKind::InvalidConstantData { op, operand: 2 });
1545 }
1546 }
1547 }
1548 }
1549 Ok(())
1550}
1551
1552fn valid_axis(axis: i32, rank: usize) -> bool {
1553 axis >= 0 && (axis as usize) < rank
1554}
1555
1556fn list_positive_bounded(values: I32List<'_>, length: usize, maximum: i32) -> bool {
1557 list_exact(values, length) && values.iter().all(|value| value >= 1 && value <= maximum)
1558}
1559
1560fn list_nonnegative_bounded(values: I32List<'_>, length: usize, maximum: i32) -> bool {
1561 list_exact(values, length) && values.iter().all(|value| value >= 0 && value <= maximum)
1562}
1563
1564fn pool_padding_valid(kernel: I32List<'_>, pad: I32List<'_>) -> bool {
1565 pad.get(0)
1566 .is_some_and(|value| value < kernel.get(0).unwrap_or(0))
1567 && pad
1568 .get(1)
1569 .is_some_and(|value| value < kernel.get(0).unwrap_or(0))
1570 && pad
1571 .get(2)
1572 .is_some_and(|value| value < kernel.get(1).unwrap_or(0))
1573 && pad
1574 .get(3)
1575 .is_some_and(|value| value < kernel.get(1).unwrap_or(0))
1576}
1577
1578fn valid_permutation(perms: I32List<'_>, rank: usize) -> bool {
1579 if perms.len() != rank {
1580 return false;
1581 }
1582 for index in 0..rank {
1583 let Some(value) = perms.get(index) else {
1584 return false;
1585 };
1586 if !valid_axis(value, rank) || (0..index).any(|prior| perms.get(prior) == Some(value)) {
1587 return false;
1588 }
1589 }
1590 true
1591}
1592
1593fn scalar_bytes_valid(dtype: DType, minimum: &[u8], maximum: &[u8]) -> bool {
1594 let Some(width) = dtype_width(dtype) else {
1595 return false;
1596 };
1597 if minimum.len() != width || maximum.len() != width {
1598 return false;
1599 }
1600 match dtype {
1601 DType::INT8 => (minimum[0] as i8) <= (maximum[0] as i8),
1602 DType::INT16 => {
1603 i16::from_le_bytes([minimum[0], minimum[1]])
1604 <= i16::from_le_bytes([maximum[0], maximum[1]])
1605 }
1606 DType::FP16 => {
1607 let min = f16_to_f32(u16::from_le_bytes([minimum[0], minimum[1]]));
1608 let max = f16_to_f32(u16::from_le_bytes([maximum[0], maximum[1]]));
1609 !min.is_nan() && !max.is_nan() && min <= max
1610 }
1611 DType::BF16 => {
1612 let min = f32::from_bits(u32::from(u16::from_le_bytes([minimum[0], minimum[1]])) << 16);
1613 let max = f32::from_bits(u32::from(u16::from_le_bytes([maximum[0], maximum[1]])) << 16);
1614 !min.is_nan() && !max.is_nan() && min <= max
1615 }
1616 DType::FP32 => {
1617 let min = f32::from_le_bytes(minimum.try_into().expect("length validated"));
1618 let max = f32::from_le_bytes(maximum.try_into().expect("length validated"));
1619 !min.is_nan() && !max.is_nan() && min <= max
1620 }
1621 _ => false,
1622 }
1623}
1624
1625fn f16_to_f32(bits: u16) -> f32 {
1626 let sign = u32::from(bits & 0x8000) << 16;
1627 let exponent = (bits >> 10) & 0x1f;
1628 let fraction = u32::from(bits & 0x03ff);
1629 let converted = match exponent {
1630 0 if fraction == 0 => sign,
1631 0 => {
1632 let shift = fraction.leading_zeros() - 21;
1633 let normalized = fraction << shift;
1634 sign | ((127 - 15 - shift + 1) << 23) | ((normalized & 0x03ff) << 13)
1635 }
1636 0x1f => sign | 0x7f80_0000 | (fraction << 13),
1637 _ => sign | ((u32::from(exponent) + 127 - 15) << 23) | (fraction << 13),
1638 };
1639 f32::from_bits(converted)
1640}
1641
1642fn constant_integer(model: &Model<'_>, value: Tensor<'_>, index: usize) -> Option<i64> {
1643 let data = tensor_data(model, value);
1644 let start = index.checked_mul(dtype_width(value.dtype())?)?;
1645 match value.dtype() {
1646 DType::INT4 => Some(i64::from(crate::unpack_int4(data, index)?)),
1647 DType::INT8 => Some(i64::from(*data.get(start)? as i8)),
1648 DType::INT16 => Some(i64::from(i16::from_le_bytes(
1649 data.get(start..start + 2)?.try_into().ok()?,
1650 ))),
1651 DType::INT32 => Some(i64::from(i32::from_le_bytes(
1652 data.get(start..start + 4)?.try_into().ok()?,
1653 ))),
1654 DType::INT48 => Some(i64::from_le_bytes(
1655 data.get(start..start + 8)?.try_into().ok()?,
1656 )),
1657 _ => None,
1658 }
1659}
1660
1661fn constant_is_zero(model: &Model<'_>, value: Tensor<'_>, index: usize) -> bool {
1662 let data = tensor_data(model, value);
1663 let Some(width) = dtype_width(value.dtype()) else {
1664 return false;
1665 };
1666 let Some(start) = index.checked_mul(width) else {
1667 return false;
1668 };
1669 match value.dtype() {
1670 DType::INT4 | DType::INT8 | DType::INT16 | DType::INT32 | DType::INT48 => {
1671 constant_integer(model, value, index) == Some(0)
1672 }
1673 DType::FP8E4M3 | DType::FP8E5M2 => data.get(start).is_some_and(|bits| bits & 0x7f == 0),
1674 DType::FP16 | DType::BF16 => data
1675 .get(start..start + 2)
1676 .and_then(|bytes| <[u8; 2]>::try_from(bytes).ok())
1677 .is_some_and(|bytes| u16::from_le_bytes(bytes) & 0x7fff == 0),
1678 DType::FP32 => data
1679 .get(start..start + 4)
1680 .and_then(|bytes| <[u8; 4]>::try_from(bytes).ok())
1681 .is_some_and(|bytes| u32::from_le_bytes(bytes) & 0x7fff_ffff == 0),
1682 _ => false,
1683 }
1684}
1685
1686fn validate_shapes(
1687 model: &Model<'_>,
1688 op: Op,
1689 attributes: OpAttributes<'_>,
1690 inputs: &[Symbol<'_>],
1691 outputs: &[Symbol<'_>],
1692 target: Target,
1693 region_names: &[&str],
1694) -> Result<(), SemanticErrorKind> {
1695 let i = |index| tensor(inputs[index]);
1696 let o = |index| tensor(outputs[index]);
1697 let dynamic = target.extensions.contains(ExtensionSet::DYNAMIC);
1698 let valid = match op.get() {
1699 1 => {
1700 let OpAttributes::ArgMax { axis, .. } = attributes else {
1701 unreachable!()
1702 };
1703 shape_without_axis(i(0), o(0), axis as usize)
1704 }
1705 2 => {
1706 let OpAttributes::AvgPool2d {
1707 kernel,
1708 stride,
1709 pad,
1710 ..
1711 } = attributes
1712 else {
1713 unreachable!()
1714 };
1715 pool_shape(i(0), o(0), kernel, stride, pad) && scalar_shape(i(1)) && scalar_shape(i(2))
1716 }
1717 3 | 5 => {
1718 let (pad, stride, dilation) = match attributes {
1719 OpAttributes::Conv2d {
1720 pad,
1721 stride,
1722 dilation,
1723 ..
1724 }
1725 | OpAttributes::DepthwiseConv2d {
1726 pad,
1727 stride,
1728 dilation,
1729 ..
1730 } => (pad, stride, dilation),
1731 _ => unreachable!(),
1732 };
1733 conv2d_shape(
1734 Conv2dGeometry {
1735 op,
1736 input: i(0),
1737 weight: i(1),
1738 bias: i(2),
1739 output: o(0),
1740 pad,
1741 stride,
1742 dilation,
1743 },
1744 target,
1745 ) && scalar_shape(i(3))
1746 && scalar_shape(i(4))
1747 }
1748 4 => {
1749 let OpAttributes::Conv3d {
1750 pad,
1751 stride,
1752 dilation,
1753 ..
1754 } = attributes
1755 else {
1756 unreachable!()
1757 };
1758 conv3d_shape(
1759 Conv3dGeometry {
1760 input: i(0),
1761 weight: i(1),
1762 bias: i(2),
1763 output: o(0),
1764 pad,
1765 stride,
1766 dilation,
1767 },
1768 target,
1769 ) && scalar_shape(i(3))
1770 && scalar_shape(i(4))
1771 }
1772 6 => {
1773 same_shape(i(0), i(1))
1774 && same_shape(i(0), o(0))
1775 && same_shape(i(0), o(1))
1776 && fft_shape(i(0), target.level.limits())
1777 }
1778 7 => {
1779 dimension(i(0), 0) == dimension(i(1), 0)
1780 && dimension(i(0), 2) == dimension(i(1), 1)
1781 && dimensions_equal(
1782 o(0),
1783 &[dimension(i(0), 0), dimension(i(0), 1), dimension(i(1), 2)],
1784 )
1785 && scalar_shape(i(2))
1786 && scalar_shape(i(3))
1787 }
1788 8 => {
1789 let OpAttributes::MaxPool2d {
1790 kernel,
1791 stride,
1792 pad,
1793 ..
1794 } = attributes
1795 else {
1796 unreachable!()
1797 };
1798 pool_shape(i(0), o(0), kernel, stride, pad)
1799 }
1800 9 => {
1801 same_prefix(i(0), o(0), 2)
1802 && same_shape(o(0), o(1))
1803 && dimension(o(0), 2) == dimension(i(0), 2) / 2 + 1
1804 && fft_shape(i(0), target.level.limits())
1805 }
1806 10 => {
1807 let OpAttributes::TransposeConv2d {
1808 out_pad, stride, ..
1809 } = attributes
1810 else {
1811 unreachable!()
1812 };
1813 transpose_conv2d_shape(i(0), i(1), i(2), o(0), out_pad, stride, target)
1814 && scalar_shape(i(3))
1815 && scalar_shape(i(4))
1816 }
1817 11..=14 | 32..=44 | 58 | 65 | 66 | 68 => same_shape(i(0), o(0)),
1818 15..=27 | 29 | 30 => broadcast_shape(&[i(0), i(1)], o(0)),
1819 28 => broadcast_shape(&[i(0), i(1)], o(0)) && scalar_shape(i(2)),
1820 31 => {
1821 same_shape(i(0), o(0))
1822 && dimension(i(1), 0)
1823 == if i(0).dtype() == DType::INT8 {
1824 256
1825 } else {
1826 513
1827 }
1828 }
1829 45 => broadcast_shape(&[i(0), i(1), i(2)], o(0)),
1830 46..=48 => broadcast_shape(&[i(0), i(1)], o(0)),
1831 49..=54 => {
1832 let axis = reduction_axis(attributes);
1833 reduced_shape(i(0), o(0), axis as usize)
1834 }
1835 55 => {
1836 let OpAttributes::Concat { axis } = attributes else {
1837 unreachable!()
1838 };
1839 concat_shape(inputs, o(0), axis as usize)
1840 }
1841 56 => pad_shape(i(0), shape(inputs[1]), i(2), o(0), dynamic),
1842 57 => reshape_shape(i(0), shape(inputs[1]), o(0), dynamic),
1843 59 => slice_shape(i(0), shape(inputs[1]), shape(inputs[2]), o(0), dynamic),
1844 60 => tile_shape(i(0), shape(inputs[1]), o(0), dynamic),
1845 61 => {
1846 let OpAttributes::Transpose { perms } = attributes else {
1847 unreachable!()
1848 };
1849 transpose_shape(i(0), o(0), perms)
1850 }
1851 62 => {
1852 dimensions_equal(
1853 o(0),
1854 &[dimension(i(0), 0), dimension(i(1), 1), dimension(i(0), 2)],
1855 ) && dimension(i(0), 0) == dimension(i(1), 0)
1856 }
1857 63 => {
1858 same_shape(i(0), o(0))
1859 && dimension(i(0), 0) == dimension(i(1), 0)
1860 && dimensions_equal(
1861 i(2),
1862 &[dimension(i(1), 0), dimension(i(1), 1), dimension(i(0), 2)],
1863 )
1864 }
1865 64 => resize_shape(
1866 i(0),
1867 shape(inputs[1]),
1868 shape(inputs[2]),
1869 shape(inputs[3]),
1870 o(0),
1871 target.level.limits(),
1872 dynamic,
1873 ),
1874 67 => !tensor_data(model, o(0)).is_empty(),
1875 69 => true,
1876 70 => {
1877 let OpAttributes::CondIf {
1878 then_graph,
1879 else_graph,
1880 } = attributes
1881 else {
1882 unreachable!()
1883 };
1884 tensor_elements(i(0)) == Some(1)
1885 && control_region_exists(region_names, then_graph)
1886 && control_region_exists(region_names, else_graph)
1887 && region_matches(model, then_graph.unwrap(), &inputs[1..], outputs)
1888 && region_matches(model, else_graph.unwrap(), &inputs[1..], outputs)
1889 }
1890 71 => {
1891 let OpAttributes::WhileLoop {
1892 cond_graph,
1893 body_graph,
1894 } = attributes
1895 else {
1896 unreachable!()
1897 };
1898 control_region_exists(region_names, cond_graph)
1899 && control_region_exists(region_names, body_graph)
1900 && tensor_lists_same(inputs, outputs)
1901 && condition_region_matches(model, cond_graph.unwrap(), inputs)
1902 && region_matches(model, body_graph.unwrap(), inputs, outputs)
1903 }
1904 72 => true,
1905 73 => i(0).is_variable() && i(0).variable_name().is_some(),
1906 74 => o(0).is_variable() && o(0).variable_name().is_some(),
1907 75 => shape(outputs[0]).values().is_some(),
1908 _ => false,
1909 };
1910 if valid {
1911 Ok(())
1912 } else if matches!(op, Op::COND_IF | Op::WHILE_LOOP)
1913 && match attributes {
1914 OpAttributes::CondIf {
1915 then_graph,
1916 else_graph,
1917 } => {
1918 !control_region_exists(region_names, then_graph)
1919 || !control_region_exists(region_names, else_graph)
1920 }
1921 OpAttributes::WhileLoop {
1922 cond_graph,
1923 body_graph,
1924 } => {
1925 !control_region_exists(region_names, cond_graph)
1926 || !control_region_exists(region_names, body_graph)
1927 }
1928 _ => false,
1929 }
1930 {
1931 Err(SemanticErrorKind::UnknownControlFlowRegion(op))
1932 } else if matches!(op, Op::COND_IF | Op::WHILE_LOOP) {
1933 Err(SemanticErrorKind::ControlFlowSignature(op))
1934 } else {
1935 Err(SemanticErrorKind::InvalidShape(op))
1936 }
1937}
1938
1939fn dimension(tensor: Tensor<'_>, index: usize) -> i32 {
1940 tensor
1941 .dimensions()
1942 .nth(index)
1943 .expect("rank and dimension validated")
1944}
1945
1946fn dimensions_equal(tensor: Tensor<'_>, expected: &[i32]) -> bool {
1947 tensor.rank() == Some(expected.len()) && tensor.dimensions().eq(expected.iter().copied())
1948}
1949
1950fn same_shape(left: Tensor<'_>, right: Tensor<'_>) -> bool {
1951 left.rank() == right.rank() && left.dimensions().eq(right.dimensions())
1952}
1953
1954fn same_prefix(left: Tensor<'_>, right: Tensor<'_>, length: usize) -> bool {
1955 (0..length).all(|index| dimension(left, index) == dimension(right, index))
1956}
1957
1958fn scalar_shape(value: Tensor<'_>) -> bool {
1959 dimensions_equal(value, &[1])
1960}
1961
1962fn shape_without_axis(input: Tensor<'_>, output: Tensor<'_>, axis: usize) -> bool {
1963 output.rank() == input.rank().map(|rank| rank - 1)
1964 && input
1965 .dimensions()
1966 .enumerate()
1967 .filter_map(|(index, value)| (index != axis).then_some(value))
1968 .eq(output.dimensions())
1969}
1970
1971fn checked_window_output(
1972 input: i32,
1973 before: i32,
1974 after: i32,
1975 kernel: i64,
1976 stride: i32,
1977) -> Option<i32> {
1978 let numerator = i64::from(input) + i64::from(before) + i64::from(after) - kernel;
1979 (numerator >= 0 && numerator % i64::from(stride) == 0)
1980 .then(|| numerator / i64::from(stride) + 1)
1981 .and_then(|value| i32::try_from(value).ok())
1982}
1983
1984fn pool_shape(
1985 input: Tensor<'_>,
1986 output: Tensor<'_>,
1987 kernel: I32List<'_>,
1988 stride: I32List<'_>,
1989 pad: I32List<'_>,
1990) -> bool {
1991 let Some(height) = checked_window_output(
1992 dimension(input, 1),
1993 pad.get(0).unwrap(),
1994 pad.get(1).unwrap(),
1995 i64::from(kernel.get(0).unwrap()),
1996 stride.get(0).unwrap(),
1997 ) else {
1998 return false;
1999 };
2000 let Some(width) = checked_window_output(
2001 dimension(input, 2),
2002 pad.get(2).unwrap(),
2003 pad.get(3).unwrap(),
2004 i64::from(kernel.get(1).unwrap()),
2005 stride.get(1).unwrap(),
2006 ) else {
2007 return false;
2008 };
2009 dimensions_equal(
2010 output,
2011 &[dimension(input, 0), height, width, dimension(input, 3)],
2012 )
2013}
2014
2015fn effective_kernel(size: i32, dilation: i32) -> Option<i64> {
2016 i64::from(size)
2017 .checked_sub(1)?
2018 .checked_mul(i64::from(dilation))?
2019 .checked_add(1)
2020}
2021
2022fn kernel_within_level(weight: Tensor<'_>, indexes: &[usize], limits: LevelLimits) -> bool {
2023 indexes.iter().all(|index| {
2024 let value = dimension(weight, *index);
2025 value <= limits.max_kernel
2026 })
2027}
2028
2029fn dilated_kernel_within_level(
2030 weight: Tensor<'_>,
2031 indexes: &[usize],
2032 dilation: I32List<'_>,
2033 limits: LevelLimits,
2034) -> bool {
2035 indexes.iter().enumerate().all(|(dilation_index, index)| {
2036 i64::from(dimension(weight, *index)) * i64::from(dilation.get(dilation_index).unwrap())
2037 <= i64::from(limits.max_kernel)
2038 })
2039}
2040
2041struct Conv2dGeometry<'a> {
2042 op: Op,
2043 input: Tensor<'a>,
2044 weight: Tensor<'a>,
2045 bias: Tensor<'a>,
2046 output: Tensor<'a>,
2047 pad: I32List<'a>,
2048 stride: I32List<'a>,
2049 dilation: I32List<'a>,
2050}
2051
2052fn conv2d_shape(geometry: Conv2dGeometry<'_>, target: Target) -> bool {
2053 let Conv2dGeometry {
2054 op,
2055 input,
2056 weight,
2057 bias,
2058 output,
2059 pad,
2060 stride,
2061 dilation,
2062 } = geometry;
2063 let limits = target.level.limits();
2064 if !dilated_kernel_within_level(weight, &[0, 1], dilation, limits) && op == Op::DEPTHWISE_CONV2D
2065 || !dilated_kernel_within_level(weight, &[1, 2], dilation, limits) && op == Op::CONV2D
2066 {
2067 return false;
2068 }
2069 let (kernel_h, kernel_w, channels, output_channels) = if op == Op::DEPTHWISE_CONV2D {
2070 let channels = dimension(weight, 2);
2071 let Some(output_channels) = channels.checked_mul(dimension(weight, 3)) else {
2072 return false;
2073 };
2074 (
2075 dimension(weight, 0),
2076 dimension(weight, 1),
2077 channels,
2078 output_channels,
2079 )
2080 } else {
2081 (
2082 dimension(weight, 1),
2083 dimension(weight, 2),
2084 dimension(weight, 3),
2085 dimension(weight, 0),
2086 )
2087 };
2088 let Some(height) = effective_kernel(kernel_h, dilation.get(0).unwrap()).and_then(|kernel| {
2089 checked_window_output(
2090 dimension(input, 1),
2091 pad.get(0).unwrap(),
2092 pad.get(1).unwrap(),
2093 kernel,
2094 stride.get(0).unwrap(),
2095 )
2096 }) else {
2097 return false;
2098 };
2099 let Some(width) = effective_kernel(kernel_w, dilation.get(1).unwrap()).and_then(|kernel| {
2100 checked_window_output(
2101 dimension(input, 2),
2102 pad.get(2).unwrap(),
2103 pad.get(3).unwrap(),
2104 kernel,
2105 stride.get(1).unwrap(),
2106 )
2107 }) else {
2108 return false;
2109 };
2110 dimension(input, 3) == channels
2111 && (dimension(bias, 0) == 1 || dimension(bias, 0) == output_channels)
2112 && dimensions_equal(
2113 output,
2114 &[dimension(input, 0), height, width, output_channels],
2115 )
2116}
2117
2118struct Conv3dGeometry<'a> {
2119 input: Tensor<'a>,
2120 weight: Tensor<'a>,
2121 bias: Tensor<'a>,
2122 output: Tensor<'a>,
2123 pad: I32List<'a>,
2124 stride: I32List<'a>,
2125 dilation: I32List<'a>,
2126}
2127
2128fn conv3d_shape(geometry: Conv3dGeometry<'_>, target: Target) -> bool {
2129 let Conv3dGeometry {
2130 input,
2131 weight,
2132 bias,
2133 output,
2134 pad,
2135 stride,
2136 dilation,
2137 } = geometry;
2138 if !dilated_kernel_within_level(weight, &[1, 2, 3], dilation, target.level.limits()) {
2139 return false;
2140 }
2141 let mut spatial = [0_i32; 3];
2142 for (index, spatial_dimension) in spatial.iter_mut().enumerate() {
2143 let Some(value) =
2144 effective_kernel(dimension(weight, index + 1), dilation.get(index).unwrap()).and_then(
2145 |kernel| {
2146 checked_window_output(
2147 dimension(input, index + 1),
2148 pad.get(index * 2).unwrap(),
2149 pad.get(index * 2 + 1).unwrap(),
2150 kernel,
2151 stride.get(index).unwrap(),
2152 )
2153 },
2154 )
2155 else {
2156 return false;
2157 };
2158 *spatial_dimension = value;
2159 }
2160 let channels = dimension(weight, 4);
2161 let output_channels = dimension(weight, 0);
2162 dimension(input, 4) == channels
2163 && (dimension(bias, 0) == 1 || dimension(bias, 0) == output_channels)
2164 && dimensions_equal(
2165 output,
2166 &[
2167 dimension(input, 0),
2168 spatial[0],
2169 spatial[1],
2170 spatial[2],
2171 output_channels,
2172 ],
2173 )
2174}
2175
2176fn fft_shape(input: Tensor<'_>, limits: LevelLimits) -> bool {
2177 [dimension(input, 1), dimension(input, 2)]
2178 .into_iter()
2179 .all(|value| {
2180 let value = value as u32;
2181 value <= limits.max_kernel as u32
2182 && value.is_power_of_two()
2183 && value.ilog2() <= limits.max_log2_size
2184 })
2185}
2186
2187fn transpose_conv2d_shape(
2188 input: Tensor<'_>,
2189 weight: Tensor<'_>,
2190 bias: Tensor<'_>,
2191 output: Tensor<'_>,
2192 out_pad: I32List<'_>,
2193 stride: I32List<'_>,
2194 target: Target,
2195) -> bool {
2196 if !kernel_within_level(weight, &[1, 2], target.level.limits()) {
2197 return false;
2198 }
2199 let kernel_h = dimension(weight, 1);
2200 let kernel_w = dimension(weight, 2);
2201 if out_pad.get(0).unwrap() <= -kernel_h
2202 || out_pad.get(1).unwrap() <= -kernel_h
2203 || out_pad.get(2).unwrap() <= -kernel_w
2204 || out_pad.get(3).unwrap() <= -kernel_w
2205 {
2206 return false;
2207 }
2208 let calculate = |size: i32, stride: i32, before: i32, after: i32, kernel: i32| {
2209 i64::from(size - 1)
2210 .checked_mul(i64::from(stride))?
2211 .checked_add(i64::from(before))?
2212 .checked_add(i64::from(after))?
2213 .checked_add(i64::from(kernel))
2214 .and_then(|value| i32::try_from(value).ok())
2215 };
2216 let Some(height) = calculate(
2217 dimension(input, 1),
2218 stride.get(0).unwrap(),
2219 out_pad.get(0).unwrap(),
2220 out_pad.get(1).unwrap(),
2221 kernel_h,
2222 ) else {
2223 return false;
2224 };
2225 let Some(width) = calculate(
2226 dimension(input, 2),
2227 stride.get(1).unwrap(),
2228 out_pad.get(2).unwrap(),
2229 out_pad.get(3).unwrap(),
2230 kernel_w,
2231 ) else {
2232 return false;
2233 };
2234 let output_channels = dimension(weight, 0);
2235 dimension(input, 3) == dimension(weight, 3)
2236 && (dimension(bias, 0) == 1 || dimension(bias, 0) == output_channels)
2237 && dimensions_equal(
2238 output,
2239 &[dimension(input, 0), height, width, output_channels],
2240 )
2241}
2242
2243fn broadcast_shape(inputs: &[Tensor<'_>], output: Tensor<'_>) -> bool {
2244 let output_rank = output.rank().expect("rank validated");
2245 if inputs
2246 .iter()
2247 .any(|input| input.rank().expect("rank validated") > output_rank)
2248 {
2249 return false;
2250 }
2251 for output_axis in 0..output_rank {
2252 let output_dimension = dimension(output, output_axis);
2253 let mut expected = 1;
2254 for input in inputs {
2255 let input_rank = input.rank().unwrap();
2256 if output_axis + input_rank >= output_rank {
2257 let input_axis = output_axis + input_rank - output_rank;
2258 let value = dimension(*input, input_axis);
2259 if value != 1 && expected != 1 && value != expected {
2260 return false;
2261 }
2262 expected = expected.max(value);
2263 }
2264 }
2265 if output_dimension != expected {
2266 return false;
2267 }
2268 }
2269 true
2270}
2271
2272fn reduction_axis(attributes: OpAttributes<'_>) -> i32 {
2273 match attributes {
2274 OpAttributes::ReduceAll { axis }
2275 | OpAttributes::ReduceAny { axis }
2276 | OpAttributes::ReduceMax { axis, .. }
2277 | OpAttributes::ReduceMin { axis, .. }
2278 | OpAttributes::ReduceProduct { axis }
2279 | OpAttributes::ReduceSum { axis } => axis,
2280 _ => unreachable!(),
2281 }
2282}
2283
2284fn reduced_shape(input: Tensor<'_>, output: Tensor<'_>, axis: usize) -> bool {
2285 input.rank() == output.rank()
2286 && input
2287 .dimensions()
2288 .enumerate()
2289 .all(|(index, value)| dimension(output, index) == if index == axis { 1 } else { value })
2290}
2291
2292fn concat_shape(inputs: &[Symbol<'_>], output: Tensor<'_>, axis: usize) -> bool {
2293 let rank = output.rank().unwrap();
2294 let mut axis_size = 0_i64;
2295 for operand in inputs {
2296 let value = tensor(*operand);
2297 if value.rank() != Some(rank) {
2298 return false;
2299 }
2300 for index in 0..rank {
2301 if index == axis {
2302 axis_size += i64::from(dimension(value, index));
2303 } else if dimension(value, index) != dimension(output, index) {
2304 return false;
2305 }
2306 }
2307 }
2308 i64::from(dimension(output, axis)) == axis_size
2309}
2310
2311fn shape_value(value: Shape<'_>, index: usize) -> Option<i64> {
2312 value.values()?.nth(index)
2313}
2314
2315fn pad_shape(
2316 input: Tensor<'_>,
2317 padding: Shape<'_>,
2318 pad: Tensor<'_>,
2319 output: Tensor<'_>,
2320 dynamic: bool,
2321) -> bool {
2322 let rank = input.rank().unwrap();
2323 if padding.rank() as usize != rank * 2 || !scalar_shape(pad) || output.rank() != Some(rank) {
2324 return false;
2325 }
2326 if padding.values().is_none() {
2327 return dynamic;
2328 }
2329 (0..rank).all(|index| {
2330 let before = shape_value(padding, index * 2);
2331 let after = shape_value(padding, index * 2 + 1);
2332 matches!((before, after), (Some(before), Some(after)) if before >= 0
2333 && after >= 0
2334 && i64::from(dimension(output, index))
2335 == i64::from(dimension(input, index)) + before + after)
2336 })
2337}
2338
2339fn tensor_elements(value: Tensor<'_>) -> Option<u64> {
2340 value.dimensions().try_fold(1_u64, |count, dimension| {
2341 count.checked_mul(dimension as u64)
2342 })
2343}
2344
2345fn reshape_shape(
2346 input: Tensor<'_>,
2347 new_shape: Shape<'_>,
2348 output: Tensor<'_>,
2349 dynamic: bool,
2350) -> bool {
2351 new_shape.rank() as usize == output.rank().unwrap()
2352 && tensor_elements(input) == tensor_elements(output)
2353 && if new_shape.values().is_some() {
2354 (0..output.rank().unwrap()).all(|index| {
2355 shape_value(new_shape, index) == Some(i64::from(dimension(output, index)))
2356 })
2357 } else {
2358 dynamic
2359 }
2360}
2361
2362fn slice_shape(
2363 input: Tensor<'_>,
2364 start: Shape<'_>,
2365 size: Shape<'_>,
2366 output: Tensor<'_>,
2367 dynamic: bool,
2368) -> bool {
2369 let rank = input.rank().unwrap();
2370 if start.rank() as usize != rank || size.rank() as usize != rank || output.rank() != Some(rank)
2371 {
2372 return false;
2373 }
2374 if start.values().is_none() || size.values().is_none() {
2375 return dynamic;
2376 }
2377 (0..rank).all(|index| {
2378 let (Some(start), Some(size)) = (shape_value(start, index), shape_value(size, index))
2379 else {
2380 return false;
2381 };
2382 start >= 0
2383 && size > 0
2384 && start + size <= i64::from(dimension(input, index))
2385 && size == i64::from(dimension(output, index))
2386 })
2387}
2388
2389fn tile_shape(input: Tensor<'_>, multiples: Shape<'_>, output: Tensor<'_>, dynamic: bool) -> bool {
2390 let rank = input.rank().unwrap();
2391 if multiples.rank() as usize != rank || output.rank() != Some(rank) {
2392 return false;
2393 }
2394 if multiples.values().is_none() {
2395 return dynamic;
2396 }
2397 (0..rank).all(|index| {
2398 shape_value(multiples, index).is_some_and(|multiple| {
2399 multiple >= 1
2400 && i64::from(dimension(input, index)) * multiple
2401 == i64::from(dimension(output, index))
2402 })
2403 })
2404}
2405
2406fn transpose_shape(input: Tensor<'_>, output: Tensor<'_>, perms: I32List<'_>) -> bool {
2407 input.rank() == output.rank()
2408 && (0..output.rank().unwrap()).all(|index| {
2409 dimension(input, perms.get(index).unwrap() as usize) == dimension(output, index)
2410 })
2411}
2412
2413fn resize_shape(
2414 input: Tensor<'_>,
2415 scale: Shape<'_>,
2416 offset: Shape<'_>,
2417 border: Shape<'_>,
2418 output: Tensor<'_>,
2419 limits: LevelLimits,
2420 dynamic: bool,
2421) -> bool {
2422 if scale.rank() != 4 || offset.rank() != 2 || border.rank() != 2 {
2423 return false;
2424 }
2425 if scale.values().is_none() || offset.values().is_none() || border.values().is_none() {
2426 return dynamic
2427 && dimension(input, 0) == dimension(output, 0)
2428 && dimension(input, 3) == dimension(output, 3)
2429 && [
2430 dimension(input, 1),
2431 dimension(input, 2),
2432 dimension(output, 1),
2433 dimension(output, 2),
2434 ]
2435 .into_iter()
2436 .all(|value| value < 16_384);
2437 }
2438 let (Some(yn_), Some(yd), Some(xn), Some(xd), Some(oy), Some(ox), Some(by), Some(bx)) = (
2439 shape_value(scale, 0),
2440 shape_value(scale, 1),
2441 shape_value(scale, 2),
2442 shape_value(scale, 3),
2443 shape_value(offset, 0),
2444 shape_value(offset, 1),
2445 shape_value(border, 0),
2446 shape_value(border, 1),
2447 ) else {
2448 return false;
2449 };
2450 if yn_ <= 0
2451 || yd <= 0
2452 || xn <= 0
2453 || xd <= 0
2454 || yn_ > 2_048
2455 || xn > 2_048
2456 || i128::from(yn_) > i128::from(limits.max_scale) * i128::from(yd)
2457 || i128::from(xn) > i128::from(limits.max_scale) * i128::from(xd)
2458 || yd >= 16 * yn_
2459 || xd >= 16 * xn
2460 || !(-yn_..16 * yn_).contains(&oy)
2461 || !(-xn..16 * xn).contains(&ox)
2462 || !(-16 * yn_..yn_).contains(&by)
2463 || !(-16 * xn..xn).contains(&bx)
2464 || [
2465 dimension(input, 1),
2466 dimension(input, 2),
2467 dimension(output, 1),
2468 dimension(output, 2),
2469 ]
2470 .into_iter()
2471 .any(|value| value >= 16_384)
2472 {
2473 return false;
2474 }
2475 let height_numerator = (i64::from(dimension(input, 1)) - 1) * yn_ - oy + by;
2476 let width_numerator = (i64::from(dimension(input, 2)) - 1) * xn - ox + bx;
2477 height_numerator % yd == 0
2478 && width_numerator % xd == 0
2479 && dimensions_equal(
2480 output,
2481 &[
2482 dimension(input, 0),
2483 i32::try_from(height_numerator / yd + 1).unwrap_or(0),
2484 i32::try_from(width_numerator / xd + 1).unwrap_or(0),
2485 dimension(input, 3),
2486 ],
2487 )
2488}
2489
2490fn control_region_exists(region_names: &[&str], name: Option<&str>) -> bool {
2491 name.is_some_and(|name| region_names.binary_search(&name).is_ok())
2492}
2493
2494fn region_block<'a>(model: &Model<'a>, name: &str) -> Option<BasicBlock<'a>> {
2495 let region = model.regions().find(|region| region.name() == name)?;
2496 let mut blocks = region.blocks();
2497 let first = blocks.next()?;
2498 if first.name() == name {
2499 Some(first)
2500 } else {
2501 blocks.find(|block| block.name() == name).or(Some(first))
2502 }
2503}
2504
2505fn block_tensor<'a>(block: BasicBlock<'a>, name: &str) -> Option<Tensor<'a>> {
2506 block.tensors().find(|value| value.name() == name)
2507}
2508
2509fn signature_matches(block: BasicBlock<'_>, operands: &[Symbol<'_>], inputs: bool) -> bool {
2510 let names = if inputs {
2511 block.inputs()
2512 } else {
2513 block.outputs()
2514 };
2515 if names.len() != operands.len() {
2516 return false;
2517 }
2518 names.zip(operands).all(|(name, operand)| {
2519 block_tensor(block, name).is_some_and(|value| {
2520 let expected = tensor(*operand);
2521 value.dtype() == expected.dtype() && same_shape(value, expected)
2522 })
2523 })
2524}
2525
2526fn region_matches(
2527 model: &Model<'_>,
2528 name: &str,
2529 inputs: &[Symbol<'_>],
2530 outputs: &[Symbol<'_>],
2531) -> bool {
2532 region_block(model, name).is_some_and(|block| {
2533 signature_matches(block, inputs, true) && signature_matches(block, outputs, false)
2534 })
2535}
2536
2537fn condition_region_matches(model: &Model<'_>, name: &str, inputs: &[Symbol<'_>]) -> bool {
2538 let Some(block) = region_block(model, name) else {
2539 return false;
2540 };
2541 if !signature_matches(block, inputs, true) || block.outputs().len() != 1 {
2542 return false;
2543 }
2544 let Some(condition) = block
2545 .outputs()
2546 .next()
2547 .and_then(|name| block_tensor(block, name))
2548 else {
2549 return false;
2550 };
2551 condition.dtype() == DType::BOOL && tensor_elements(condition) == Some(1)
2552}
2553
2554fn tensor_lists_same(inputs: &[Symbol<'_>], outputs: &[Symbol<'_>]) -> bool {
2555 inputs.len() == outputs.len()
2556 && inputs.iter().zip(outputs).all(|(input, output)| {
2557 let input = tensor(*input);
2558 let output = tensor(*output);
2559 input.dtype() == output.dtype() && same_shape(input, output)
2560 })
2561}
2562
2563#[allow(dead_code)]
2564fn valid_nan_mode(mode: NanPropagationMode) -> bool {
2565 mode == NanPropagationMode::PROPAGATE || mode == NanPropagationMode::IGNORE
2566}
2567
2568#[allow(dead_code)]
2569fn valid_resize_mode(mode: ResizeMode) -> bool {
2570 mode == ResizeMode::NEAREST || mode == ResizeMode::BILINEAR
2571}
2572
2573#[allow(dead_code)]
2574fn valid_rounding_mode(mode: RoundingMode) -> bool {
2575 mode == RoundingMode::SINGLE_ROUND
2576 || mode == RoundingMode::INEXACT_ROUND
2577 || mode == RoundingMode::DOUBLE_ROUND
2578}
2579
2580#[allow(dead_code)]
2581fn list_exact(values: I32List<'_>, length: usize) -> bool {
2582 values.len() == length
2583}
2584
2585#[cfg(test)]
2586mod tests {
2587 use super::*;
2588 use crate::{Level, Version};
2589
2590 fn complete_target() -> Target {
2591 Target::new(
2592 Version::TOSA_1_0,
2593 ProfileSet::ALL,
2594 Level::Unbounded,
2595 ExtensionSet::ALL,
2596 )
2597 }
2598
2599 #[test]
2600 fn stable_opcode_and_arity_tables_are_exhaustive() {
2601 assert_eq!(Op::ALL.len(), 75);
2602 for (index, op) in Op::ALL.iter().copied().enumerate() {
2603 assert_eq!(op.get() as usize, index + 1);
2604 assert!(op.arity().is_some(), "missing arity for {op:?}");
2605 }
2606 }
2607
2608 #[test]
2609 fn cast_matrix_matches_the_stable_specification_rows() {
2610 let target = complete_target();
2611 let dtypes = [
2612 DType::BOOL,
2613 DType::INT4,
2614 DType::INT8,
2615 DType::INT16,
2616 DType::INT32,
2617 DType::INT48,
2618 DType::FP32,
2619 DType::FP16,
2620 DType::BF16,
2621 DType::SHAPE,
2622 DType::FP8E4M3,
2623 DType::FP8E5M2,
2624 ];
2625 let rows = dtypes
2626 .into_iter()
2627 .flat_map(|input| dtypes.into_iter().map(move |output| (input, output)))
2628 .filter(|(input, output)| supports_cast(*input, *output, target))
2629 .count();
2630 assert_eq!(rows, 46);
2631 }
2632
2633 #[test]
2634 fn half_conversion_preserves_zero_order_and_nan() {
2635 assert_eq!(f16_to_f32(0), 0.0);
2636 assert_eq!(f16_to_f32(0x3c00), 1.0);
2637 assert_eq!(f16_to_f32(0xbc00), -1.0);
2638 assert!(f16_to_f32(0x7e00).is_nan());
2639 assert!(f16_to_f32(0x0001) > 0.0);
2640 }
2641}