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