Skip to main content

tenferro_runtime/program/
builder.rs

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
24/// Mutable validation boundary for one semantic program.
25pub 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    /// Construct an empty builder with a fresh opaque identity.
42    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    /// Attach a tensor default or large constant to one external input.
54    ///
55    /// # Errors
56    ///
57    /// Returns [`ProgramBuildError::ForeignValue`] for a token from another
58    /// builder, [`ProgramBuildError::BindingTargetNotInput`] for a computed
59    /// value, or [`ProgramBuildError::DuplicateBinding`] when the input already
60    /// has a binding.
61    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    /// Add one ordered external input.
79    ///
80    /// # Errors
81    ///
82    /// Returns [`ProgramBuildError::TooManyValues`] if the builder cannot
83    /// represent another value slot.
84    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    /// Validate that a value belongs to this builder.
94    ///
95    /// # Errors
96    ///
97    /// Returns [`ProgramBuildError::ForeignValue`] for a token from another
98    /// builder or one that does not name an existing value.
99    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    /// Borrow metadata for a builder-local value.
107    ///
108    /// # Errors
109    ///
110    /// Returns [`ProgramBuildError::ForeignValue`] for a foreign token.
111    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    /// Return the number of semantic operations added so far.
120    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    /// Import the dependency closure of ordered source roots atomically.
142    ///
143    /// Empty and duplicate roots are preserved. Tensor bindings remain
144    /// separate and are remapped only for imported source inputs.
145    ///
146    /// # Errors
147    ///
148    /// Returns [`ProgramBuildError::ForeignImportRoot`] for a root outside the
149    /// source program, [`ProgramBuildError::ForeignBindings`] for bindings
150    /// frozen with another program, [`ProgramBuildError::InvalidImport`] for
151    /// invalid source structure, or [`ProgramBuildError::TooManyValues`] when
152    /// the destination cannot represent the imported values. On error this
153    /// builder is unchanged.
154    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    /// Consume this builder and atomically freeze semantic structure and bindings.
169    ///
170    /// # Errors
171    ///
172    /// Returns [`ProgramFinishError::ForeignOutput`] for an output outside this
173    /// builder, [`ProgramFinishError::StructuralValidation`] for invalid SSA
174    /// structure, or [`ProgramFinishError::BindingFinalization`] when a tensor
175    /// binding does not match its input declaration.
176    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    /// Add one canonical core semantic operation.
228    ///
229    /// # Errors
230    ///
231    /// Returns a typed build error for foreign values, wrong arity, invalid
232    /// metadata, or an unrepresentable output count.
233    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    /// Add one extension semantic operation with explicit effects and aliases.
255    ///
256    /// # Examples
257    ///
258    /// ```
259    /// use std::any::Any;
260    /// use std::hash::Hasher;
261    /// use std::sync::Arc;
262    /// use tenferro_ops::dim_expr::DimExpr;
263    /// use tenferro_ops::ext_op::{
264    ///     ExtensionAliasDeclaration, ExtensionEffectDeclaration, ExtensionOp,
265    /// };
266    /// use tenferro_ops::{ExtensionShapeContext, SymDim};
267    /// use tenferro_runtime::program::{ProgramInputSpec, SemanticProgramBuilder};
268    /// use tenferro_tensor::DType;
269    ///
270    /// #[derive(Clone, Debug)]
271    /// struct Identity;
272    /// impl ExtensionOp for Identity {
273    ///     fn family_id(&self) -> &'static str { "example.identity.v1" }
274    ///     fn payload_hash(&self, hasher: &mut dyn Hasher) {
275    ///         hasher.write_u8(1);
276    ///     }
277    ///     fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
278    ///         other.as_any().is::<Self>()
279    ///     }
280    ///     fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
281    ///         Arc::new(self.clone())
282    ///     }
283    ///     fn as_any(&self) -> &dyn Any { self }
284    ///     fn input_count(&self) -> usize { 1 }
285    ///     fn output_count(&self) -> usize { 1 }
286    ///     fn infer_output_meta(
287    ///         &self,
288    ///         context: &mut ExtensionShapeContext<'_>,
289    ///     ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
290    ///         Ok(vec![(
291    ///             context.input_dtype(0)?,
292    ///             context.input_shape(0)?.to_vec(),
293    ///         )])
294    ///     }
295    ///     fn semantic_effects(&self) -> ExtensionEffectDeclaration<'_> {
296    ///         ExtensionEffectDeclaration::Declared(&[])
297    ///     }
298    ///     fn semantic_aliases(&self) -> ExtensionAliasDeclaration<'_> {
299    ///         ExtensionAliasDeclaration::AllFresh
300    ///     }
301    /// }
302    ///
303    /// let mut builder = SemanticProgramBuilder::new();
304    /// let input = builder.input(ProgramInputSpec::new(
305    ///     DType::F64,
306    ///     [DimExpr::Const(2)],
307    /// ))?;
308    /// let output = builder.add_extension(Arc::new(Identity), &[input])?[0];
309    /// let frozen = builder.finish(&[output])?;
310    /// assert_eq!(frozen.program.operations().count(), 1);
311    /// # Ok::<(), Box<dyn std::error::Error>>(())
312    /// ```
313    ///
314    /// # Errors
315    ///
316    /// Returns a typed build error when the payload leaves effects or aliases
317    /// undeclared, metadata inference fails, or any value/arity/alias is
318    /// invalid.
319    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    /// Return ordered input metadata after resolving bound input dimensions.
525    ///
526    /// Bound tensor shapes are process-local and intentionally live outside the
527    /// semantic identity, but AD/runtime caches that prepare shape-specialized
528    /// programs must distinguish those concrete shapes.
529    #[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    /// Return a clone of this frozen program with tensor bindings copied from
568    /// `source` onto this program's input prefix.
569    ///
570    /// This is intentionally narrow: semantic AD transforms import every source
571    /// primal input first, then append derivative seed inputs. Cached derivative
572    /// program structure can therefore be reused across source programs with the
573    /// same normalized semantics while still carrying the current source's
574    /// process-local tensor defaults.
575    #[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        // Pre-resolve concrete input shapes from tensor bindings for InputDim
723        // resolution in imported metadata.
724        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
1163/// Resolve [`DimExpr::InputDim`] references using concrete bound-input shapes.
1164///
1165/// When an input has a known tensor binding, its concrete shape replaces the
1166/// symbolic `InputDim { input_idx, axis }` reference. References to unbound
1167/// inputs are left unchanged.
1168fn 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}