1use std::sync::Arc;
2
3use tenferro_ops::dim_expr::DimExpr;
4use tenferro_ops::ext_op::{
5 ExtensionAlias, ExtensionAliasDeclaration, ExtensionEffectAccess, ExtensionEffectDeclaration,
6 ExtensionOp,
7};
8use tenferro_ops::shape_extent::ShapeExtent;
9use tenferro_tensor::Tensor;
10
11use super::bindings::PendingBinding;
12use super::identity::SemanticIdentity;
13use super::metadata::SemanticProvenance;
14use super::op::{SemanticOp, SemanticOperation};
15use super::value::ProgramBuilderNonce;
16use super::{
17 Alias, BindingKey, CoreSemanticOp, Effect, EffectAccess, EffectResource, FrozenProgram,
18 ImportedProgramValues, ProgramBindingError, ProgramBindings, ProgramBuildError,
19 ProgramFinishError, ProgramImport, ProgramInputSpec, ProgramShapeRelation,
20 ProgramStructuralError, ProgramValue, ProgramValueMetadata, SemanticPlacementConstraint,
21 SemanticProgram, ShapeGuard,
22};
23
24pub struct SemanticProgramBuilder {
26 owner: ProgramBuilderNonce,
27 inputs: Vec<ProgramValue>,
28 input_specs: Vec<ProgramInputSpec>,
29 values: Vec<ProgramValueMetadata>,
30 operations: Vec<SemanticOperation>,
31 bindings: Vec<PendingBinding>,
32}
33
34impl Default for SemanticProgramBuilder {
35 fn default() -> Self {
36 Self::new()
37 }
38}
39
40impl SemanticProgramBuilder {
41 pub fn new() -> Self {
43 Self {
44 owner: ProgramBuilderNonce::fresh(),
45 inputs: Vec::new(),
46 input_specs: Vec::new(),
47 values: Vec::new(),
48 operations: Vec::new(),
49 bindings: Vec::new(),
50 }
51 }
52
53 pub fn bind_input(
62 &mut self,
63 input: ProgramValue,
64 tensor: Arc<Tensor>,
65 ) -> Result<BindingKey, ProgramBuildError> {
66 self.validate_value(input)?;
67 if !self.inputs.contains(&input) {
68 return Err(ProgramBuildError::BindingTargetNotInput);
69 }
70 if self.bindings.iter().any(|binding| binding.input == input) {
71 return Err(ProgramBuildError::DuplicateBinding);
72 }
73 let key = BindingKey::new(input.slot, self.owner);
74 self.bindings.push(PendingBinding { key, input, tensor });
75 Ok(key)
76 }
77
78 pub fn input(&mut self, spec: ProgramInputSpec) -> Result<ProgramValue, ProgramBuildError> {
85 let slot = self.next_value_slot()?;
86 let value = ProgramValue::new(slot, self.owner);
87 self.values.push(spec.metadata().clone());
88 self.inputs.push(value);
89 self.input_specs.push(spec);
90 Ok(value)
91 }
92
93 pub fn validate_value(&self, value: ProgramValue) -> Result<(), ProgramBuildError> {
100 if value.owner != self.owner || value.slot as usize >= self.values.len() {
101 return Err(ProgramBuildError::ForeignValue);
102 }
103 Ok(())
104 }
105
106 pub fn value_metadata(
112 &self,
113 value: ProgramValue,
114 ) -> Result<&ProgramValueMetadata, ProgramBuildError> {
115 self.validate_value(value)?;
116 Ok(&self.values[value.slot as usize])
117 }
118
119 pub fn operation_count(&self) -> usize {
121 self.operations.len()
122 }
123
124 pub(crate) fn add_shape_guards_to_output(
125 &mut self,
126 output: ProgramValue,
127 guards: impl IntoIterator<Item = ShapeGuard>,
128 ) -> Result<(), ProgramBuildError> {
129 self.validate_value(output)?;
130 let operation = self
131 .operations
132 .iter_mut()
133 .find(|operation| operation.outputs.contains(&output))
134 .ok_or(ProgramBuildError::GuardTargetNotOperationOutput)?;
135 let mut combined = operation.shape_guards.to_vec();
136 combined.extend(guards);
137 operation.shape_guards = combined.into_boxed_slice();
138 Ok(())
139 }
140
141 pub fn import(
155 &mut self,
156 request: ProgramImport<'_>,
157 ) -> Result<ImportedProgramValues, ProgramBuildError> {
158 let transaction = ImportTransaction::prepare(self, request)?;
159 let roots = transaction.roots.clone();
160 self.inputs.extend(transaction.inputs);
161 self.input_specs.extend(transaction.input_specs);
162 self.values.extend(transaction.values);
163 self.operations.extend(transaction.operations);
164 self.bindings.extend(transaction.bindings);
165 Ok(ImportedProgramValues::new(roots))
166 }
167
168 pub fn finish(self, outputs: &[ProgramValue]) -> Result<FrozenProgram, ProgramFinishError> {
177 if outputs
178 .iter()
179 .any(|output| output.owner != self.owner || output.slot as usize >= self.values.len())
180 {
181 return Err(ProgramFinishError::ForeignOutput);
182 }
183
184 validate_structure(
185 self.owner,
186 &self.inputs,
187 self.values.len(),
188 &self.operations,
189 )?;
190 validate_bindings(&self.inputs, &self.input_specs, &self.bindings)?;
191
192 let inputs = self.inputs.into_boxed_slice();
193 let outputs: Box<[ProgramValue]> = outputs.into();
194 let values = self.values.into_boxed_slice();
195 let operations = self.operations.into_boxed_slice();
196 let shape_guards: Box<[ShapeGuard]> = operations
197 .iter()
198 .flat_map(|operation| operation.shape_guards.iter().cloned())
199 .collect();
200 let identity =
201 SemanticIdentity::build(&inputs, &outputs, &values, &operations, &shape_guards);
202 let bindings = ProgramBindings::freeze(self.owner, self.bindings);
203 let program = SemanticProgram {
204 owner: self.owner,
205 inputs,
206 outputs,
207 values,
208 operations,
209 shape_guards,
210 identity,
211 };
212 Ok(FrozenProgram {
213 program: Arc::new(program),
214 bindings,
215 })
216 }
217
218 #[cfg(test)]
219 pub(crate) fn operation_views_for_test(
220 &self,
221 ) -> impl ExactSizeIterator<Item = super::SemanticOperationView<'_>> + '_ {
222 self.operations
223 .iter()
224 .map(super::SemanticOperationView::new)
225 }
226
227 pub fn add_op(
234 &mut self,
235 op: CoreSemanticOp,
236 inputs: &[ProgramValue],
237 ) -> Result<Box<[ProgramValue]>, ProgramBuildError> {
238 self.validate_inputs(inputs)?;
239 validate_arity(op.input_count(), inputs.len())?;
240 let output_count = op.output_count();
241 let metadata = self.infer_core_metadata(&op, inputs)?;
242 validate_output_count(output_count, metadata.len())?;
243 let aliases = (0..output_count).map(Alias::fresh).collect();
244 self.append_operation(
245 SemanticOp::Core(op),
246 inputs,
247 metadata,
248 Vec::new(),
249 aliases,
250 Vec::new(),
251 )
252 }
253
254 pub fn add_extension(
320 &mut self,
321 op: Arc<dyn ExtensionOp>,
322 inputs: &[ProgramValue],
323 ) -> Result<Box<[ProgramValue]>, ProgramBuildError> {
324 self.validate_inputs(inputs)?;
325 validate_arity(op.input_count(), inputs.len())?;
326 let effects = extension_effects(op.as_ref())?;
327 let aliases = extension_aliases(op.as_ref())?;
328 validate_aliases(&aliases, inputs.len(), op.output_count())?;
329 let (metadata, guards) = self.infer_extension_metadata(op.as_ref(), inputs)?;
330 validate_output_count(op.output_count(), metadata.len())?;
331 self.append_operation(
332 SemanticOp::Extension(op),
333 inputs,
334 metadata,
335 effects,
336 aliases,
337 guards,
338 )
339 }
340
341 fn next_value_slot(&self) -> Result<u32, ProgramBuildError> {
342 u32::try_from(self.values.len()).map_err(|_| ProgramBuildError::TooManyValues)
343 }
344
345 fn validate_inputs(&self, inputs: &[ProgramValue]) -> Result<(), ProgramBuildError> {
346 inputs
347 .iter()
348 .try_for_each(|&value| self.validate_value(value))
349 }
350
351 fn input_metadata(
352 &self,
353 inputs: &[ProgramValue],
354 ) -> Result<Vec<&ProgramValueMetadata>, ProgramBuildError> {
355 inputs
356 .iter()
357 .map(|&value| self.value_metadata(value))
358 .collect()
359 }
360
361 fn infer_core_metadata(
362 &self,
363 op: &CoreSemanticOp,
364 inputs: &[ProgramValue],
365 ) -> Result<Vec<ProgramValueMetadata>, ProgramBuildError> {
366 let input_metadata = self.input_metadata(inputs)?;
367 let precision = input_extent_precision(&input_metadata);
368 let input_dtypes: Vec<_> = input_metadata
369 .iter()
370 .map(|metadata| metadata.dtype())
371 .collect();
372 let input_shapes = inference_shapes(&input_metadata);
373 let input_shape_refs: Vec<_> = input_shapes.iter().map(Vec::as_slice).collect();
374 let standard = tenferro_ops::std_tensor_op::StdTensorOp::from(op);
375 let dtype = crate::shape_infer::infer_output_dtype(&standard, &input_dtypes)
376 .map_err(metadata_error)?;
377 if core_output_uses_local_shape_coordinates(op) {
378 let local_input_shapes: Vec<_> = input_metadata
379 .iter()
380 .enumerate()
381 .map(|(input_idx, metadata)| {
382 DimExpr::input_shape(input_idx, metadata.shape().len())
383 })
384 .collect();
385 let local_input_shape_refs: Vec<_> =
386 local_input_shapes.iter().map(Vec::as_slice).collect();
387 let output_extents =
388 crate::shape_infer::infer_output_extents(&standard, &local_input_shape_refs)
389 .map_err(metadata_error)?;
390 output_extents
391 .into_iter()
392 .map(|shape| {
393 resolve_inferred_extents(shape, precision, &input_shape_refs)
394 .map(|shape| ProgramValueMetadata::from_extents(dtype, shape))
395 })
396 .collect()
397 } else {
398 let output_extents =
399 crate::shape_infer::infer_output_extents(&standard, &input_shape_refs)
400 .map_err(metadata_error)?;
401 Ok(output_extents
402 .into_iter()
403 .map(|shape| {
404 ProgramValueMetadata::from_extents(
405 dtype,
406 conservatively_bound_extents(shape, precision),
407 )
408 })
409 .collect())
410 }
411 }
412
413 fn infer_extension_metadata(
414 &self,
415 op: &dyn ExtensionOp,
416 inputs: &[ProgramValue],
417 ) -> Result<(Vec<ProgramValueMetadata>, Vec<ShapeGuard>), ProgramBuildError> {
418 let input_metadata = self.input_metadata(inputs)?;
419 let precision = input_extent_precision(&input_metadata);
420 let input_dtypes: Vec<_> = input_metadata
421 .iter()
422 .map(|metadata| metadata.dtype())
423 .collect();
424 let input_shapes = inference_shapes(&input_metadata);
425 let input_shape_refs: Vec<_> = input_shapes.iter().map(Vec::as_slice).collect();
426 let inferred = crate::shape_infer::infer_extension_output_meta_with_constraints(
427 op,
428 &input_dtypes,
429 &input_shape_refs,
430 )
431 .map_err(metadata_error)?;
432 let metadata = inferred
433 .output_metas
434 .into_iter()
435 .map(|(dtype, shape)| {
436 ProgramValueMetadata::from_extents(
437 dtype,
438 conservatively_bound_extents(
439 shape.into_iter().map(ShapeExtent::Exact),
440 precision,
441 ),
442 )
443 })
444 .collect();
445 let guards = inferred
446 .constraints
447 .into_iter()
448 .map(|constraint| {
449 let relation = match constraint.relation {
450 tenferro_ops::ShapeRelation::Equal => ProgramShapeRelation::Equal,
451 };
452 ShapeGuard::new(relation, constraint.lhs, constraint.rhs)
453 })
454 .collect();
455 Ok((metadata, guards))
456 }
457
458 fn append_operation(
459 &mut self,
460 op: SemanticOp,
461 inputs: &[ProgramValue],
462 metadata: Vec<ProgramValueMetadata>,
463 effects: Vec<Effect>,
464 aliases: Vec<Alias>,
465 shape_guards: Vec<ShapeGuard>,
466 ) -> Result<Box<[ProgramValue]>, ProgramBuildError> {
467 let provenance = match &op {
468 SemanticOp::Core(_) => SemanticProvenance::builder(None),
469 SemanticOp::Extension(extension) => {
470 SemanticProvenance::builder(Some(extension.family_id()))
471 }
472 };
473 let start = self.values.len();
474 let end = start
475 .checked_add(metadata.len())
476 .ok_or(ProgramBuildError::TooManyValues)?;
477 if end > u32::MAX as usize {
478 return Err(ProgramBuildError::TooManyValues);
479 }
480 let outputs: Box<[_]> = (start..end)
481 .map(|slot| ProgramValue::new(slot as u32, self.owner))
482 .collect();
483 self.values.extend(metadata);
484 self.operations.push(SemanticOperation {
485 op,
486 inputs: inputs.into(),
487 outputs: outputs.clone(),
488 effects: effects.into(),
489 aliases: aliases.into(),
490 shape_guards: shape_guards.into(),
491 placement: SemanticPlacementConstraint::any(),
492 provenance,
493 });
494 Ok(outputs)
495 }
496}
497
498fn core_output_uses_local_shape_coordinates(op: &CoreSemanticOp) -> bool {
499 matches!(
500 op,
501 CoreSemanticOp::Reshape { .. }
502 | CoreSemanticOp::BroadcastInDim { .. }
503 | CoreSemanticOp::GatherDynamicSliceSizes { .. }
504 )
505}
506
507fn validate_arity(expected: usize, actual: usize) -> Result<(), ProgramBuildError> {
508 if expected == actual {
509 Ok(())
510 } else {
511 Err(ProgramBuildError::Arity { expected, actual })
512 }
513}
514
515fn validate_output_count(expected: usize, actual: usize) -> Result<(), ProgramBuildError> {
516 if expected == actual {
517 Ok(())
518 } else {
519 Err(ProgramBuildError::OutputMetadataCount { expected, actual })
520 }
521}
522
523impl FrozenProgram {
524 #[doc(hidden)]
530 pub fn input_metadata_with_bound_shapes(&self) -> Box<[ProgramValueMetadata]> {
531 let bound_input_shapes: Vec<Option<Vec<DimExpr>>> = self
532 .program
533 .inputs
534 .iter()
535 .map(|input| {
536 self.bindings.tensor_for_input(*input).map(|tensor| {
537 tensor
538 .shape()
539 .iter()
540 .map(|&size| DimExpr::Const(size))
541 .collect()
542 })
543 })
544 .collect();
545
546 self.program
547 .inputs
548 .iter()
549 .map(|input| {
550 let metadata = self.program.values[input.slot as usize].clone();
551 ProgramValueMetadata::from_extents(
552 metadata.dtype(),
553 metadata.shape().iter().map(|extent| match extent {
554 ShapeExtent::Exact(expr) => ShapeExtent::Exact(
555 resolve_dim_expr_from_input_shapes(expr, &bound_input_shapes),
556 ),
557 ShapeExtent::UpperBound(expr) => ShapeExtent::UpperBound(
558 resolve_dim_expr_from_input_shapes(expr, &bound_input_shapes),
559 ),
560 ShapeExtent::Unknown => ShapeExtent::Unknown,
561 }),
562 )
563 })
564 .collect()
565 }
566
567 #[doc(hidden)]
576 pub fn with_input_prefix_bindings_from(
577 &self,
578 source: &FrozenProgram,
579 ) -> Result<FrozenProgram, ProgramFinishError> {
580 if self.program.inputs.len() < source.program.inputs.len() {
581 return Err(ProgramFinishError::StructuralValidation {
582 source: ProgramStructuralError::InvalidValueReference,
583 });
584 }
585
586 let mut bindings = Vec::new();
587 for (source_input, destination_input) in
588 source.program.inputs.iter().zip(self.program.inputs.iter())
589 {
590 if let Some(tensor) = source.bindings.tensor_for_input(*source_input) {
591 bindings.push(PendingBinding {
592 key: BindingKey::new(destination_input.slot, self.program.owner),
593 input: *destination_input,
594 tensor,
595 });
596 }
597 }
598
599 let input_specs: Vec<_> = self
600 .program
601 .inputs
602 .iter()
603 .map(|input| {
604 ProgramInputSpec::from_metadata(self.program.values[input.slot as usize].clone())
605 })
606 .collect();
607 validate_bindings(&self.program.inputs, &input_specs, &bindings)?;
608
609 Ok(FrozenProgram {
610 program: Arc::clone(&self.program),
611 bindings: ProgramBindings::freeze(self.program.owner, bindings),
612 })
613 }
614}
615
616struct ImportTransaction {
617 inputs: Vec<ProgramValue>,
618 input_specs: Vec<ProgramInputSpec>,
619 values: Vec<ProgramValueMetadata>,
620 operations: Vec<SemanticOperation>,
621 bindings: Vec<PendingBinding>,
622 roots: Box<[ProgramValue]>,
623}
624
625impl ImportTransaction {
626 fn prepare(
627 destination: &SemanticProgramBuilder,
628 request: ProgramImport<'_>,
629 ) -> Result<Self, ProgramBuildError> {
630 let source = request.program;
631 if !request.bindings.belongs_to(source.owner) {
632 return Err(ProgramBuildError::ForeignBindings);
633 }
634 if request
635 .roots
636 .iter()
637 .any(|root| root.owner != source.owner || root.slot as usize >= source.values.len())
638 {
639 return Err(ProgramBuildError::ForeignImportRoot);
640 }
641
642 let mut producer = vec![None; source.values.len()];
643 for (operation_index, operation) in source.operations.iter().enumerate() {
644 for output in &operation.outputs {
645 producer[output.slot as usize] = Some(operation_index);
646 }
647 }
648
649 let mut needed_values = vec![false; source.values.len()];
650 let mut needed_operations = vec![false; source.operations.len()];
651 let mut pending: Vec<_> = request
652 .roots
653 .iter()
654 .map(|root| root.slot as usize)
655 .collect();
656 pending.extend(
657 request
658 .bindings
659 .bound_inputs()
660 .map(|input| input.slot as usize),
661 );
662 for (operation_index, operation) in source.operations.iter().enumerate() {
663 if !operation.effects.is_empty() {
664 needed_operations[operation_index] = true;
665 for output in &operation.outputs {
666 needed_values[output.slot as usize] = true;
667 }
668 pending.extend(operation.inputs.iter().map(|input| input.slot as usize));
669 }
670 }
671 while let Some(slot) = pending.pop() {
672 if needed_values[slot] {
673 continue;
674 }
675 needed_values[slot] = true;
676 if let Some(operation_index) = producer[slot] {
677 if !needed_operations[operation_index] {
678 needed_operations[operation_index] = true;
679 let operation = &source.operations[operation_index];
680 for output in &operation.outputs {
681 needed_values[output.slot as usize] = true;
682 }
683 pending.extend(operation.inputs.iter().map(|input| input.slot as usize));
684 }
685 }
686 }
687
688 let imported_input_count = source
689 .inputs
690 .iter()
691 .filter(|input| needed_values[input.slot as usize])
692 .count();
693 let imported_output_count: usize = source
694 .operations
695 .iter()
696 .zip(&needed_operations)
697 .filter(|(_, needed)| **needed)
698 .map(|(operation, _)| operation.outputs.len())
699 .sum();
700 let imported_value_count = imported_input_count
701 .checked_add(imported_output_count)
702 .ok_or(ProgramBuildError::TooManyValues)?;
703 let final_value_count = destination
704 .values
705 .len()
706 .checked_add(imported_value_count)
707 .ok_or(ProgramBuildError::TooManyValues)?;
708 if final_value_count > u32::MAX as usize {
709 return Err(ProgramBuildError::TooManyValues);
710 }
711
712 let mut transaction = Self {
713 inputs: Vec::with_capacity(imported_input_count),
714 input_specs: Vec::with_capacity(imported_input_count),
715 values: Vec::with_capacity(imported_value_count),
716 operations: Vec::with_capacity(
717 needed_operations.iter().filter(|needed| **needed).count(),
718 ),
719 bindings: Vec::new(),
720 roots: Box::new([]),
721 };
722 let bound_input_shapes: Vec<Option<Vec<DimExpr>>> = source
725 .inputs
726 .iter()
727 .map(|input| {
728 request.bindings.tensor_for_input(*input).map(|tensor| {
729 tensor
730 .shape()
731 .iter()
732 .map(|&size| DimExpr::Const(size))
733 .collect()
734 })
735 })
736 .collect();
737
738 let resolve_extent = |extent: &ShapeExtent<DimExpr>| -> ShapeExtent<DimExpr> {
739 match extent {
740 ShapeExtent::Exact(expr) => ShapeExtent::Exact(resolve_dim_expr_from_input_shapes(
741 expr,
742 &bound_input_shapes,
743 )),
744 ShapeExtent::UpperBound(expr) => ShapeExtent::UpperBound(
745 resolve_dim_expr_from_input_shapes(expr, &bound_input_shapes),
746 ),
747 ShapeExtent::Unknown => ShapeExtent::Unknown,
748 }
749 };
750
751 let mut remap = vec![None; source.values.len()];
752
753 for &input in &source.inputs {
754 if !needed_values[input.slot as usize] {
755 continue;
756 }
757 let metadata = source.values[input.slot as usize].clone();
758 let metadata = ProgramValueMetadata::from_extents(
759 metadata.dtype(),
760 metadata
761 .shape()
762 .iter()
763 .map(&resolve_extent)
764 .collect::<Vec<_>>(),
765 );
766 let imported = transaction.next_value(destination.values.len(), destination.owner)?;
767 transaction.inputs.push(imported);
768 transaction
769 .input_specs
770 .push(ProgramInputSpec::from_metadata(metadata.clone()));
771 transaction.values.push(metadata);
772 remap[input.slot as usize] = Some(imported);
773 if let Some(tensor) = request.bindings.tensor_for_input(input) {
774 transaction.bindings.push(PendingBinding {
775 key: BindingKey::new(imported.slot, destination.owner),
776 input: imported,
777 tensor,
778 });
779 }
780 }
781
782 for (operation, needed) in source.operations.iter().zip(needed_operations) {
783 if !needed {
784 continue;
785 }
786 let inputs: Box<[_]> = operation
787 .inputs
788 .iter()
789 .map(|input| {
790 remap[input.slot as usize].ok_or(ProgramBuildError::InvalidImport {
791 source: ProgramStructuralError::InvalidSsaOrder,
792 })
793 })
794 .collect::<Result<_, _>>()?;
795 let mut outputs = Vec::with_capacity(operation.outputs.len());
796 for output in &operation.outputs {
797 let imported =
798 transaction.next_value(destination.values.len(), destination.owner)?;
799 let meta = source.values[output.slot as usize].clone();
800 let resolved = ProgramValueMetadata::from_extents(
801 meta.dtype(),
802 meta.shape().iter().map(&resolve_extent).collect::<Vec<_>>(),
803 );
804 transaction.values.push(resolved);
805 remap[output.slot as usize] = Some(imported);
806 outputs.push(imported);
807 }
808 let op = match &operation.op {
809 SemanticOp::Core(op) => SemanticOp::Core(op.clone()),
810 SemanticOp::Extension(op) => SemanticOp::Extension(op.clone_arc()),
811 };
812 transaction.operations.push(SemanticOperation {
813 op,
814 inputs,
815 outputs: outputs.into(),
816 effects: operation.effects.clone(),
817 aliases: operation.aliases.clone(),
818 shape_guards: operation.shape_guards.clone(),
819 placement: operation.placement,
820 provenance: operation.provenance.clone(),
821 });
822 }
823
824 transaction.roots = request
825 .roots
826 .iter()
827 .map(|root| {
828 remap[root.slot as usize].ok_or(ProgramBuildError::InvalidImport {
829 source: ProgramStructuralError::InvalidValueReference,
830 })
831 })
832 .collect::<Result<_, _>>()?;
833 Ok(transaction)
834 }
835
836 fn next_value(
837 &self,
838 destination_value_count: usize,
839 owner: ProgramBuilderNonce,
840 ) -> Result<ProgramValue, ProgramBuildError> {
841 let slot = destination_value_count
842 .checked_add(self.values.len())
843 .ok_or(ProgramBuildError::TooManyValues)?;
844 let slot = u32::try_from(slot).map_err(|_| ProgramBuildError::TooManyValues)?;
845 Ok(ProgramValue::new(slot, owner))
846 }
847}
848
849#[derive(Clone, Copy)]
850enum InputExtentPrecision {
851 Exact,
852 Bounded,
853 Unknown,
854}
855
856fn input_extent_precision(metadata: &[&ProgramValueMetadata]) -> InputExtentPrecision {
857 let mut precision = InputExtentPrecision::Exact;
858 for extent in metadata.iter().flat_map(|metadata| metadata.shape()) {
859 match extent {
860 ShapeExtent::Unknown => return InputExtentPrecision::Unknown,
861 ShapeExtent::UpperBound(_) => precision = InputExtentPrecision::Bounded,
862 ShapeExtent::Exact(_) => {}
863 }
864 }
865 precision
866}
867
868fn conservatively_bound_extents(
869 extents: impl IntoIterator<Item = ShapeExtent<DimExpr>>,
870 precision: InputExtentPrecision,
871) -> impl Iterator<Item = ShapeExtent<DimExpr>> {
872 extents.into_iter().map(move |extent| match precision {
873 InputExtentPrecision::Exact => extent,
874 InputExtentPrecision::Bounded => match extent {
875 ShapeExtent::Exact(expression) | ShapeExtent::UpperBound(expression) => {
876 ShapeExtent::UpperBound(expression)
877 }
878 ShapeExtent::Unknown => ShapeExtent::Unknown,
879 },
880 InputExtentPrecision::Unknown => ShapeExtent::Unknown,
881 })
882}
883
884fn resolve_inferred_extents(
885 extents: impl IntoIterator<Item = ShapeExtent<DimExpr>>,
886 precision: InputExtentPrecision,
887 input_shapes: &[&[DimExpr]],
888) -> Result<Vec<ShapeExtent<DimExpr>>, ProgramBuildError> {
889 extents
890 .into_iter()
891 .map(|extent| {
892 if matches!(precision, InputExtentPrecision::Unknown) {
893 return Ok(ShapeExtent::Unknown);
894 }
895 let resolved = match extent {
896 ShapeExtent::Exact(expression) => ShapeExtent::Exact(
897 crate::shape_infer::resolve_dim_expr_from_shapes(&expression, input_shapes)
898 .map_err(metadata_error)?,
899 ),
900 ShapeExtent::UpperBound(expression) => ShapeExtent::UpperBound(
901 crate::shape_infer::resolve_dim_expr_from_shapes(&expression, input_shapes)
902 .map_err(metadata_error)?,
903 ),
904 ShapeExtent::Unknown => ShapeExtent::Unknown,
905 };
906 Ok(match precision {
907 InputExtentPrecision::Exact => resolved,
908 InputExtentPrecision::Bounded => match resolved {
909 ShapeExtent::Exact(expression) | ShapeExtent::UpperBound(expression) => {
910 ShapeExtent::UpperBound(expression)
911 }
912 ShapeExtent::Unknown => ShapeExtent::Unknown,
913 },
914 InputExtentPrecision::Unknown => unreachable!("handled above"),
915 })
916 })
917 .collect()
918}
919
920fn inference_shapes(metadata: &[&ProgramValueMetadata]) -> Vec<Vec<DimExpr>> {
921 metadata
922 .iter()
923 .enumerate()
924 .map(|(input_idx, metadata)| {
925 metadata
926 .shape()
927 .iter()
928 .enumerate()
929 .map(|(axis, extent)| match extent {
930 ShapeExtent::Exact(expression) | ShapeExtent::UpperBound(expression) => {
931 expression.clone()
932 }
933 ShapeExtent::Unknown => DimExpr::InputDim { input_idx, axis },
934 })
935 .collect()
936 })
937 .collect()
938}
939
940fn metadata_error(source: crate::Error) -> ProgramBuildError {
941 ProgramBuildError::Metadata {
942 source: Box::new(source),
943 }
944}
945
946fn extension_effects(op: &dyn ExtensionOp) -> Result<Vec<Effect>, ProgramBuildError> {
947 let family = op.family_id();
948 let effects = match op.semantic_effects() {
949 ExtensionEffectDeclaration::Undeclared => {
950 return Err(ProgramBuildError::UndeclaredExtensionEffects { family })
951 }
952 ExtensionEffectDeclaration::Declared(effects) => effects,
953 };
954 effects
955 .iter()
956 .map(|effect| {
957 let resource = EffectResource::new(effect.family, effect.key)
958 .map_err(|source| ProgramBuildError::InvalidEffectResource { family, source })?;
959 let access = match effect.access {
960 ExtensionEffectAccess::Read => EffectAccess::Read,
961 ExtensionEffectAccess::Write => EffectAccess::Write,
962 };
963 Ok(Effect::new(resource, access))
964 })
965 .collect()
966}
967
968fn extension_aliases(op: &dyn ExtensionOp) -> Result<Vec<Alias>, ProgramBuildError> {
969 let family = op.family_id();
970 match op.semantic_aliases() {
971 ExtensionAliasDeclaration::Undeclared => {
972 Err(ProgramBuildError::UndeclaredExtensionAliases { family })
973 }
974 ExtensionAliasDeclaration::AllFresh => {
975 Ok((0..op.output_count()).map(Alias::fresh).collect())
976 }
977 ExtensionAliasDeclaration::Declared(aliases) => aliases
978 .iter()
979 .map(|alias| match *alias {
980 ExtensionAlias::Fresh { output } => Ok(Alias::fresh(output)),
981 ExtensionAlias::ViewOf { output, input } => Ok(Alias::view_of(output, input)),
982 ExtensionAlias::MustAlias { output, input } => Ok(Alias::must_alias(output, input)),
983 ExtensionAlias::ExternalAlias {
984 output,
985 family: resource_family,
986 key,
987 } => EffectResource::new(resource_family, key)
988 .map(|resource| Alias::external(output, resource))
989 .map_err(|source| ProgramBuildError::InvalidEffectResource { family, source }),
990 })
991 .collect(),
992 }
993}
994
995fn validate_aliases(
996 aliases: &[Alias],
997 input_count: usize,
998 output_count: usize,
999) -> Result<(), ProgramBuildError> {
1000 let mut seen = vec![false; output_count];
1001 for &alias in aliases {
1002 let output = alias.output();
1003 let input = alias.input();
1004 if output >= output_count || input.is_some_and(|input| input >= input_count) {
1005 return Err(ProgramBuildError::AliasOutOfBounds {
1006 output,
1007 output_count,
1008 input,
1009 input_count,
1010 });
1011 }
1012 if seen[output] {
1013 return Err(ProgramBuildError::AliasCoverage {
1014 expected: output_count,
1015 actual: seen.iter().filter(|&&present| present).count(),
1016 });
1017 }
1018 seen[output] = true;
1019 }
1020 let actual = seen.iter().filter(|&&present| present).count();
1021 if actual != output_count {
1022 return Err(ProgramBuildError::AliasCoverage {
1023 expected: output_count,
1024 actual,
1025 });
1026 }
1027 Ok(())
1028}
1029
1030fn validate_structure(
1031 owner: ProgramBuilderNonce,
1032 inputs: &[ProgramValue],
1033 value_count: usize,
1034 operations: &[SemanticOperation],
1035) -> Result<(), ProgramFinishError> {
1036 let mut covered = vec![false; value_count];
1037 for input in inputs {
1038 if input.owner != owner
1039 || input.slot as usize >= value_count
1040 || covered[input.slot as usize]
1041 {
1042 return Err(ProgramFinishError::StructuralValidation {
1043 source: ProgramStructuralError::InvalidValueReference,
1044 });
1045 }
1046 covered[input.slot as usize] = true;
1047 }
1048 let mut previous_output = None;
1049 for operation in operations {
1050 let Some(first_output) = operation.outputs.first() else {
1051 if operation.inputs.iter().any(|value| {
1052 value.owner != owner
1053 || value.slot as usize >= value_count
1054 || !covered[value.slot as usize]
1055 }) {
1056 return Err(ProgramFinishError::StructuralValidation {
1057 source: ProgramStructuralError::InvalidValueReference,
1058 });
1059 }
1060 continue;
1061 };
1062 let output_start = first_output.slot as usize;
1063 let valid_input = operation.inputs.iter().all(|value| {
1064 value.owner == owner
1065 && (value.slot as usize) < output_start
1066 && (value.slot as usize) < value_count
1067 && covered[value.slot as usize]
1068 });
1069 let valid_output = operation.outputs.iter().enumerate().all(|(offset, value)| {
1070 value.owner == owner
1071 && value.slot as usize == output_start + offset
1072 && (value.slot as usize) < value_count
1073 && !covered[value.slot as usize]
1074 });
1075 let ordered = previous_output.is_none_or(|previous| output_start > previous);
1076 if !valid_input || !valid_output || !ordered {
1077 let source = if operation
1078 .inputs
1079 .iter()
1080 .chain(operation.outputs.iter())
1081 .any(|value| value.owner != owner || value.slot as usize >= value_count)
1082 {
1083 ProgramStructuralError::InvalidValueReference
1084 } else {
1085 ProgramStructuralError::InvalidSsaOrder
1086 };
1087 return Err(ProgramFinishError::StructuralValidation { source });
1088 }
1089 for output in &operation.outputs {
1090 covered[output.slot as usize] = true;
1091 }
1092 previous_output = operation.outputs.last().map(|value| value.slot as usize);
1093 }
1094 if covered.iter().any(|covered| !covered) {
1095 return Err(ProgramFinishError::StructuralValidation {
1096 source: ProgramStructuralError::InvalidSsaOrder,
1097 });
1098 }
1099 Ok(())
1100}
1101
1102fn validate_bindings(
1103 inputs: &[ProgramValue],
1104 input_specs: &[ProgramInputSpec],
1105 bindings: &[PendingBinding],
1106) -> Result<(), ProgramFinishError> {
1107 for binding in bindings {
1108 let input_index = inputs
1109 .iter()
1110 .position(|input| *input == binding.input)
1111 .ok_or(ProgramFinishError::BindingFinalization {
1112 source: ProgramBindingError::InvalidTarget,
1113 })?;
1114 let spec = &input_specs[input_index];
1115 let metadata = spec.metadata();
1116 let actual_dtype = binding.tensor.dtype();
1117 if actual_dtype != metadata.dtype() {
1118 return Err(ProgramFinishError::BindingFinalization {
1119 source: ProgramBindingError::DTypeMismatch {
1120 expected: metadata.dtype(),
1121 actual: actual_dtype,
1122 },
1123 });
1124 }
1125 let actual_shape = binding.tensor.shape();
1126 if actual_shape.len() != metadata.shape().len() {
1127 return Err(ProgramFinishError::BindingFinalization {
1128 source: ProgramBindingError::RankMismatch {
1129 expected: metadata.shape().len(),
1130 actual: actual_shape.len(),
1131 },
1132 });
1133 }
1134 for (axis, (extent, &actual)) in
1135 metadata.shape().iter().zip(actual_shape.iter()).enumerate()
1136 {
1137 match extent {
1138 ShapeExtent::Exact(DimExpr::Const(expected)) if *expected != actual => {
1139 return Err(ProgramFinishError::BindingFinalization {
1140 source: ProgramBindingError::ExactExtentMismatch {
1141 axis,
1142 expected: *expected,
1143 actual,
1144 },
1145 });
1146 }
1147 ShapeExtent::UpperBound(DimExpr::Const(bound)) if actual > *bound => {
1148 return Err(ProgramFinishError::BindingFinalization {
1149 source: ProgramBindingError::UpperBoundExceeded {
1150 axis,
1151 bound: *bound,
1152 actual,
1153 },
1154 });
1155 }
1156 _ => {}
1157 }
1158 }
1159 }
1160 Ok(())
1161}
1162
1163fn resolve_dim_expr_from_input_shapes(
1169 expr: &DimExpr,
1170 bound_input_shapes: &[Option<Vec<DimExpr>>],
1171) -> DimExpr {
1172 match expr {
1173 DimExpr::Const(_) => expr.clone(),
1174 DimExpr::InputDim { input_idx, axis } => {
1175 if let Some(Some(shape)) = bound_input_shapes.get(*input_idx) {
1176 if let Some(dim) = shape.get(*axis) {
1177 return dim.clone();
1178 }
1179 }
1180 expr.clone()
1181 }
1182 DimExpr::Add(a, b) => DimExpr::add(
1183 resolve_dim_expr_from_input_shapes(a, bound_input_shapes),
1184 resolve_dim_expr_from_input_shapes(b, bound_input_shapes),
1185 ),
1186 DimExpr::Sub(a, b) => DimExpr::sub(
1187 resolve_dim_expr_from_input_shapes(a, bound_input_shapes),
1188 resolve_dim_expr_from_input_shapes(b, bound_input_shapes),
1189 ),
1190 DimExpr::Mul(a, b) => DimExpr::mul(
1191 resolve_dim_expr_from_input_shapes(a, bound_input_shapes),
1192 resolve_dim_expr_from_input_shapes(b, bound_input_shapes),
1193 ),
1194 DimExpr::FloorDiv(a, b) => DimExpr::floor_div(
1195 resolve_dim_expr_from_input_shapes(a, bound_input_shapes),
1196 resolve_dim_expr_from_input_shapes(b, bound_input_shapes),
1197 ),
1198 DimExpr::Min(a, b) => DimExpr::min(
1199 resolve_dim_expr_from_input_shapes(a, bound_input_shapes),
1200 resolve_dim_expr_from_input_shapes(b, bound_input_shapes),
1201 ),
1202 DimExpr::Max(a, b) => DimExpr::max(
1203 resolve_dim_expr_from_input_shapes(a, bound_input_shapes),
1204 resolve_dim_expr_from_input_shapes(b, bound_input_shapes),
1205 ),
1206 }
1207}