Skip to main content

virtio_accel_tosa/
analysis.rs

1use alloc::vec::Vec;
2use core::fmt;
3
4use crate::{
5    BasicBlock, DType, ExtensionSet, Model, Op, Operator, SemanticError, Shape, Target, Tensor,
6    validate_semantics,
7};
8
9macro_rules! index_type {
10    ($name:ident) => {
11        #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
12        #[repr(transparent)]
13        pub struct $name(u32);
14
15        impl $name {
16            pub const fn get(self) -> u32 {
17                self.0
18            }
19
20            fn from_usize(value: usize) -> Result<Self, AnalysisError> {
21                match u32::try_from(value) {
22                    Ok(value) => Ok(Self(value)),
23                    Err(_) => Err(AnalysisError::TooManyObjects),
24                }
25            }
26
27            const fn index(self) -> usize {
28                self.0 as usize
29            }
30        }
31    };
32}
33
34index_type!(RegionId);
35index_type!(BlockId);
36index_type!(ValueId);
37index_type!(OperatorId);
38
39#[cfg(test)]
40impl ValueId {
41    pub(crate) const fn from_raw(raw: u32) -> Self {
42        Self(raw)
43    }
44}
45
46/// Compact half-open span into one of [`TosaAnalysis`]'s indexed slices.
47#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
48pub struct AnalysisSpan {
49    start: u32,
50    len: u32,
51}
52
53impl AnalysisSpan {
54    pub const fn start(self) -> u32 {
55        self.start
56    }
57
58    pub const fn len(self) -> usize {
59        self.len as usize
60    }
61
62    pub const fn is_empty(self) -> bool {
63        self.len == 0
64    }
65
66    fn new(start: usize, len: usize) -> Result<Self, AnalysisError> {
67        Ok(Self {
68            start: u32::try_from(start).map_err(|_| AnalysisError::TooManyObjects)?,
69            len: u32::try_from(len).map_err(|_| AnalysisError::TooManyObjects)?,
70        })
71    }
72
73    fn range(self) -> core::ops::Range<usize> {
74        let start = self.start as usize;
75        start..start + self.len as usize
76    }
77}
78
79/// Limits for optional analysis work. They do not relax semantic or parser limits.
80#[derive(Clone, Copy, Debug, PartialEq, Eq)]
81pub struct AnalysisOptions {
82    /// Largest aggregate output that may be proposed for constant folding.
83    pub max_folded_constant_bytes: u64,
84}
85
86impl Default for AnalysisOptions {
87    fn default() -> Self {
88        Self {
89            max_folded_constant_bytes: 1024 * 1024,
90        }
91    }
92}
93
94/// Failure while producing a compact lowering plan.
95#[derive(Clone, Copy, Debug, PartialEq, Eq)]
96pub enum AnalysisError {
97    Semantic(SemanticError),
98    AllocationFailed,
99    TooManyObjects,
100}
101
102impl fmt::Display for AnalysisError {
103    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
104        write!(formatter, "{self:?}")
105    }
106}
107
108impl From<SemanticError> for AnalysisError {
109    fn from(value: SemanticError) -> Self {
110        Self::Semantic(value)
111    }
112}
113
114/// Zero-copy source value retained by an analyzed plan.
115#[derive(Clone, Copy, Debug)]
116pub enum AnalyzedValueKind<'a> {
117    Tensor(Tensor<'a>),
118    Shape(Shape<'a>),
119}
120
121/// How an analyzed value becomes constant for lowering.
122#[derive(Clone, Copy, Debug, PartialEq, Eq)]
123pub enum ConstantState {
124    NonConstant,
125    /// Bytes are already serialized in the artifact by `CONST` or `CONST_SHAPE`.
126    Serialized,
127    /// A pure, bounded constant subgraph can be evaluated once by the provider.
128    Foldable,
129}
130
131/// One globally indexed value in a model. Names and serialized payloads remain borrowed.
132#[derive(Clone, Copy, Debug)]
133pub struct AnalyzedValue<'a> {
134    id: ValueId,
135    block: BlockId,
136    name: &'a str,
137    kind: AnalyzedValueKind<'a>,
138    producer: Option<OperatorId>,
139    constant: ConstantState,
140    byte_size: Option<u64>,
141    consumers: u32,
142    first_use: Option<u32>,
143    last_use: Option<u32>,
144}
145
146impl<'a> AnalyzedValue<'a> {
147    pub const fn id(&self) -> ValueId {
148        self.id
149    }
150
151    pub const fn block(&self) -> BlockId {
152        self.block
153    }
154
155    pub const fn name(&self) -> &'a str {
156        self.name
157    }
158
159    pub const fn kind(&self) -> AnalyzedValueKind<'a> {
160        self.kind
161    }
162
163    pub const fn producer(&self) -> Option<OperatorId> {
164        self.producer
165    }
166
167    pub const fn constant(&self) -> ConstantState {
168        self.constant
169    }
170
171    pub const fn byte_size(&self) -> Option<u64> {
172        self.byte_size
173    }
174
175    pub const fn consumers(&self) -> u32 {
176        self.consumers
177    }
178
179    /// First operator execution position that reads this value within its block.
180    pub const fn first_use(&self) -> Option<u32> {
181        self.first_use
182    }
183
184    /// Last operator execution position that reads this value within its block.
185    /// Block outputs remain live through the position immediately after the final operator.
186    pub const fn last_use(&self) -> Option<u32> {
187        self.last_use
188    }
189}
190
191/// Conditions that a provider may need to lower or inspect at execution time.
192///
193/// `REQUIRE` failures make a TOSA graph unpredictable and need not be detected. Dynamic CTC
194/// inputs can affect mandatory `ERROR_IF` checks and therefore are classified separately.
195#[derive(Clone, Copy, Debug, PartialEq, Eq)]
196pub enum RuntimeCondition<'a> {
197    DynamicCompileTimeInput {
198        operator: OperatorId,
199        input_index: u16,
200        value: ValueId,
201        required_error_check: bool,
202    },
203    ShiftInRange {
204        operator: OperatorId,
205        value: ValueId,
206        maximum: u8,
207    },
208    NonZero {
209        operator: OperatorId,
210        value: ValueId,
211    },
212    Int32MultiplyInRange {
213        operator: OperatorId,
214        left: ValueId,
215        right: ValueId,
216        shift: ValueId,
217    },
218    PowDomain {
219        operator: OperatorId,
220        base: ValueId,
221        exponent: ValueId,
222    },
223    IndicesInRange {
224        operator: OperatorId,
225        indices: ValueId,
226        upper_bound: u64,
227    },
228    ScatterIndicesUnique {
229        operator: OperatorId,
230        indices: ValueId,
231    },
232    VariableState {
233        operator: OperatorId,
234        value: ValueId,
235    },
236    Custom {
237        operator: OperatorId,
238        domain: &'a str,
239        name: &'a str,
240    },
241}
242
243impl RuntimeCondition<'_> {
244    /// Whether a failing condition must be surfaced as a TOSA error for predictable inputs.
245    pub const fn error_detection_required(self) -> bool {
246        matches!(
247            self,
248            Self::DynamicCompileTimeInput {
249                required_error_check: true,
250                ..
251            }
252        )
253    }
254}
255
256/// Numerically safe and provider-conditional lowering opportunities.
257#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
258#[repr(transparent)]
259pub struct OptimizationHints(u16);
260
261impl OptimizationHints {
262    pub const NONE: Self = Self(0);
263    /// The operator and its pure dependency chain do not contribute to block outputs.
264    pub const DEAD: Self = Self(1 << 0);
265    /// Output storage may alias input zero when the provider supports internal views.
266    pub const ALIAS_INPUT_ZERO: Self = Self(1 << 1);
267    /// All inputs are constant and bounded by [`AnalysisOptions::max_folded_constant_bytes`].
268    pub const FOLD_CONSTANT: Self = Self(1 << 2);
269    /// A preceding single-use reshape can be composed without changing tensor values.
270    pub const COMPOSE_RESHAPE: Self = Self(1 << 3);
271    /// A preceding single-use transpose can be composed by combining permutations.
272    pub const COMPOSE_TRANSPOSE: Self = Self(1 << 4);
273    /// Provider-specific policy and lowering are required.
274    pub const CUSTOM: Self = Self(1 << 5);
275    /// Provider control-flow lowering is required.
276    pub const CONTROL_FLOW: Self = Self(1 << 6);
277
278    pub const fn contains(self, other: Self) -> bool {
279        self.0 & other.0 == other.0
280    }
281
282    const fn with(self, other: Self) -> Self {
283        Self(self.0 | other.0)
284    }
285}
286
287/// One operator indexed independently of serialized names.
288#[derive(Clone, Copy, Debug)]
289pub struct AnalyzedOperator<'a> {
290    id: OperatorId,
291    block: BlockId,
292    source_index: u32,
293    execution_index: u32,
294    op: Op,
295    source: Operator<'a>,
296    inputs: AnalysisSpan,
297    outputs: AnalysisSpan,
298    conditions: AnalysisSpan,
299    hints: OptimizationHints,
300}
301
302impl<'a> AnalyzedOperator<'a> {
303    pub const fn id(&self) -> OperatorId {
304        self.id
305    }
306
307    pub const fn block(&self) -> BlockId {
308        self.block
309    }
310
311    pub const fn source_index(&self) -> u32 {
312        self.source_index
313    }
314
315    pub const fn execution_index(&self) -> u32 {
316        self.execution_index
317    }
318
319    pub const fn op(&self) -> Op {
320        self.op
321    }
322
323    /// Original verified operator view for zero-copy attribute and location access.
324    pub const fn source(&self) -> Operator<'a> {
325        self.source
326    }
327
328    pub const fn hints(&self) -> OptimizationHints {
329        self.hints
330    }
331}
332
333/// One basic block and its precomputed execution order.
334#[derive(Clone, Copy, Debug)]
335pub struct AnalyzedBlock<'a> {
336    id: BlockId,
337    region: RegionId,
338    name: &'a str,
339    values: AnalysisSpan,
340    operators: AnalysisSpan,
341    inputs: AnalysisSpan,
342    outputs: AnalysisSpan,
343    execution_order: AnalysisSpan,
344}
345
346impl<'a> AnalyzedBlock<'a> {
347    pub const fn id(&self) -> BlockId {
348        self.id
349    }
350
351    pub const fn region(&self) -> RegionId {
352        self.region
353    }
354
355    pub const fn name(&self) -> &'a str {
356        self.name
357    }
358}
359
360/// One region and its contiguous block span.
361#[derive(Clone, Copy, Debug)]
362pub struct AnalyzedRegion<'a> {
363    id: RegionId,
364    name: &'a str,
365    blocks: AnalysisSpan,
366}
367
368impl<'a> AnalyzedRegion<'a> {
369    pub const fn id(&self) -> RegionId {
370        self.id
371    }
372
373    pub const fn name(&self) -> &'a str {
374        self.name
375    }
376}
377
378/// Compact provider-neutral overlay used by lowering and specialization code.
379///
380/// The plan owns only bounded indexes and bookkeeping. Names, tensor metadata, shape metadata, and
381/// constant payloads remain borrowed from the verified model.
382#[derive(Debug)]
383pub struct TosaAnalysis<'a> {
384    target: Target,
385    model_bytes: &'a [u8],
386    regions: Vec<AnalyzedRegion<'a>>,
387    blocks: Vec<AnalyzedBlock<'a>>,
388    values: Vec<AnalyzedValue<'a>>,
389    operators: Vec<AnalyzedOperator<'a>>,
390    operands: Vec<ValueId>,
391    block_io: Vec<ValueId>,
392    execution_order: Vec<OperatorId>,
393    conditions: Vec<RuntimeCondition<'a>>,
394}
395
396impl<'a> TosaAnalysis<'a> {
397    pub fn build(model: &Model<'a>, target: Target) -> Result<Self, AnalysisError> {
398        Self::build_with_options(model, target, AnalysisOptions::default())
399    }
400
401    pub fn build_with_options(
402        model: &Model<'a>,
403        target: Target,
404        options: AnalysisOptions,
405    ) -> Result<Self, AnalysisError> {
406        validate_semantics(model, target)?;
407        Builder::new(model, target, options).build()
408    }
409
410    pub const fn target(&self) -> Target {
411        self.target
412    }
413
414    /// Serialized bytes for a `CONST`/`CONST_SHAPE` value, including external constant storage.
415    pub fn serialized_constant(&self, id: ValueId) -> Option<&'a [u8]> {
416        let value = self.value(id);
417        if value.constant != ConstantState::Serialized {
418            return None;
419        }
420        match value.kind {
421            AnalyzedValueKind::Tensor(tensor) => {
422                if !tensor.data().is_empty() {
423                    return Some(tensor.data());
424                }
425                let (offset, size) = tensor.external_data_range()?;
426                let start = usize::try_from(offset).ok()?;
427                let size = usize::try_from(size).ok()?;
428                self.model_bytes.get(start..start.checked_add(size)?)
429            }
430            AnalyzedValueKind::Shape(shape) => Some(shape.data()),
431        }
432    }
433
434    pub fn regions(&self) -> &[AnalyzedRegion<'a>] {
435        &self.regions
436    }
437
438    pub fn blocks(&self) -> &[AnalyzedBlock<'a>] {
439        &self.blocks
440    }
441
442    pub fn values(&self) -> &[AnalyzedValue<'a>] {
443        &self.values
444    }
445
446    pub fn operators(&self) -> &[AnalyzedOperator<'a>] {
447        &self.operators
448    }
449
450    pub fn conditions(&self) -> &[RuntimeCondition<'a>] {
451        &self.conditions
452    }
453
454    pub fn value(&self, id: ValueId) -> &AnalyzedValue<'a> {
455        &self.values[id.index()]
456    }
457
458    pub fn operator(&self, id: OperatorId) -> &AnalyzedOperator<'a> {
459        &self.operators[id.index()]
460    }
461
462    pub fn operator_inputs(&self, operator: OperatorId) -> &[ValueId] {
463        &self.operands[self.operator(operator).inputs.range()]
464    }
465
466    pub fn operator_outputs(&self, operator: OperatorId) -> &[ValueId] {
467        &self.operands[self.operator(operator).outputs.range()]
468    }
469
470    pub fn operator_conditions(&self, operator: OperatorId) -> &[RuntimeCondition<'a>] {
471        &self.conditions[self.operator(operator).conditions.range()]
472    }
473
474    pub fn block_values(&self, block: BlockId) -> &[AnalyzedValue<'a>] {
475        &self.values[self.blocks[block.index()].values.range()]
476    }
477
478    pub fn block_operators(&self, block: BlockId) -> &[AnalyzedOperator<'a>] {
479        &self.operators[self.blocks[block.index()].operators.range()]
480    }
481
482    pub fn block_inputs(&self, block: BlockId) -> &[ValueId] {
483        &self.block_io[self.blocks[block.index()].inputs.range()]
484    }
485
486    pub fn block_outputs(&self, block: BlockId) -> &[ValueId] {
487        &self.block_io[self.blocks[block.index()].outputs.range()]
488    }
489
490    pub fn execution_order(&self, block: BlockId) -> &[OperatorId] {
491        &self.execution_order[self.blocks[block.index()].execution_order.range()]
492    }
493
494    pub fn region_blocks(&self, region: RegionId) -> &[AnalyzedBlock<'a>] {
495        &self.blocks[self.regions[region.index()].blocks.range()]
496    }
497}
498
499struct Builder<'model, 'a> {
500    model: &'model Model<'a>,
501    target: Target,
502    options: AnalysisOptions,
503    analysis: TosaAnalysis<'a>,
504}
505
506impl<'model, 'a> Builder<'model, 'a> {
507    fn new(model: &'model Model<'a>, target: Target, options: AnalysisOptions) -> Self {
508        Self {
509            model,
510            target,
511            options,
512            analysis: TosaAnalysis {
513                target,
514                model_bytes: model.as_bytes(),
515                regions: Vec::new(),
516                blocks: Vec::new(),
517                values: Vec::new(),
518                operators: Vec::new(),
519                operands: Vec::new(),
520                block_io: Vec::new(),
521                execution_order: Vec::new(),
522                conditions: Vec::new(),
523            },
524        }
525    }
526
527    fn build(mut self) -> Result<TosaAnalysis<'a>, AnalysisError> {
528        reserve(&mut self.analysis.regions, self.model.regions().len())?;
529        reserve(&mut self.analysis.blocks, self.model.stats().blocks)?;
530        reserve(
531            &mut self.analysis.values,
532            self.model.stats().tensors + self.model.stats().shapes,
533        )?;
534        reserve(&mut self.analysis.operators, self.model.stats().operators)?;
535        reserve(&mut self.analysis.operands, self.model.stats().edges)?;
536        reserve(&mut self.analysis.block_io, self.model.stats().edges)?;
537        reserve(
538            &mut self.analysis.execution_order,
539            self.model.stats().operators,
540        )?;
541
542        for region in self.model.regions() {
543            let region_id = RegionId::from_usize(self.analysis.regions.len())?;
544            let block_start = self.analysis.blocks.len();
545            for block in region.blocks() {
546                self.build_block(region_id, block)?;
547            }
548            self.analysis.regions.push(AnalyzedRegion {
549                id: region_id,
550                name: region.name(),
551                blocks: AnalysisSpan::new(block_start, self.analysis.blocks.len() - block_start)?,
552            });
553        }
554        Ok(self.analysis)
555    }
556
557    fn build_block(
558        &mut self,
559        region_id: RegionId,
560        block: BasicBlock<'a>,
561    ) -> Result<(), AnalysisError> {
562        let block_id = BlockId::from_usize(self.analysis.blocks.len())?;
563        let value_start = self.analysis.values.len();
564        let operator_start = self.analysis.operators.len();
565        let mut symbols = Vec::new();
566        reserve(&mut symbols, block.tensors().len() + block.shapes().len())?;
567
568        for tensor in block.tensors() {
569            let id = ValueId::from_usize(self.analysis.values.len())?;
570            let byte_size = tensor_byte_size(tensor);
571            self.analysis.values.push(AnalyzedValue {
572                id,
573                block: block_id,
574                name: tensor.name(),
575                kind: AnalyzedValueKind::Tensor(tensor),
576                producer: None,
577                constant: ConstantState::NonConstant,
578                byte_size,
579                consumers: 0,
580                first_use: None,
581                last_use: None,
582            });
583            symbols.push((tensor.name(), id));
584        }
585        for shape in block.shapes() {
586            let id = ValueId::from_usize(self.analysis.values.len())?;
587            self.analysis.values.push(AnalyzedValue {
588                id,
589                block: block_id,
590                name: shape.name(),
591                kind: AnalyzedValueKind::Shape(shape),
592                producer: None,
593                constant: ConstantState::NonConstant,
594                byte_size: Some(u64::from(shape.rank()) * 8),
595                consumers: 0,
596                first_use: None,
597                last_use: None,
598            });
599            symbols.push((shape.name(), id));
600        }
601        symbols.sort_unstable_by_key(|(name, _)| *name);
602
603        let input_start = self.analysis.block_io.len();
604        for name in block.inputs() {
605            self.analysis.block_io.push(resolve(&symbols, name));
606        }
607        let input_span =
608            AnalysisSpan::new(input_start, self.analysis.block_io.len() - input_start)?;
609        let output_start = self.analysis.block_io.len();
610        for name in block.outputs() {
611            self.analysis.block_io.push(resolve(&symbols, name));
612        }
613        let output_span =
614            AnalysisSpan::new(output_start, self.analysis.block_io.len() - output_start)?;
615
616        let local_operator_count = block.operators().len();
617        let mut local_inputs = Vec::new();
618        let mut local_outputs = Vec::new();
619        reserve(&mut local_inputs, local_operator_count)?;
620        reserve(&mut local_outputs, local_operator_count)?;
621        for (source_index, operator) in block.operators().enumerate() {
622            let id = OperatorId::from_usize(self.analysis.operators.len())?;
623            let inputs_start = self.analysis.operands.len();
624            for name in operator.inputs() {
625                self.analysis.operands.push(resolve(&symbols, name));
626            }
627            let inputs =
628                AnalysisSpan::new(inputs_start, self.analysis.operands.len() - inputs_start)?;
629            let outputs_start = self.analysis.operands.len();
630            for name in operator.outputs() {
631                let value = resolve(&symbols, name);
632                self.analysis.operands.push(value);
633                if !matches!(self.analysis.values[value.index()].kind, AnalyzedValueKind::Tensor(tensor) if tensor.is_variable())
634                {
635                    self.analysis.values[value.index()].producer = Some(id);
636                }
637            }
638            let outputs =
639                AnalysisSpan::new(outputs_start, self.analysis.operands.len() - outputs_start)?;
640            if matches!(operator.op(), Op::CONST | Op::CONST_SHAPE) {
641                for value in &self.analysis.operands[outputs.range()] {
642                    self.analysis.values[value.index()].constant = ConstantState::Serialized;
643                }
644            }
645            self.analysis.operators.push(AnalyzedOperator {
646                id,
647                block: block_id,
648                source_index: u32::try_from(source_index)
649                    .map_err(|_| AnalysisError::TooManyObjects)?,
650                execution_index: 0,
651                op: operator.op(),
652                source: operator,
653                inputs,
654                outputs,
655                conditions: AnalysisSpan::default(),
656                hints: OptimizationHints::NONE,
657            });
658            local_inputs.push(inputs);
659            local_outputs.push(outputs);
660        }
661
662        for (local_index, operator) in block.operators().enumerate() {
663            let operator_index = operator_start + local_index;
664            let analyzed = self.analysis.operators[operator_index];
665            let condition_start = self.analysis.conditions.len();
666            add_conditions(
667                ConditionContext {
668                    target: self.target,
669                    operator: analyzed.id,
670                    op: analyzed.op,
671                    attributes: operator.attributes(),
672                    inputs: &self.analysis.operands[analyzed.inputs.range()],
673                    outputs: &self.analysis.operands[analyzed.outputs.range()],
674                    values: &self.analysis.values,
675                },
676                &mut self.analysis.conditions,
677            )?;
678            self.analysis.operators[operator_index].conditions = AnalysisSpan::new(
679                condition_start,
680                self.analysis.conditions.len() - condition_start,
681            )?;
682        }
683
684        let order = topological_order(
685            operator_start,
686            &local_inputs,
687            &self.analysis.values,
688            &self.analysis.operands,
689            local_operator_count,
690        )?;
691        let execution_start = self.analysis.execution_order.len();
692        for (position, local_index) in order.iter().copied().enumerate() {
693            let operator_index = operator_start + local_index;
694            let operator_id = self.analysis.operators[operator_index].id;
695            self.analysis.operators[operator_index].execution_index =
696                u32::try_from(position).map_err(|_| AnalysisError::TooManyObjects)?;
697            self.analysis.execution_order.push(operator_id);
698            let input_span = self.analysis.operators[operator_index].inputs;
699            for value in &self.analysis.operands[input_span.range()] {
700                let value = &mut self.analysis.values[value.index()];
701                value.consumers = value
702                    .consumers
703                    .checked_add(1)
704                    .ok_or(AnalysisError::TooManyObjects)?;
705                let position =
706                    u32::try_from(position).map_err(|_| AnalysisError::TooManyObjects)?;
707                value.first_use.get_or_insert(position);
708                value.last_use = Some(position);
709            }
710            self.propagate_constants(operator_index)?;
711        }
712        let block_end =
713            u32::try_from(local_operator_count).map_err(|_| AnalysisError::TooManyObjects)?;
714        for value in &self.analysis.block_io[output_span.range()] {
715            self.analysis.values[value.index()].last_use = Some(block_end);
716        }
717        let mut output_values = Vec::new();
718        reserve(&mut output_values, output_span.len())?;
719        output_values.extend_from_slice(&self.analysis.block_io[output_span.range()]);
720        self.mark_live_and_hints(operator_start, &order, &output_values)?;
721
722        self.analysis.blocks.push(AnalyzedBlock {
723            id: block_id,
724            region: region_id,
725            name: block.name(),
726            values: AnalysisSpan::new(value_start, self.analysis.values.len() - value_start)?,
727            operators: AnalysisSpan::new(
728                operator_start,
729                self.analysis.operators.len() - operator_start,
730            )?,
731            inputs: input_span,
732            outputs: output_span,
733            execution_order: AnalysisSpan::new(
734                execution_start,
735                self.analysis.execution_order.len() - execution_start,
736            )?,
737        });
738        Ok(())
739    }
740
741    fn propagate_constants(&mut self, operator_index: usize) -> Result<(), AnalysisError> {
742        let operator = self.analysis.operators[operator_index];
743        if !foldable_operator(operator.op)
744            || !self.analysis.operands[operator.inputs.range()]
745                .iter()
746                .all(|value| {
747                    self.analysis.values[value.index()].constant != ConstantState::NonConstant
748                })
749        {
750            return Ok(());
751        }
752        let total_bytes = self.analysis.operands[operator.outputs.range()]
753            .iter()
754            .try_fold(0_u64, |total, value| {
755                total.checked_add(self.analysis.values[value.index()].byte_size?)
756            });
757        if total_bytes.is_some_and(|bytes| bytes <= self.options.max_folded_constant_bytes) {
758            self.analysis.operators[operator_index].hints = self.analysis.operators[operator_index]
759                .hints
760                .with(OptimizationHints::FOLD_CONSTANT);
761            for value in &self.analysis.operands[operator.outputs.range()] {
762                self.analysis.values[value.index()].constant = ConstantState::Foldable;
763            }
764        }
765        Ok(())
766    }
767
768    fn mark_live_and_hints(
769        &mut self,
770        operator_start: usize,
771        order: &[usize],
772        block_outputs: &[ValueId],
773    ) -> Result<(), AnalysisError> {
774        let mut live = Vec::new();
775        reserve(&mut live, order.len())?;
776        live.resize(order.len(), false);
777        for value in block_outputs {
778            if let Some(producer) = self.analysis.values[value.index()].producer {
779                live[producer.index() - operator_start] = true;
780            }
781        }
782        for &local_index in order {
783            let op = self.analysis.operators[operator_start + local_index].op;
784            if has_side_effects(op) {
785                live[local_index] = true;
786            }
787        }
788        for &local_index in order.iter().rev() {
789            let operator_index = operator_start + local_index;
790            if live[local_index] {
791                let inputs = self.analysis.operators[operator_index].inputs;
792                for value in &self.analysis.operands[inputs.range()] {
793                    if let Some(producer) = self.analysis.values[value.index()].producer {
794                        live[producer.index() - operator_start] = true;
795                    }
796                }
797            } else {
798                self.analysis.operators[operator_index].hints = self.analysis.operators
799                    [operator_index]
800                    .hints
801                    .with(OptimizationHints::DEAD);
802            }
803        }
804
805        for &local_index in order {
806            let operator_index = operator_start + local_index;
807            let operator = self.analysis.operators[operator_index];
808            let mut hints = operator.hints;
809            if matches!(operator.op, Op::IDENTITY | Op::RESHAPE) {
810                hints = hints.with(OptimizationHints::ALIAS_INPUT_ZERO);
811            }
812            if operator.op == Op::CUSTOM {
813                hints = hints.with(OptimizationHints::CUSTOM);
814            }
815            if matches!(operator.op, Op::COND_IF | Op::WHILE_LOOP) {
816                hints = hints.with(OptimizationHints::CONTROL_FLOW);
817            }
818            let inputs = &self.analysis.operands[operator.inputs.range()];
819            if let Some(input) = inputs.first() {
820                if self.analysis.values[input.index()].consumers == 1 {
821                    if let Some(producer) = self.analysis.values[input.index()].producer {
822                        let producer_op = self.analysis.operators[producer.index()].op;
823                        if operator.op == Op::RESHAPE && producer_op == Op::RESHAPE {
824                            hints = hints.with(OptimizationHints::COMPOSE_RESHAPE);
825                        }
826                        if operator.op == Op::TRANSPOSE && producer_op == Op::TRANSPOSE {
827                            hints = hints.with(OptimizationHints::COMPOSE_TRANSPOSE);
828                        }
829                    }
830                }
831            }
832            self.analysis.operators[operator_index].hints = hints;
833        }
834        Ok(())
835    }
836}
837
838fn reserve<T>(values: &mut Vec<T>, additional: usize) -> Result<(), AnalysisError> {
839    values
840        .try_reserve_exact(additional)
841        .map_err(|_| AnalysisError::AllocationFailed)
842}
843
844fn resolve(symbols: &[(&str, ValueId)], name: &str) -> ValueId {
845    symbols[symbols
846        .binary_search_by_key(&name, |(candidate, _)| *candidate)
847        .expect("semantically validated symbol")]
848    .1
849}
850
851fn topological_order(
852    operator_start: usize,
853    inputs: &[AnalysisSpan],
854    values: &[AnalyzedValue<'_>],
855    operands: &[ValueId],
856    operator_count: usize,
857) -> Result<Vec<usize>, AnalysisError> {
858    let mut indegrees = Vec::new();
859    reserve(&mut indegrees, operator_count)?;
860    indegrees.resize(operator_count, 0_usize);
861    let mut edges = Vec::new();
862    let edge_count = inputs
863        .iter()
864        .try_fold(0_usize, |count, span| count.checked_add(span.len()))
865        .ok_or(AnalysisError::TooManyObjects)?;
866    reserve(&mut edges, edge_count)?;
867    for (consumer, span) in inputs.iter().copied().enumerate() {
868        for value in &operands[span.range()] {
869            if let Some(producer) = values[value.index()].producer {
870                let producer = producer
871                    .index()
872                    .checked_sub(operator_start)
873                    .ok_or(AnalysisError::TooManyObjects)?;
874                edges.push((producer, consumer));
875                indegrees[consumer] = indegrees[consumer]
876                    .checked_add(1)
877                    .ok_or(AnalysisError::TooManyObjects)?;
878            }
879        }
880    }
881    edges.sort_unstable();
882    let mut order = Vec::new();
883    reserve(&mut order, operator_count)?;
884    order.extend(
885        indegrees
886            .iter()
887            .enumerate()
888            .filter_map(|(index, indegree)| (*indegree == 0).then_some(index)),
889    );
890    let mut cursor = 0;
891    while cursor < order.len() {
892        let producer = order[cursor];
893        cursor += 1;
894        let start = edges.partition_point(|(candidate, _)| *candidate < producer);
895        let end = edges.partition_point(|(candidate, _)| *candidate <= producer);
896        for &(_, consumer) in &edges[start..end] {
897            indegrees[consumer] -= 1;
898            if indegrees[consumer] == 0 {
899                order.push(consumer);
900            }
901        }
902    }
903    if order.len() == operator_count {
904        Ok(order)
905    } else {
906        // The semantic pass has already rejected cycles.
907        Err(AnalysisError::TooManyObjects)
908    }
909}
910
911fn tensor_byte_size(tensor: Tensor<'_>) -> Option<u64> {
912    let _ = tensor.rank()?;
913    let elements = tensor.dimensions().try_fold(1_u64, |count, dimension| {
914        count.checked_mul(u64::try_from(dimension).ok()?)
915    })?;
916    let bytes = match tensor.dtype() {
917        DType::INT4 => return elements.checked_add(1).map(|value| value / 2),
918        DType::BOOL | DType::INT8 | DType::FP8E4M3 | DType::FP8E5M2 => 1,
919        DType::INT16 | DType::FP16 | DType::BF16 => 2,
920        DType::INT32 | DType::FP32 => 4,
921        DType::INT48 => 6,
922        _ => return None,
923    };
924    elements.checked_mul(bytes)
925}
926
927fn foldable_operator(op: Op) -> bool {
928    matches!(
929        op,
930        Op::IDENTITY
931            | Op::RESHAPE
932            | Op::TRANSPOSE
933            | Op::REVERSE
934            | Op::SLICE
935            | Op::TILE
936            | Op::CONCAT
937    )
938}
939
940fn has_side_effects(op: Op) -> bool {
941    matches!(
942        op,
943        Op::CUSTOM
944            | Op::COND_IF
945            | Op::WHILE_LOOP
946            | Op::VARIABLE
947            | Op::VARIABLE_READ
948            | Op::VARIABLE_WRITE
949    )
950}
951
952struct ConditionContext<'a, 'plan> {
953    target: Target,
954    operator: OperatorId,
955    op: Op,
956    attributes: crate::OpAttributes<'a>,
957    inputs: &'plan [ValueId],
958    outputs: &'plan [ValueId],
959    values: &'plan [AnalyzedValue<'a>],
960}
961
962fn add_conditions<'a>(
963    context: ConditionContext<'a, '_>,
964    conditions: &mut Vec<RuntimeCondition<'a>>,
965) -> Result<(), AnalysisError> {
966    let ConditionContext {
967        target,
968        operator,
969        op,
970        attributes,
971        inputs,
972        outputs,
973        values,
974    } = context;
975    let dynamic = target.extensions.contains(ExtensionSet::DYNAMIC);
976    if dynamic {
977        for &input_index in ctc_inputs(op) {
978            let value = inputs[input_index];
979            if values[value.index()].constant != ConstantState::Serialized {
980                conditions
981                    .try_reserve(1)
982                    .map_err(|_| AnalysisError::AllocationFailed)?;
983                conditions.push(RuntimeCondition::DynamicCompileTimeInput {
984                    operator,
985                    input_index: u16::try_from(input_index)
986                        .map_err(|_| AnalysisError::TooManyObjects)?,
987                    value,
988                    required_error_check: dynamic_ctc_requires_error(op, input_index),
989                });
990            }
991        }
992    }
993    let mut push = |condition| -> Result<(), AnalysisError> {
994        conditions
995            .try_reserve(1)
996            .map_err(|_| AnalysisError::AllocationFailed)?;
997        conditions.push(condition);
998        Ok(())
999    };
1000    match op {
1001        Op::ARITHMETIC_RIGHT_SHIFT | Op::LOGICAL_LEFT_SHIFT | Op::LOGICAL_RIGHT_SHIFT => {
1002            let maximum = match values[inputs[0].index()].kind {
1003                AnalyzedValueKind::Tensor(tensor) if tensor.dtype() == DType::INT32 => 31,
1004                AnalyzedValueKind::Tensor(tensor) if tensor.dtype() == DType::INT16 => 15,
1005                _ => 7,
1006            };
1007            push(RuntimeCondition::ShiftInRange {
1008                operator,
1009                value: inputs[1],
1010                maximum,
1011            })?;
1012        }
1013        Op::INTDIV => push(RuntimeCondition::NonZero {
1014            operator,
1015            value: inputs[1],
1016        })?,
1017        Op::MUL
1018            if matches!(
1019                values[inputs[0].index()].kind,
1020                AnalyzedValueKind::Tensor(tensor) if tensor.dtype() == DType::INT32
1021            ) =>
1022        {
1023            push(RuntimeCondition::Int32MultiplyInRange {
1024                operator,
1025                left: inputs[0],
1026                right: inputs[1],
1027                shift: inputs[2],
1028            })?;
1029        }
1030        Op::POW => push(RuntimeCondition::PowDomain {
1031            operator,
1032            base: inputs[0],
1033            exponent: inputs[1],
1034        })?,
1035        Op::GATHER | Op::SCATTER => {
1036            let data = match values[inputs[0].index()].kind {
1037                AnalyzedValueKind::Tensor(tensor) => tensor,
1038                AnalyzedValueKind::Shape(_) => unreachable!(),
1039            };
1040            let upper_bound = data.dimensions().nth(1).unwrap_or(0) as u64;
1041            push(RuntimeCondition::IndicesInRange {
1042                operator,
1043                indices: inputs[1],
1044                upper_bound,
1045            })?;
1046            if op == Op::SCATTER {
1047                push(RuntimeCondition::ScatterIndicesUnique {
1048                    operator,
1049                    indices: inputs[1],
1050                })?;
1051            }
1052        }
1053        Op::VARIABLE_WRITE => push(RuntimeCondition::VariableState {
1054            operator,
1055            value: inputs[0],
1056        })?,
1057        Op::VARIABLE_READ => {
1058            if let Some(value) = outputs.first() {
1059                push(RuntimeCondition::VariableState {
1060                    operator,
1061                    value: *value,
1062                })?;
1063            }
1064        }
1065        Op::CUSTOM => {
1066            let crate::OpAttributes::Custom {
1067                operator_name,
1068                domain_name,
1069                ..
1070            } = attributes
1071            else {
1072                unreachable!()
1073            };
1074            push(RuntimeCondition::Custom {
1075                operator,
1076                domain: domain_name.unwrap_or(""),
1077                name: operator_name.unwrap_or(""),
1078            })?;
1079        }
1080        _ => {}
1081    }
1082    Ok(())
1083}
1084
1085fn ctc_inputs(op: Op) -> &'static [usize] {
1086    match op.get() {
1087        2 => &[1, 2],
1088        3..=5 | 10 => &[3, 4],
1089        7 => &[2, 3],
1090        28 => &[2],
1091        31 => &[1],
1092        41 => &[1, 2],
1093        56 => &[1, 2],
1094        57 => &[1],
1095        59 => &[1, 2],
1096        60 => &[1],
1097        64 => &[1, 2, 3],
1098        66 => &[1, 2, 3, 4],
1099        _ => &[],
1100    }
1101}
1102
1103const fn dynamic_ctc_requires_error(op: Op, input_index: usize) -> bool {
1104    match op.get() {
1105        2..=5 | 7 | 10 | 41 | 56..=57 | 59..=60 | 64 => true,
1106        66 => input_index >= 3,
1107        28 | 31 => false,
1108        _ => false,
1109    }
1110}