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#[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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
81pub struct AnalysisOptions {
82 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#[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#[derive(Clone, Copy, Debug)]
116pub enum AnalyzedValueKind<'a> {
117 Tensor(Tensor<'a>),
118 Shape(Shape<'a>),
119}
120
121#[derive(Clone, Copy, Debug, PartialEq, Eq)]
123pub enum ConstantState {
124 NonConstant,
125 Serialized,
127 Foldable,
129}
130
131#[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 pub const fn first_use(&self) -> Option<u32> {
181 self.first_use
182 }
183
184 pub const fn last_use(&self) -> Option<u32> {
187 self.last_use
188 }
189}
190
191#[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 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#[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 pub const DEAD: Self = Self(1 << 0);
265 pub const ALIAS_INPUT_ZERO: Self = Self(1 << 1);
267 pub const FOLD_CONSTANT: Self = Self(1 << 2);
269 pub const COMPOSE_RESHAPE: Self = Self(1 << 3);
271 pub const COMPOSE_TRANSPOSE: Self = Self(1 << 4);
273 pub const CUSTOM: Self = Self(1 << 5);
275 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#[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 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#[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#[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#[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 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 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}