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 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
26/// Mutable validation boundary for one semantic program.
27pub 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    /// Construct an empty builder with a fresh opaque identity.
44    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    /// Attach a tensor default or large constant to one external input.
56    ///
57    /// # Errors
58    ///
59    /// Returns [`ProgramBuildError::ForeignValue`] for a token from another
60    /// builder, [`ProgramBuildError::BindingTargetNotInput`] for a computed
61    /// value, or [`ProgramBuildError::DuplicateBinding`] when the input already
62    /// has a binding.
63    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    /// Add one ordered external input.
96    ///
97    /// # Errors
98    ///
99    /// Returns [`ProgramBuildError::TooManyValues`] if the builder cannot
100    /// represent another value slot.
101    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    /// Validate that a value belongs to this builder.
112    ///
113    /// # Errors
114    ///
115    /// Returns [`ProgramBuildError::ForeignValue`] for a token from another
116    /// builder or one that does not name an existing value.
117    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    /// Borrow metadata for a builder-local value.
125    ///
126    /// # Errors
127    ///
128    /// Returns [`ProgramBuildError::ForeignValue`] for a foreign token.
129    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    /// Return the number of semantic operations added so far.
138    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    /// Import the dependency closure of ordered source roots atomically.
160    ///
161    /// Empty and duplicate roots are preserved. Tensor bindings remain
162    /// separate and are remapped only for imported source inputs.
163    ///
164    /// # Errors
165    ///
166    /// Returns [`ProgramBuildError::ForeignImportRoot`] for a root outside the
167    /// source program, [`ProgramBuildError::ForeignBindings`] for bindings
168    /// frozen with another program, [`ProgramBuildError::InvalidImport`] for
169    /// invalid source structure, or [`ProgramBuildError::TooManyValues`] when
170    /// the destination cannot represent the imported values. On error this
171    /// builder is unchanged.
172    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    /// Consume this builder and atomically freeze semantic structure and bindings.
187    ///
188    /// # Errors
189    ///
190    /// Returns [`ProgramFinishError::ForeignOutput`] for an output outside this
191    /// builder, [`ProgramFinishError::StructuralValidation`] for invalid SSA
192    /// structure, or [`ProgramFinishError::BindingFinalization`] when a tensor
193    /// binding does not match its input declaration.
194    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    /// Add one canonical core semantic operation.
246    ///
247    /// # Errors
248    ///
249    /// Returns a typed build error for foreign values, wrong arity, invalid
250    /// metadata, or an unrepresentable output count.
251    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    /// Add one extension semantic operation with explicit effects and aliases.
274    ///
275    /// # Examples
276    ///
277    /// ```
278    /// use std::any::Any;
279    /// use std::hash::Hasher;
280    /// use std::sync::Arc;
281    /// use tenferro_ops::dim_expr::DimExpr;
282    /// use tenferro_ops::ext_op::{
283    ///     ExtensionAliasDeclaration, ExtensionEffectDeclaration, ExtensionOp,
284    /// };
285    /// use tenferro_ops::{ExtensionShapeContext, SymDim};
286    /// use tenferro_runtime::program::{ProgramInputSpec, SemanticProgramBuilder};
287    /// use tenferro_tensor::DType;
288    ///
289    /// #[derive(Clone, Debug)]
290    /// struct Identity;
291    /// impl ExtensionOp for Identity {
292    ///     fn family_id(&self) -> &'static str { "example.identity.v1" }
293    ///     fn payload_hash(&self, hasher: &mut dyn Hasher) {
294    ///         hasher.write_u8(1);
295    ///     }
296    ///     fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
297    ///         other.as_any().is::<Self>()
298    ///     }
299    ///     fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
300    ///         Arc::new(self.clone())
301    ///     }
302    ///     fn as_any(&self) -> &dyn Any { self }
303    ///     fn input_count(&self) -> usize { 1 }
304    ///     fn output_count(&self) -> usize { 1 }
305    ///     fn infer_output_meta(
306    ///         &self,
307    ///         context: &mut ExtensionShapeContext<'_>,
308    ///     ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
309    ///         Ok(vec![(
310    ///             context.input_dtype(0)?,
311    ///             context.input_shape(0)?.to_vec(),
312    ///         )])
313    ///     }
314    ///     fn semantic_effects(&self) -> ExtensionEffectDeclaration<'_> {
315    ///         ExtensionEffectDeclaration::Declared(&[])
316    ///     }
317    ///     fn semantic_aliases(&self) -> ExtensionAliasDeclaration<'_> {
318    ///         ExtensionAliasDeclaration::AllFresh
319    ///     }
320    /// }
321    ///
322    /// let mut builder = SemanticProgramBuilder::new();
323    /// let input = builder.input(ProgramInputSpec::new(
324    ///     DType::F64,
325    ///     [DimExpr::Const(2)],
326    /// ))?;
327    /// let output = builder.add_extension(Arc::new(Identity), &[input])?[0];
328    /// let frozen = builder.finish(&[output])?;
329    /// assert_eq!(frozen.program.operations().count(), 1);
330    /// # Ok::<(), Box<dyn std::error::Error>>(())
331    /// ```
332    ///
333    /// # Errors
334    ///
335    /// Returns a typed build error when the payload leaves effects or aliases
336    /// undeclared, metadata inference fails, or any value/arity/alias is
337    /// invalid.
338    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        // An extension that carries an externally defined scalar declares its
487        // canonical identity, which every external value of this operation carries
488        // so the program can be given a reproducible identity.
489        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
542/// Reject a core operation that names an externally defined scalar.
543///
544/// A core operation has no contribution-owned kernel for a scalar tenferro does not
545/// declare, so no core operation may name one. Carrying an external scalar through a
546/// program is what an extension operation is for.
547pub(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
567/// Reject an external scalar the program cannot give a canonical identity.
568///
569/// A semantic program's identity must be reproducible across processes, and an
570/// externally defined tag is a process-local `TypeId`. A program that carries one
571/// declares the stable name, and a value without one is rejected explicitly instead
572/// of being encoded as an unstable code.
573pub(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
588/// Keep the canonical identity a source program declared on a rebuilt value.
589///
590/// An imported value's metadata is resolved against the destination program, so it is
591/// rebuilt from a dtype and shape; the declared identity of an externally defined
592/// scalar is not derivable from those and has to be carried across.
593fn 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    /// Return ordered input metadata after resolving bound input dimensions.
621    ///
622    /// Bound tensor shapes are process-local and intentionally live outside the
623    /// semantic identity, but AD/runtime caches that prepare shape-specialized
624    /// programs must distinguish those concrete shapes.
625    #[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    /// Return a clone of this frozen program with tensor bindings copied from
664    /// `source` onto this program's input prefix.
665    ///
666    /// This is intentionally narrow: semantic AD transforms import every source
667    /// primal input first, then append derivative seed inputs. Cached derivative
668    /// program structure can therefore be reused across source programs with the
669    /// same normalized semantics while still carrying the current source's
670    /// process-local tensor defaults.
671    #[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        // Pre-resolve concrete input shapes from tensor bindings for InputDim
819        // resolution in imported metadata.
820        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
1264/// Resolve [`DimExpr::InputDim`] references using concrete bound-input shapes.
1265///
1266/// When an input has a known tensor binding, its concrete shape replaces the
1267/// symbolic `InputDim { input_idx, axis }` reference. References to unbound
1268/// inputs are left unchanged.
1269fn 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}