Skip to main content

virtio_accel_tosa/
semantic.rs

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/// Operand side used in semantic diagnostics.
11#[derive(Clone, Copy, Debug, PartialEq, Eq)]
12pub enum OperandRole {
13    Input,
14    Output,
15}
16
17/// Operator-level failure after the serialization envelope has already been validated.
18#[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/// Located semantic validation failure.
77#[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/// Complete stable-TOSA semantic pass for a declared device-neutral target.
96#[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
119/// Validate stable TOSA operator, profile, extension, level, shape, and numerical constraints.
120pub 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, &region_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                &region_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
752// Implemented below in rule-focused sections so every stable opcode is mechanically covered.
753fn 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}