Skip to main content

tenferro_runtime/graph/
compiler.rs

1use std::collections::HashMap;
2use std::fmt;
3use std::sync::Arc;
4
5use computegraph::compile::{compile, CompiledProgram};
6use computegraph::materialize::{
7    materialize_merge, MaterializedGraph, MaterializedOperation, MaterializedValue,
8};
9use computegraph::resolve::{resolve, ResolvedView, ValueDef};
10use computegraph::types::{OperationKey, ValueKey};
11use computegraph::GraphOperation;
12use num_complex::{Complex32, Complex64};
13use tenferro_ops::dim_expr::{DimExpr, DimExprEvalError};
14use tenferro_ops::input_key::TensorInputKey;
15use tenferro_ops::std_tensor_op::StdTensorOp;
16use tenferro_ops::{ShapeExtent, ShapeRelation, SymDim};
17#[cfg(test)]
18use tenferro_tensor::Tensor;
19use tenferro_tensor::{CacheStats, DType, SliceConfig, TensorScalar};
20
21use super::program::CompiledGraph;
22use crate::checkpoint::RetainedValue;
23#[cfg(test)]
24use crate::compiler::semantic_staging::stage_semantic_program;
25use crate::compiler::{lower_scoped_dim_expr, CompilerOptions};
26use crate::error::{Error, Result};
27use crate::extension_cache::{ExtensionCacheSelector, ExtensionCacheStore};
28use crate::metadata::registered_meta;
29use crate::program::{
30    CoreSemanticOp, FrozenProgram, ProgramInputSpec, ProgramShapeRelation, ProgramValueMetadata,
31    SemanticOpRef, SemanticProgramBuilder, ShapeGuard as ProgramShapeGuard,
32};
33use crate::shape_constraint::{discharge, LocalShapeConstraint, SlotScopedShapeConstraint};
34use crate::shape_infer::{infer_extension_output_meta, infer_output_shapes};
35use crate::trace::TracedGraph;
36use crate::traced::{try_concrete_shape, TracedTensor};
37
38#[derive(Clone)]
39struct InputDescriptor {
40    dtype: DType,
41    shape: Vec<usize>,
42    extent_identity: InputExtentIdentity,
43    default_tensor: Option<Arc<RetainedValue>>,
44    /// Canonical identity declared for an externally defined scalar input.
45    scalar_identity: Option<&'static str>,
46}
47
48#[derive(Clone, Copy)]
49enum InputExtentIdentity {
50    Concrete,
51    Symbolic,
52}
53
54impl InputDescriptor {
55    fn semantic_shape(&self, input_idx: usize) -> Vec<DimExpr> {
56        match self.extent_identity {
57            InputExtentIdentity::Concrete => DimExpr::from_concrete(&self.shape),
58            InputExtentIdentity::Symbolic => (0..self.shape.len())
59                .map(|axis| DimExpr::InputDim { input_idx, axis })
60                .collect(),
61        }
62    }
63
64    fn constraint_guard_shape(&self, input_idx: usize) -> Vec<DimExpr> {
65        if self.default_tensor.is_some() {
66            return input_dim_shape(input_idx, self.shape.len());
67        }
68        self.semantic_shape(input_idx)
69    }
70}
71
72fn input_dim_shape(input_idx: usize, rank: usize) -> Vec<DimExpr> {
73    (0..rank)
74        .map(|axis| DimExpr::InputDim { input_idx, axis })
75        .collect()
76}
77
78/// Compiler for traced tensor graphs.
79///
80/// A graph compiler lowers one or more [`TracedTensor`] outputs to a reusable
81/// [`CompiledGraph`] without requiring a backend.
82///
83/// # Examples
84///
85/// ```
86/// use tenferro_runtime::{GraphCompiler, TracedTensor};
87///
88/// let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
89/// let y = (&x + &x).unwrap();
90/// let mut compiler = GraphCompiler::new();
91/// let program = compiler.compile(&y).unwrap();
92/// assert_eq!(program.output_count(), 1);
93/// ```
94pub struct GraphCompiler {
95    extension_cache: ExtensionCacheStore,
96    compiler_options: CompilerOptions,
97}
98
99impl fmt::Debug for GraphCompiler {
100    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
101        f.debug_struct("GraphCompiler")
102            .field("extension_cache_stats", &self.cache_stats())
103            .field("compiler_options", &self.compiler_options)
104            .field("extension_cache", &self.extension_cache)
105            .finish_non_exhaustive()
106    }
107}
108
109impl GraphCompiler {
110    /// Create a compiler with bounded default caches.
111    ///
112    /// # Examples
113    ///
114    /// ```
115    /// use tenferro_runtime::GraphCompiler;
116    ///
117    /// let compiler = GraphCompiler::new();
118    /// assert!(compiler.extension_caches().is_empty());
119    /// ```
120    pub fn new() -> Self {
121        Self {
122            extension_cache: ExtensionCacheStore::new(),
123            compiler_options: CompilerOptions::default(),
124        }
125    }
126
127    /// Create a compiler with explicit lowering and optimizer options.
128    ///
129    /// # Examples
130    ///
131    /// ```
132    /// use tenferro_runtime::{CompilerOptions, OptimizerConfig};
133    /// use tenferro_runtime::GraphCompiler;
134    ///
135    /// let compiler = GraphCompiler::with_compiler_options(CompilerOptions {
136    ///     optimizer: OptimizerConfig {
137    ///         dot_decomposer: true,
138    ///         ..OptimizerConfig::default()
139    ///     },
140    /// });
141    /// assert!(compiler.compiler_options().optimizer.dot_decomposer);
142    /// ```
143    pub fn with_compiler_options(compiler_options: CompilerOptions) -> Self {
144        Self {
145            extension_cache: ExtensionCacheStore::new(),
146            compiler_options,
147        }
148    }
149
150    /// Compile one traced output into a graph program.
151    ///
152    /// # Examples
153    ///
154    /// ```
155    /// use tenferro_runtime::{GraphCompiler, TracedTensor};
156    ///
157    /// let x = TracedTensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap();
158    /// let mut compiler = GraphCompiler::new();
159    /// let y = x.neg().unwrap();
160    /// let program = compiler.compile(&y).unwrap();
161    /// assert_eq!(program.input_count(), 1);
162    /// ```
163    ///
164    /// # Errors
165    ///
166    /// Returns [`Error::Validation`] with `ShapeMismatch`, `RankMismatch`,
167    /// `DTypeMismatch`, or `InvalidArgument` for invalid graph metadata or
168    /// shape constraints, [`Error::RuntimeState`] for missing/inconsistent
169    /// metadata or cache state, and [`Error::Internal`] when the graph
170    /// violates a compiler invariant. Extension lowering failures retain
171    /// their typed [`Error::Extension`] source.
172    pub fn compile(&mut self, output: &TracedTensor) -> Result<CompiledGraph> {
173        self.compile_many(&[output])
174    }
175
176    /// Compile an immutable semantic trace without consulting a backend.
177    ///
178    /// This is the forward-only trace boundary. The compiler preserves the
179    /// frozen semantic program and bindings without preparing backend/runtime
180    /// staging. Runtime preparation owns backend-private staging and plan caches.
181    ///
182    /// # Errors
183    ///
184    /// Returns [`Error::Validation`] for invalid metadata or shape constraints,
185    /// [`Error::Extension`] when extension lowering fails,
186    /// [`Error::RuntimeState`] for inconsistent staging state, or
187    /// [`Error::Internal`] when compilation encounters an invariant violation.
188    pub fn compile_traced_graph(&mut self, graph: &TracedGraph) -> Result<CompiledGraph> {
189        self.compile_frozen(graph.frozen())
190    }
191
192    /// Compile an immutable semantic program for ordered execution.
193    ///
194    /// This entry is used by validation-preserving semantic transforms such as
195    /// whole-program AD. Tensor bindings remain outside semantic identity and
196    /// are preserved in the returned [`CompiledGraph`].
197    ///
198    /// # Examples
199    ///
200    /// ```
201    /// use tenferro_ops::dim_expr::DimExpr;
202    /// use tenferro_runtime::program::{
203    ///     CoreSemanticOp, ProgramInputSpec, SemanticProgramBuilder,
204    /// };
205    /// use tenferro_runtime::{DType, GraphCompiler};
206    ///
207    /// let mut builder = SemanticProgramBuilder::new();
208    /// let input = builder
209    ///     .input(ProgramInputSpec::new(DType::F64, [DimExpr::Const(2)]))
210    ///     .unwrap();
211    /// let output = builder.add_op(CoreSemanticOp::Neg, &[input]).unwrap()[0];
212    /// let frozen = builder.finish(&[output]).unwrap();
213    /// let compiled = GraphCompiler::new()
214    ///     .compile_frozen_program(&frozen)
215    ///     .unwrap();
216    /// assert_eq!(compiled.input_count(), 1);
217    /// ```
218    ///
219    /// # Errors
220    ///
221    /// Returns [`Error::Validation`] for invalid metadata or shape constraints,
222    /// [`Error::Extension`] when extension lowering fails,
223    /// [`Error::RuntimeState`] for inconsistent staging state, or
224    /// [`Error::Internal`] when compilation encounters an invariant violation.
225    pub fn compile_frozen_program(&mut self, frozen: &FrozenProgram) -> Result<CompiledGraph> {
226        self.compile_frozen(frozen)
227    }
228
229    fn compile_frozen(&mut self, frozen: &FrozenProgram) -> Result<CompiledGraph> {
230        validate_bound_shape_guards(frozen)?;
231        Ok(CompiledGraph::new(
232            frozen.clone(),
233            self.compiler_options,
234            [],
235        ))
236    }
237
238    /// Compile multiple traced outputs into one graph program.
239    ///
240    /// # Examples
241    ///
242    /// ```
243    /// use tenferro_runtime::{GraphCompiler, TracedTensor};
244    ///
245    /// let x = TracedTensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap();
246    /// let y = x.neg().unwrap();
247    /// let mut compiler = GraphCompiler::new();
248    /// let program = compiler.compile_many(&[&x, &y]).unwrap();
249    /// assert_eq!(program.output_count(), 2);
250    /// ```
251    ///
252    /// # Errors
253    ///
254    /// Returns [`Error::Validation`] with `ShapeMismatch`, `RankMismatch`,
255    /// `DTypeMismatch`, or `InvalidArgument` for invalid graph metadata or
256    /// shape constraints, [`Error::RuntimeState`] for missing/inconsistent
257    /// metadata or cache state, and [`Error::Internal`] when the graph
258    /// violates a compiler invariant. Extension lowering failures retain
259    /// their typed [`Error::Extension`] source.
260    pub fn compile_many(&mut self, outputs: &[&TracedTensor]) -> Result<CompiledGraph> {
261        let all_inputs = collect_default_inputs(outputs)?;
262        self.compile_many_with_descriptors(
263            outputs,
264            &HashMap::new(),
265            &all_inputs,
266            None,
267            false,
268            false,
269        )
270    }
271
272    pub(crate) fn compile_ad_source(&mut self, output: &TracedTensor) -> Result<CompiledGraph> {
273        self.compile_ad_source_many(&[output])
274    }
275
276    pub(crate) fn compile_ad_source_many(
277        &mut self,
278        outputs: &[&TracedTensor],
279    ) -> Result<CompiledGraph> {
280        let all_inputs = collect_default_inputs(outputs)?;
281        self.compile_many_with_descriptors(outputs, &HashMap::new(), &all_inputs, None, true, true)
282    }
283
284    /// Compile one traced output with concrete placeholder specs.
285    ///
286    /// The compiled program's explicit inputs follow the order of `bindings`,
287    /// and every declared placeholder must be one the output depends on. A
288    /// declared placeholder that the output does not use (for example the
289    /// variable of a derivative that is constant) is rejected at compile time:
290    /// it would not be a program input, so a tensor passed for it at run time
291    /// could not be matched to any placeholder.
292    ///
293    /// # Examples
294    ///
295    /// ```
296    /// use tenferro_runtime::{DType, GraphCompiler, TracedTensor};
297    ///
298    /// let x = TracedTensor::input_symbolic_shape(DType::F64, 1).unwrap();
299    /// let mut compiler = GraphCompiler::new();
300    /// let y = x.neg().unwrap();
301    /// let program = compiler
302    ///     .compile_with_input_specs(&y, &[(&x, DType::F64, &[3])])
303    ///     .unwrap();
304    /// assert_eq!(program.input_count(), 1);
305    ///
306    ///
307    /// // `z` does not depend on `x`, so declaring `x` is rejected.
308    /// let z = TracedTensor::from_vec_col_major(vec![3], vec![1.0_f64; 3]).unwrap();
309    /// assert!(compiler
310    ///     .compile_with_input_specs(&z, &[(&x, DType::F64, &[3])])
311    ///     .is_err());
312    /// ```
313    ///
314    /// # Errors
315    ///
316    /// Returns [`Error::UnexpectedBinding`] for a data-carrying tensor,
317    /// [`Error::DuplicateBinding`] for repeated placeholders,
318    /// [`Error::PlaceholderDtypeMismatch`],
319    /// [`Error::PlaceholderShapeMismatch`], or
320    /// [`Error::PlaceholderRankMismatch`] for incompatible specs,
321    /// [`Error::Validation`] with `InvalidArgument` (phase
322    /// [`ErrorPhase::Compile`](crate::ErrorPhase::Compile)) for a declared placeholder the output does not
323    /// depend on, and [`Error::Validation`] with `ShapeMismatch`,
324    /// `RankMismatch`, `DTypeMismatch`, or `InvalidArgument` /
325    /// [`Error::RuntimeState`] when compilation or metadata lowering fails.
326    pub fn compile_with_input_specs(
327        &mut self,
328        output: &TracedTensor,
329        bindings: &[(&TracedTensor, DType, &[usize])],
330    ) -> Result<CompiledGraph> {
331        let mut binding_specs = HashMap::new();
332        let mut input_order = Vec::with_capacity(bindings.len());
333        for (index, (placeholder, dtype, shape)) in bindings.iter().enumerate() {
334            validate_placeholder_spec(index, placeholder, *dtype, shape)?;
335            let key = placeholder.input_key().ok_or(Error::UnexpectedBinding {
336                binding_index: index,
337            })?;
338            if binding_specs
339                .insert(
340                    key.clone(),
341                    InputDescriptor {
342                        dtype: *dtype,
343                        shape: (*shape).to_vec(),
344                        extent_identity: InputExtentIdentity::Concrete,
345                        default_tensor: None,
346                        // A caller-supplied binding declares its own identity through
347                        // the traced value's metadata instead.
348                        scalar_identity: None,
349                    },
350                )
351                .is_some()
352            {
353                return Err(Error::DuplicateBinding {
354                    input_key: format!("{:?}", key),
355                });
356            }
357            input_order.push(key);
358        }
359
360        let program = self.compile_many_with_descriptors(
361            &[output],
362            &binding_specs,
363            output.inputs_map.as_ref(),
364            Some(&input_order),
365            false,
366            false,
367        )?;
368        if let Some(binding_index) = input_order
369            .iter()
370            .position(|key| program.input_key_index(key).is_none())
371        {
372            return Err(Error::invalid_argument(
373                "GraphCompiler::compile_with_input_specs",
374                crate::ErrorPhase::Compile,
375                "bindings",
376                format!(
377                    "binding {binding_index} declares a placeholder the output does not \
378                     depend on; remove it from the bindings"
379                ),
380            ));
381        }
382        Ok(program)
383    }
384
385    /// Return the compiler options used for future graph lowerings.
386    ///
387    /// # Examples
388    ///
389    /// ```
390    /// use tenferro_runtime::CompilerOptions;
391    /// use tenferro_runtime::GraphCompiler;
392    ///
393    /// let compiler = GraphCompiler::new();
394    /// assert_eq!(compiler.compiler_options(), CompilerOptions::default());
395    /// ```
396    pub fn compiler_options(&self) -> CompilerOptions {
397        self.compiler_options
398    }
399
400    /// Replace compiler options and clear compiler-owned extension cache entries.
401    ///
402    /// # Examples
403    ///
404    /// ```
405    /// use tenferro_runtime::{CompilerOptions, OptimizerConfig};
406    /// use tenferro_runtime::GraphCompiler;
407    ///
408    /// let mut compiler = GraphCompiler::new();
409    /// let options = CompilerOptions {
410    ///     optimizer: OptimizerConfig {
411    ///         dot_decomposer: true,
412    ///         ..OptimizerConfig::default()
413    ///     },
414    /// };
415    /// compiler.set_compiler_options(options);
416    /// assert_eq!(compiler.compiler_options(), options);
417    /// assert_eq!(compiler.cache_stats().entries, 0);
418    /// ```
419    pub fn set_compiler_options(&mut self, compiler_options: CompilerOptions) {
420        if self.compiler_options == compiler_options {
421            return;
422        }
423        self.compiler_options = compiler_options;
424        self.clear_extension_caches();
425    }
426
427    /// Clear generic extension compile-time cache entries.
428    ///
429    /// # Examples
430    ///
431    /// ```
432    /// use tenferro_runtime::GraphCompiler;
433    ///
434    /// let mut compiler = GraphCompiler::new();
435    /// compiler.clear_extension_caches();
436    /// assert_eq!(compiler.cache_stats().entries, 0);
437    /// ```
438    pub fn clear_extension_caches(&mut self) {
439        self.extension_cache.clear();
440    }
441
442    /// Clear every cache owned by the compiler.
443    ///
444    /// # Examples
445    ///
446    /// ```
447    /// use tenferro_runtime::GraphCompiler;
448    ///
449    /// let mut compiler = GraphCompiler::new();
450    /// compiler.clear_caches();
451    /// assert_eq!(compiler.cache_stats().entries, 0);
452    /// ```
453    pub fn clear_caches(&mut self) {
454        self.clear_extension_caches();
455    }
456
457    /// Return compiler-owned extension cache-entry and retained-byte stats.
458    ///
459    /// # Examples
460    ///
461    /// ```
462    /// use tenferro_runtime::GraphCompiler;
463    ///
464    /// let compiler = GraphCompiler::new();
465    /// let stats = compiler.cache_stats();
466    /// assert_eq!(stats.entries, 0);
467    /// ```
468    pub fn cache_stats(&self) -> CacheStats {
469        self.extension_cache.stats(ExtensionCacheSelector::All)
470    }
471
472    /// Borrow generic compiler-owned extension cache storage.
473    ///
474    /// # Examples
475    ///
476    /// ```
477    /// use tenferro_runtime::GraphCompiler;
478    ///
479    /// let compiler = GraphCompiler::new();
480    /// assert!(compiler.extension_caches().is_empty());
481    /// ```
482    pub fn extension_caches(&self) -> &ExtensionCacheStore {
483        &self.extension_cache
484    }
485
486    /// Mutably borrow generic compiler-owned extension cache storage.
487    ///
488    /// # Examples
489    ///
490    /// ```
491    /// use tenferro_runtime::GraphCompiler;
492    ///
493    /// let mut compiler = GraphCompiler::new();
494    /// compiler.extension_caches_mut().clear();
495    /// ```
496    pub fn extension_caches_mut(&mut self) -> &mut ExtensionCacheStore {
497        &mut self.extension_cache
498    }
499
500    fn compile_many_with_descriptors(
501        &mut self,
502        outputs: &[&TracedTensor],
503        binding_specs: &HashMap<TensorInputKey, InputDescriptor>,
504        default_inputs: &HashMap<TensorInputKey, Arc<RetainedValue>>,
505        explicit_input_order: Option<&[TensorInputKey]>,
506        include_checkpoint_aliases: bool,
507        allow_unbound_placeholders: bool,
508    ) -> Result<CompiledGraph> {
509        let mut constraint_scopes = Vec::new();
510        let mut seen_constraint_scopes = std::collections::HashSet::new();
511        for output in outputs {
512            for scope in output.constraint_scopes.as_slice() {
513                if seen_constraint_scopes.insert(Arc::as_ptr(scope)) {
514                    #[cfg(test)]
515                    test_support::record_constraint_scope_clones(1);
516                    constraint_scopes.push(Arc::clone(scope));
517                }
518            }
519        }
520
521        let mut roots = Vec::new();
522        let mut checkpoint_aliases = HashMap::new();
523        let mut output_keys = Vec::with_capacity(outputs.len());
524        for output in outputs {
525            roots.extend(output.resolve_roots());
526            if include_checkpoint_aliases && let Some(chain) = &output.checkpoint_chain {
527                roots.extend(chain.collect_graphs());
528                for (alias_key, target_key) in chain.collect_aliases() {
529                    insert_checkpoint_alias(&mut checkpoint_aliases, alias_key, target_key)?;
530                }
531            }
532            output_keys.push(output.graph.values()[output.val].key.clone());
533        }
534
535        let view = resolve(roots);
536        let graph = if checkpoint_aliases.is_empty() {
537            materialize_merge(&view, &output_keys)
538        } else {
539            let checkpoint_alias_shapes = checkpoint_aliases
540                .keys()
541                .filter_map(|key| {
542                    default_inputs
543                        .get(key)
544                        .map(|tensor| (key.clone(), tensor.shape().to_vec()))
545                })
546                .collect::<HashMap<_, _>>();
547            materialize_merge_with_input_aliases(
548                &view,
549                &output_keys,
550                &checkpoint_aliases,
551                &checkpoint_alias_shapes,
552            )
553        };
554        let mut compiled = compile(&graph);
555        prune_compiled_extension_outputs(&mut compiled)?;
556        let slot_by_key: HashMap<_, _> = graph
557            .values
558            .iter()
559            .enumerate()
560            .map(|(slot, value)| (value.key.clone(), slot))
561            .collect();
562        let mut scoped_constraints = Vec::new();
563        for scope in constraint_scopes {
564            for scoped in scope.constraints() {
565                let origin_slots: Vec<_> = scoped
566                    .origins
567                    .iter()
568                    .filter_map(|key| slot_by_key.get(key).copied())
569                    .collect();
570                if origin_slots.is_empty() {
571                    continue;
572                }
573                let origin_instruction = origin_slots.iter().find_map(|&slot| {
574                    graph
575                        .values
576                        .get(slot)
577                        .and_then(|value| value.producer.map(|p| p.0))
578                });
579                let mut local = scoped.local.clone();
580                if let Some(instruction_index) = origin_instruction {
581                    local.source = local.source.with_instruction(instruction_index);
582                }
583                let mut input_slots = Vec::with_capacity(scoped.inputs.len());
584                for (input_idx, key) in scoped.inputs.iter().enumerate() {
585                    let Some(slot) = slot_by_key.get(key).copied() else {
586                        return Err(Error::ShapeConstraintEvaluation {
587                            family: local.source.family_id,
588                            instruction_index: local.source.instruction_index,
589                            relation: local.relation,
590                            expression: format!("{:?}", local.lhs),
591                            cause: crate::ShapeConstraintEvalError::MissingInput {
592                                input_idx,
593                                input_count: scoped.inputs.len(),
594                            },
595                        });
596                    };
597                    input_slots.push(slot);
598                }
599                scoped_constraints.push(SlotScopedShapeConstraint {
600                    origin_slots,
601                    input_slots,
602                    local,
603                });
604            }
605        }
606
607        let mut descriptors = Vec::with_capacity(graph.inputs.len());
608        let mut input_keys = Vec::with_capacity(graph.inputs.len());
609        for key in &graph.inputs {
610            let ValueKey::Input(input_key) = key else {
611                return Err(Error::Internal(
612                    "expected Input key in graph inputs".to_string(),
613                ));
614            };
615            let descriptor = descriptor_for_input(
616                input_key,
617                binding_specs,
618                default_inputs,
619                allow_unbound_placeholders,
620            )?;
621            descriptors.push(descriptor);
622            input_keys.push(input_key.clone());
623        }
624        if let Some(explicit_input_order) = explicit_input_order {
625            let input_position_by_key: HashMap<_, _> = graph
626                .inputs
627                .iter()
628                .enumerate()
629                .filter_map(|(position, key)| match key {
630                    ValueKey::Input(key) => Some((key.clone(), position)),
631                    _ => None,
632                })
633                .collect();
634            let mut ordered_positions = Vec::with_capacity(graph.inputs.len());
635            let mut selected_positions = vec![false; graph.inputs.len()];
636            for key in explicit_input_order {
637                if let Some(&position) = input_position_by_key.get(key) {
638                    ordered_positions.push(position);
639                    selected_positions[position] = true;
640                }
641            }
642            for (position, selected) in selected_positions.iter().enumerate() {
643                if !selected {
644                    ordered_positions.push(position);
645                }
646            }
647            compiled.input_slots = ordered_positions
648                .iter()
649                .map(|&position| compiled.input_slots[position])
650                .collect();
651            descriptors = ordered_positions
652                .iter()
653                .map(|&position| descriptors[position].clone())
654                .collect();
655            input_keys = ordered_positions
656                .into_iter()
657                .map(|position| input_keys[position].clone())
658                .collect();
659        }
660
661        let semantic =
662            compile_materialized_semantic_program(&compiled, &descriptors, &scoped_constraints)?;
663        validate_bound_shape_guards(&semantic)?;
664        Ok(CompiledGraph::new(
665            semantic,
666            self.compiler_options,
667            input_keys,
668        ))
669    }
670}
671
672fn validate_bound_shape_guards(frozen: &FrozenProgram) -> Result<()> {
673    let input_shapes = compile_time_input_shapes(frozen)?;
674    for operation in frozen.program.operations() {
675        let fallback_family = match operation.op() {
676            SemanticOpRef::Core(_) => "tenferro-runtime.core.v1",
677            SemanticOpRef::Extension(extension) => extension.family_id(),
678        };
679        for guard in operation.shape_guards() {
680            if guard.source_family().is_some() {
681                continue;
682            }
683            validate_bound_shape_guard(guard, fallback_family, &input_shapes)?;
684        }
685    }
686    Ok(())
687}
688
689fn compile_time_input_shapes(frozen: &FrozenProgram) -> Result<Vec<Option<Vec<usize>>>> {
690    let metadata = frozen.input_metadata_with_bound_shapes();
691    frozen
692        .program
693        .inputs()
694        .iter()
695        .enumerate()
696        .map(|(input_idx, &input)| {
697            if let Some(tensor) = frozen.bindings.tensor_ref_for_input(input) {
698                return Ok(Some(tensor.shape().to_vec()));
699            }
700            let Some(metadata) = metadata.get(input_idx) else {
701                return Err(invalid_compiled_graph(format!(
702                    "semantic input metadata index {input_idx} is outside metadata table"
703                )));
704            };
705            concrete_shape_from_input_metadata(metadata)
706        })
707        .collect()
708}
709
710fn concrete_shape_from_input_metadata(
711    metadata: &ProgramValueMetadata,
712) -> Result<Option<Vec<usize>>> {
713    let mut shape = Vec::with_capacity(metadata.shape().len());
714    for extent in metadata.shape() {
715        let ShapeExtent::Exact(expression) = extent else {
716            return Ok(None);
717        };
718        let Some(value) = evaluate_static_dim_expr(expression)? else {
719            return Ok(None);
720        };
721        shape.push(value);
722    }
723    Ok(Some(shape))
724}
725
726fn validate_bound_shape_guard(
727    guard: &ProgramShapeGuard,
728    fallback_family: &'static str,
729    input_shapes: &[Option<Vec<usize>>],
730) -> Result<()> {
731    let ProgramShapeRelation::Equal = guard.relation() else {
732        return Ok(());
733    };
734    let family = guard.source_family().unwrap_or(fallback_family);
735    let relation = ShapeRelation::Equal;
736    let lhs = evaluate_bound_shape_guard_expression(family, relation, guard.lhs(), input_shapes)?;
737    let rhs = evaluate_bound_shape_guard_expression(family, relation, guard.rhs(), input_shapes)?;
738    let (Some(lhs), Some(rhs)) = (lhs, rhs) else {
739        return Ok(());
740    };
741    if lhs == rhs {
742        return Ok(());
743    }
744    Err(Error::ShapeConstraintViolation {
745        family,
746        instruction_index: None,
747        relation,
748        lhs_expr: format!("{:?}", guard.lhs()),
749        rhs_expr: format!("{:?}", guard.rhs()),
750        lhs_value: lhs,
751        rhs_value: rhs,
752    })
753}
754
755fn evaluate_bound_shape_guard_expression(
756    family: &'static str,
757    relation: ShapeRelation,
758    expression: &DimExpr,
759    input_shapes: &[Option<Vec<usize>>],
760) -> Result<Option<usize>> {
761    evaluate_static_dim_expr_with_inputs(expression, input_shapes).map_err(|cause| {
762        Error::ShapeConstraintEvaluation {
763            family,
764            instruction_index: None,
765            relation,
766            expression: format!("{expression:?}"),
767            cause: cause.into(),
768        }
769    })
770}
771
772fn evaluate_static_dim_expr(expression: &DimExpr) -> Result<Option<usize>> {
773    evaluate_static_dim_expr_without_inputs(expression).map_err(|cause| {
774        Error::ShapeConstraintEvaluation {
775            family: "tenferro-runtime.input.v1",
776            instruction_index: None,
777            relation: ShapeRelation::Equal,
778            expression: format!("{expression:?}"),
779            cause: cause.into(),
780        }
781    })
782}
783
784fn evaluate_static_dim_expr_without_inputs(
785    expression: &DimExpr,
786) -> std::result::Result<Option<usize>, DimExprEvalError> {
787    match expression {
788        DimExpr::Const(value) => Ok(Some(*value)),
789        DimExpr::InputDim { .. } => Ok(None),
790        DimExpr::Add(a, b) => {
791            let Some(lhs) = evaluate_static_dim_expr_without_inputs(a)? else {
792                return Ok(None);
793            };
794            let Some(rhs) = evaluate_static_dim_expr_without_inputs(b)? else {
795                return Ok(None);
796            };
797            lhs.checked_add(rhs)
798                .map(Some)
799                .ok_or(DimExprEvalError::AddOverflow { lhs, rhs })
800        }
801        DimExpr::Sub(a, b) => {
802            let Some(lhs) = evaluate_static_dim_expr_without_inputs(a)? else {
803                return Ok(None);
804            };
805            let Some(rhs) = evaluate_static_dim_expr_without_inputs(b)? else {
806                return Ok(None);
807            };
808            lhs.checked_sub(rhs)
809                .map(Some)
810                .ok_or(DimExprEvalError::SubUnderflow { lhs, rhs })
811        }
812        DimExpr::Mul(a, b) => {
813            let Some(lhs) = evaluate_static_dim_expr_without_inputs(a)? else {
814                return Ok(None);
815            };
816            let Some(rhs) = evaluate_static_dim_expr_without_inputs(b)? else {
817                return Ok(None);
818            };
819            lhs.checked_mul(rhs)
820                .map(Some)
821                .ok_or(DimExprEvalError::MulOverflow { lhs, rhs })
822        }
823        DimExpr::FloorDiv(a, b) => {
824            let Some(lhs) = evaluate_static_dim_expr_without_inputs(a)? else {
825                return Ok(None);
826            };
827            let Some(rhs) = evaluate_static_dim_expr_without_inputs(b)? else {
828                return Ok(None);
829            };
830            if rhs == 0 {
831                return Err(DimExprEvalError::FloorDivByZero { lhs, rhs });
832            }
833            Ok(Some(lhs / rhs))
834        }
835        DimExpr::Min(a, b) => {
836            let Some(lhs) = evaluate_static_dim_expr_without_inputs(a)? else {
837                return Ok(None);
838            };
839            let Some(rhs) = evaluate_static_dim_expr_without_inputs(b)? else {
840                return Ok(None);
841            };
842            Ok(Some(lhs.min(rhs)))
843        }
844        DimExpr::Max(a, b) => {
845            let Some(lhs) = evaluate_static_dim_expr_without_inputs(a)? else {
846                return Ok(None);
847            };
848            let Some(rhs) = evaluate_static_dim_expr_without_inputs(b)? else {
849                return Ok(None);
850            };
851            Ok(Some(lhs.max(rhs)))
852        }
853    }
854}
855
856fn evaluate_static_dim_expr_with_inputs(
857    expression: &DimExpr,
858    input_shapes: &[Option<Vec<usize>>],
859) -> std::result::Result<Option<usize>, DimExprEvalError> {
860    match expression {
861        DimExpr::Const(value) => Ok(Some(*value)),
862        DimExpr::InputDim { input_idx, axis } => match input_shapes.get(*input_idx) {
863            Some(Some(shape)) => {
864                shape
865                    .get(*axis)
866                    .copied()
867                    .map(Some)
868                    .ok_or(DimExprEvalError::AxisOutOfBounds {
869                        input_idx: *input_idx,
870                        axis: *axis,
871                        rank: shape.len(),
872                    })
873            }
874            Some(None) => Ok(None),
875            None => Err(DimExprEvalError::InputOutOfBounds {
876                input_idx: *input_idx,
877                input_count: input_shapes.len(),
878            }),
879        },
880        DimExpr::Add(a, b) => {
881            let Some(lhs) = evaluate_static_dim_expr_with_inputs(a, input_shapes)? else {
882                return Ok(None);
883            };
884            let Some(rhs) = evaluate_static_dim_expr_with_inputs(b, input_shapes)? else {
885                return Ok(None);
886            };
887            lhs.checked_add(rhs)
888                .map(Some)
889                .ok_or(DimExprEvalError::AddOverflow { lhs, rhs })
890        }
891        DimExpr::Sub(a, b) => {
892            let Some(lhs) = evaluate_static_dim_expr_with_inputs(a, input_shapes)? else {
893                return Ok(None);
894            };
895            let Some(rhs) = evaluate_static_dim_expr_with_inputs(b, input_shapes)? else {
896                return Ok(None);
897            };
898            lhs.checked_sub(rhs)
899                .map(Some)
900                .ok_or(DimExprEvalError::SubUnderflow { lhs, rhs })
901        }
902        DimExpr::Mul(a, b) => {
903            let Some(lhs) = evaluate_static_dim_expr_with_inputs(a, input_shapes)? else {
904                return Ok(None);
905            };
906            let Some(rhs) = evaluate_static_dim_expr_with_inputs(b, input_shapes)? else {
907                return Ok(None);
908            };
909            lhs.checked_mul(rhs)
910                .map(Some)
911                .ok_or(DimExprEvalError::MulOverflow { lhs, rhs })
912        }
913        DimExpr::FloorDiv(a, b) => {
914            let Some(lhs) = evaluate_static_dim_expr_with_inputs(a, input_shapes)? else {
915                return Ok(None);
916            };
917            let Some(rhs) = evaluate_static_dim_expr_with_inputs(b, input_shapes)? else {
918                return Ok(None);
919            };
920            if rhs == 0 {
921                return Err(DimExprEvalError::FloorDivByZero { lhs, rhs });
922            }
923            Ok(Some(lhs / rhs))
924        }
925        DimExpr::Min(a, b) => {
926            let Some(lhs) = evaluate_static_dim_expr_with_inputs(a, input_shapes)? else {
927                return Ok(None);
928            };
929            let Some(rhs) = evaluate_static_dim_expr_with_inputs(b, input_shapes)? else {
930                return Ok(None);
931            };
932            Ok(Some(lhs.min(rhs)))
933        }
934        DimExpr::Max(a, b) => {
935            let Some(lhs) = evaluate_static_dim_expr_with_inputs(a, input_shapes)? else {
936                return Ok(None);
937            };
938            let Some(rhs) = evaluate_static_dim_expr_with_inputs(b, input_shapes)? else {
939                return Ok(None);
940            };
941            Ok(Some(lhs.max(rhs)))
942        }
943    }
944}
945
946fn compile_materialized_semantic_program(
947    compiled: &CompiledProgram<StdTensorOp>,
948    descriptors: &[InputDescriptor],
949    scoped_constraints: &[SlotScopedShapeConstraint],
950) -> Result<FrozenProgram> {
951    if compiled.input_slots.len() != descriptors.len() {
952        return Err(Error::runtime_state(
953            "graph_compile_semantic",
954            crate::ErrorPhase::Compile,
955            "materialized input count does not match semantic descriptors",
956        ));
957    }
958
959    let mut builder = SemanticProgramBuilder::new();
960    let mut values = vec![None; compiled.n_slots];
961    let mut slot_shapes = vec![None; compiled.n_slots];
962    let mut guard_slot_shapes = vec![None; compiled.n_slots];
963    let mut slot_dtypes = vec![None; compiled.n_slots];
964    for (input_idx, (&slot, descriptor)) in compiled.input_slots.iter().zip(descriptors).enumerate()
965    {
966        let Some(value_slot) = values.get_mut(slot) else {
967            return Err(invalid_compiled_graph(format!(
968                "semantic input slot {slot} is outside slot table of length {}",
969                compiled.n_slots
970            )));
971        };
972        let semantic_shape = descriptor.semantic_shape(input_idx);
973        let spec = match descriptor.scalar_identity {
974            Some(identity) => ProgramInputSpec::new(descriptor.dtype, semantic_shape.clone())
975                .with_scalar_identity(identity),
976            None => ProgramInputSpec::new(descriptor.dtype, semantic_shape.clone()),
977        };
978        let value = builder.input(spec).map_err(semantic_build_error)?;
979        if let Some(tensor) = &descriptor.default_tensor {
980            builder
981                .bind_input_retained(value, Arc::clone(tensor))
982                .map_err(semantic_build_error)?;
983        }
984        *value_slot = Some(value);
985        slot_shapes[slot] = Some(semantic_shape);
986        guard_slot_shapes[slot] = Some(descriptor.constraint_guard_shape(input_idx));
987        slot_dtypes[slot] = Some(descriptor.dtype);
988    }
989
990    for instruction in &compiled.instructions {
991        let inputs = instruction
992            .inputs
993            .iter()
994            .map(|&slot| {
995                values.get(slot).and_then(|value| *value).ok_or_else(|| {
996                    invalid_compiled_graph(format!(
997                        "semantic operation input slot {slot} is unavailable"
998                    ))
999                })
1000            })
1001            .collect::<Result<Vec<_>>>()?;
1002        let guard_input_shapes = instruction
1003            .inputs
1004            .iter()
1005            .map(|&slot| {
1006                guard_slot_shapes
1007                    .get(slot)
1008                    .and_then(|shape| shape.as_deref())
1009                    .ok_or_else(|| {
1010                        invalid_compiled_graph(format!(
1011                            "semantic operation guard-shape input slot {slot} is unavailable"
1012                        ))
1013                    })
1014            })
1015            .collect::<Result<Vec<_>>>()?;
1016        let guard_output_shapes = match &instruction.operation {
1017            StdTensorOp::Extension(extension) => {
1018                let input_dtypes = instruction
1019                    .inputs
1020                    .iter()
1021                    .map(|&slot| {
1022                        slot_dtypes
1023                            .get(slot)
1024                            .and_then(|dtype| *dtype)
1025                            .ok_or_else(|| {
1026                                invalid_compiled_graph(format!(
1027                                    "semantic operation dtype input slot {slot} is unavailable"
1028                                ))
1029                            })
1030                    })
1031                    .collect::<Result<Vec<_>>>()?;
1032                infer_extension_output_meta(extension.as_ref(), &input_dtypes, &guard_input_shapes)?
1033                    .into_iter()
1034                    .map(|(_dtype, shape)| shape)
1035                    .collect::<Vec<_>>()
1036            }
1037            operation => infer_output_shapes(operation, &guard_input_shapes)?,
1038        };
1039        let outputs = match &instruction.operation {
1040            StdTensorOp::Extension(extension) => builder
1041                .add_extension(Arc::clone(extension), &inputs)
1042                .map_err(semantic_build_error)?,
1043            operation => builder
1044                .add_op(
1045                    CoreSemanticOp::try_from(operation).map_err(|source| {
1046                        Error::runtime_state_source(
1047                            "graph_compile_semantic",
1048                            crate::ErrorPhase::Compile,
1049                            source,
1050                        )
1051                    })?,
1052                    &inputs,
1053                )
1054                .map_err(semantic_build_error)?,
1055        };
1056        if outputs.len() != instruction.outputs.len() {
1057            return Err(invalid_compiled_graph(format!(
1058                "semantic operation produced {} outputs for {} materialized slots",
1059                outputs.len(),
1060                instruction.outputs.len()
1061            )));
1062        }
1063        if guard_output_shapes.len() != instruction.outputs.len() {
1064            return Err(invalid_compiled_graph(format!(
1065                "semantic operation inferred {} guard shapes for {} materialized slots",
1066                guard_output_shapes.len(),
1067                instruction.outputs.len()
1068            )));
1069        }
1070        for ((&slot, &value), guard_shape) in instruction
1071            .outputs
1072            .iter()
1073            .zip(outputs.iter())
1074            .zip(guard_output_shapes.iter())
1075        {
1076            let Some(value_slot) = values.get_mut(slot) else {
1077                return Err(invalid_compiled_graph(format!(
1078                    "semantic output slot {slot} is outside slot table of length {}",
1079                    compiled.n_slots
1080                )));
1081            };
1082            if value_slot.replace(value).is_some() {
1083                return Err(invalid_compiled_graph(format!(
1084                    "semantic output slot {slot} has multiple producers"
1085                )));
1086            }
1087            let metadata = builder
1088                .value_metadata(value)
1089                .map_err(semantic_build_error)?;
1090            slot_shapes[slot] = Some(
1091                metadata
1092                    .shape()
1093                    .iter()
1094                    .enumerate()
1095                    .map(|(axis, extent)| match extent {
1096                        tenferro_ops::ShapeExtent::Exact(expression)
1097                        | tenferro_ops::ShapeExtent::UpperBound(expression) => expression.clone(),
1098                        tenferro_ops::ShapeExtent::Unknown => DimExpr::InputDim {
1099                            input_idx: slot,
1100                            axis,
1101                        },
1102                    })
1103                    .collect(),
1104            );
1105            guard_slot_shapes[slot] = Some(guard_shape.clone());
1106            slot_dtypes[slot] = Some(metadata.dtype());
1107        }
1108    }
1109
1110    for scoped in scoped_constraints {
1111        let target = scoped
1112            .origin_slots
1113            .iter()
1114            .find_map(|&slot| values.get(slot).and_then(|value| *value))
1115            .ok_or_else(|| {
1116                invalid_compiled_graph(
1117                    "semantic shape constraint has no available origin".to_string(),
1118                )
1119            })?;
1120        let lhs = lower_scoped_dim_expr(
1121            &scoped.local.lhs,
1122            &scoped.input_slots,
1123            &guard_slot_shapes,
1124            &scoped.local,
1125        )?;
1126        let rhs = lower_scoped_dim_expr(
1127            &scoped.local.rhs,
1128            &scoped.input_slots,
1129            &guard_slot_shapes,
1130            &scoped.local,
1131        )?;
1132        let relation = match scoped.local.relation {
1133            tenferro_ops::ShapeRelation::Equal => ProgramShapeRelation::Equal,
1134        };
1135        let lowered = LocalShapeConstraint {
1136            source: scoped.local.source.clone(),
1137            relation: scoped.local.relation,
1138            lhs,
1139            rhs,
1140        };
1141        let retained_guards = discharge(vec![lowered])?;
1142        if retained_guards.is_empty() {
1143            continue;
1144        }
1145        let guards = retained_guards.into_iter().map(|guard| {
1146            ProgramShapeGuard::new(relation, guard.lhs, guard.rhs)
1147                .with_source_family(guard.source.family_id)
1148        });
1149        builder
1150            .add_shape_guards_to_output(target, guards)
1151            .map_err(semantic_build_error)?;
1152    }
1153
1154    let outputs = compiled
1155        .output_slots
1156        .iter()
1157        .map(|&slot| {
1158            values.get(slot).and_then(|value| *value).ok_or_else(|| {
1159                invalid_compiled_graph(format!(
1160                    "semantic program output slot {slot} is unavailable"
1161                ))
1162            })
1163        })
1164        .collect::<Result<Vec<_>>>()?;
1165    builder.finish(&outputs).map_err(|source| {
1166        Error::runtime_state_source("graph_compile_semantic", crate::ErrorPhase::Compile, source)
1167    })
1168}
1169
1170fn semantic_build_error(source: crate::program::ProgramBuildError) -> Error {
1171    Error::runtime_state_source("graph_compile_semantic", crate::ErrorPhase::Compile, source)
1172}
1173
1174impl Default for GraphCompiler {
1175    fn default() -> Self {
1176        Self::new()
1177    }
1178}
1179
1180fn collect_default_inputs(
1181    outputs: &[&TracedTensor],
1182) -> Result<HashMap<TensorInputKey, Arc<RetainedValue>>> {
1183    let mut all_inputs = HashMap::new();
1184    for output in outputs {
1185        for (key, tensor) in output.inputs_map.iter() {
1186            if let Some(existing) = all_inputs.get(key) {
1187                if !default_tensors_equivalent(existing, tensor) {
1188                    return Err(Error::DuplicateBinding {
1189                        input_key: format!("{:?}", key),
1190                    });
1191                }
1192                continue;
1193            }
1194            all_inputs.insert(key.clone(), tensor.clone());
1195        }
1196    }
1197    Ok(all_inputs)
1198}
1199
1200fn insert_checkpoint_alias(
1201    aliases: &mut HashMap<TensorInputKey, ValueKey<StdTensorOp>>,
1202    alias_key: TensorInputKey,
1203    target_key: ValueKey<StdTensorOp>,
1204) -> Result<()> {
1205    if let Some(existing) = aliases.get(&alias_key) {
1206        if existing != &target_key {
1207            return Err(Error::Internal(format!(
1208                "checkpoint alias {alias_key:?} targets both {existing:?} and {target_key:?}"
1209            )));
1210        }
1211        return Ok(());
1212    }
1213    aliases.insert(alias_key, target_key);
1214    Ok(())
1215}
1216
1217struct AliasAwareMaterializer<'a> {
1218    view: &'a ResolvedView<StdTensorOp>,
1219    aliases: &'a HashMap<TensorInputKey, ValueKey<StdTensorOp>>,
1220    alias_shapes: &'a HashMap<TensorInputKey, Vec<usize>>,
1221    val_map: HashMap<ValueKey<StdTensorOp>, usize>,
1222    op_map: HashMap<Arc<OperationKey<StdTensorOp>>, usize>,
1223    values: Vec<MaterializedValue<StdTensorOp>>,
1224    operations: Vec<MaterializedOperation<StdTensorOp>>,
1225    input_keys: Vec<ValueKey<StdTensorOp>>,
1226}
1227
1228impl<'a> AliasAwareMaterializer<'a> {
1229    fn new(
1230        view: &'a ResolvedView<StdTensorOp>,
1231        aliases: &'a HashMap<TensorInputKey, ValueKey<StdTensorOp>>,
1232        alias_shapes: &'a HashMap<TensorInputKey, Vec<usize>>,
1233    ) -> Self {
1234        Self {
1235            view,
1236            aliases,
1237            alias_shapes,
1238            val_map: HashMap::new(),
1239            op_map: HashMap::new(),
1240            values: Vec::new(),
1241            operations: Vec::new(),
1242            input_keys: Vec::new(),
1243        }
1244    }
1245
1246    fn visit(&mut self, key: &ValueKey<StdTensorOp>) -> usize {
1247        if let Some(&index) = self.val_map.get(key) {
1248            return index;
1249        }
1250        if let Some(target) = self.alias_target(key).cloned() {
1251            let index = self.visit(&target);
1252            let index = self.refine_alias_to_checkpoint_shape(key, index);
1253            self.val_map.insert(key.clone(), index);
1254            return index;
1255        }
1256
1257        let resolved = self.view.resolve_value(key);
1258        assert!(
1259            resolved.is_some(),
1260            "key not found in resolved view: {:?}",
1261            key
1262        );
1263        match resolved {
1264            Some(ValueDef::Input { .. }) => self.materialize_input(key),
1265            Some(ValueDef::Produced {
1266                operation,
1267                input_keys,
1268                role,
1269                output_slot,
1270            }) => self.materialize_produced(operation, input_keys, role, output_slot),
1271            None => unreachable!("asserted above"),
1272        }
1273    }
1274
1275    fn alias_target(&self, key: &ValueKey<StdTensorOp>) -> Option<&ValueKey<StdTensorOp>> {
1276        let ValueKey::Input(input_key) = key else {
1277            return None;
1278        };
1279        self.aliases.get(input_key)
1280    }
1281
1282    fn refine_alias_to_checkpoint_shape(
1283        &mut self,
1284        key: &ValueKey<StdTensorOp>,
1285        target_index: usize,
1286    ) -> usize {
1287        let ValueKey::Input(input_key) = key else {
1288            return target_index;
1289        };
1290        let Some(shape) = self.alias_shapes.get(input_key) else {
1291            return target_index;
1292        };
1293        let target_key = self.values[target_index].key.clone();
1294        let rank = shape.len();
1295        self.materialize_produced(
1296            StdTensorOp::Slice(SliceConfig {
1297                starts: vec![0; rank],
1298                limits: shape.clone(),
1299                strides: vec![1; rank],
1300            }),
1301            vec![target_key],
1302            computegraph::types::OperationRole::Primary,
1303            0,
1304        )
1305    }
1306
1307    fn materialize_input(&mut self, key: &ValueKey<StdTensorOp>) -> usize {
1308        let index = self.values.len();
1309        self.values.push(MaterializedValue {
1310            key: key.clone(),
1311            producer: None,
1312        });
1313        self.val_map.insert(key.clone(), index);
1314        self.input_keys.push(key.clone());
1315        index
1316    }
1317
1318    fn materialize_produced(
1319        &mut self,
1320        operation: StdTensorOp,
1321        input_keys: Vec<ValueKey<StdTensorOp>>,
1322        role: computegraph::types::OperationRole,
1323        output_slot: usize,
1324    ) -> usize {
1325        let op_key = Arc::new(OperationKey::new(
1326            operation.clone(),
1327            input_keys.clone(),
1328            role.clone(),
1329        ));
1330
1331        if self.op_map.contains_key(&op_key) {
1332            let output_key = ValueKey::Derived {
1333                operation: op_key,
1334                output_slot: output_slot as u8,
1335            };
1336            let val_index = self.val_map.get(&output_key).copied();
1337            assert!(
1338                val_index.is_some(),
1339                "materialized op {:?} is missing output slot {}",
1340                operation,
1341                output_slot
1342            );
1343            return match val_index {
1344                Some(index) => index,
1345                None => unreachable!("asserted above"),
1346            };
1347        }
1348
1349        let materialized_inputs = input_keys.iter().map(|input| self.visit(input)).collect();
1350        let op_index = self.operations.len();
1351        self.op_map.insert(Arc::clone(&op_key), op_index);
1352        self.operations.push(MaterializedOperation {
1353            operation: operation.clone(),
1354            inputs: materialized_inputs,
1355            outputs: Vec::with_capacity(operation.output_count()),
1356            role,
1357        });
1358
1359        for slot in 0..operation.output_count() {
1360            let output_key = ValueKey::Derived {
1361                operation: Arc::clone(&op_key),
1362                output_slot: slot as u8,
1363            };
1364            let val_index = self.values.len();
1365            self.values.push(MaterializedValue {
1366                key: output_key.clone(),
1367                producer: Some((op_index, slot)),
1368            });
1369            self.val_map.insert(output_key, val_index);
1370            self.operations[op_index].outputs.push(val_index);
1371        }
1372
1373        self.operations[op_index].outputs[output_slot]
1374    }
1375}
1376
1377fn materialize_merge_with_input_aliases(
1378    view: &ResolvedView<StdTensorOp>,
1379    outputs: &[ValueKey<StdTensorOp>],
1380    aliases: &HashMap<TensorInputKey, ValueKey<StdTensorOp>>,
1381    alias_shapes: &HashMap<TensorInputKey, Vec<usize>>,
1382) -> MaterializedGraph<StdTensorOp> {
1383    let mut materializer = AliasAwareMaterializer::new(view, aliases, alias_shapes);
1384    let mut materialized_outputs = Vec::with_capacity(outputs.len());
1385
1386    for output in outputs {
1387        let output_slot = materializer.visit(output);
1388        materialized_outputs.push(materializer.values[output_slot].key.clone());
1389    }
1390
1391    MaterializedGraph {
1392        values: materializer.values,
1393        operations: materializer.operations,
1394        inputs: materializer.input_keys,
1395        outputs: materialized_outputs,
1396    }
1397}
1398
1399fn validate_placeholder_spec(
1400    index: usize,
1401    placeholder: &TracedTensor,
1402    dtype: DType,
1403    shape: &[usize],
1404) -> Result<()> {
1405    if placeholder.data.is_some() {
1406        return Err(Error::UnexpectedBinding {
1407            binding_index: index,
1408        });
1409    }
1410    placeholder.input_key().ok_or(Error::UnexpectedBinding {
1411        binding_index: index,
1412    })?;
1413
1414    if placeholder.dtype != dtype {
1415        return Err(Error::PlaceholderDtypeMismatch {
1416            expected: placeholder.dtype,
1417            actual: dtype,
1418        });
1419    }
1420    validate_placeholder_shape(placeholder, shape)
1421}
1422
1423fn validate_placeholder_shape(placeholder: &TracedTensor, shape: &[usize]) -> Result<()> {
1424    match try_concrete_shape(placeholder) {
1425        Some(expected_shape) => {
1426            if expected_shape.as_slice() != shape {
1427                return Err(Error::PlaceholderShapeMismatch {
1428                    expected: expected_shape,
1429                    actual: shape.to_vec(),
1430                });
1431            }
1432        }
1433        None => {
1434            if placeholder.rank != shape.len() {
1435                return Err(Error::PlaceholderRankMismatch {
1436                    expected: placeholder.rank,
1437                    actual: shape.len(),
1438                });
1439            }
1440        }
1441    }
1442    Ok(())
1443}
1444
1445fn descriptor_for_input(
1446    key: &TensorInputKey,
1447    binding_specs: &HashMap<TensorInputKey, InputDescriptor>,
1448    default_inputs: &HashMap<TensorInputKey, Arc<RetainedValue>>,
1449    allow_unbound_placeholders: bool,
1450) -> Result<InputDescriptor> {
1451    if let Some(tensor) = default_inputs.get(key) {
1452        return Ok(InputDescriptor {
1453            dtype: tensor.dtype(),
1454            shape: tensor.shape().to_vec(),
1455            extent_identity: default_input_extent_identity(key, tensor)?,
1456            default_tensor: Some(tensor.clone()),
1457            // A bound tensor carries no canonical name itself, so the declared
1458            // identity comes from the traced value's registered metadata.
1459            scalar_identity: registered_meta(&ValueKey::Input(key.clone()))
1460                .ok()
1461                .and_then(|metadata| metadata.scalar_identity()),
1462        });
1463    }
1464    if let Some(spec) = binding_specs.get(key) {
1465        return Ok(spec.clone());
1466    }
1467    if allow_unbound_placeholders {
1468        return descriptor_for_unbound_input(key);
1469    }
1470    Err(Error::UnboundPlaceholder {
1471        input_key: format!("{:?}", key),
1472    })
1473}
1474
1475fn descriptor_for_unbound_input(key: &TensorInputKey) -> Result<InputDescriptor> {
1476    let metadata = registered_meta(&ValueKey::Input(key.clone()))?;
1477    if let Some(shape) = metadata
1478        .exact_shape()
1479        .as_deref()
1480        .and_then(concrete_shape_from_sym_dims)
1481    {
1482        return Ok(InputDescriptor {
1483            dtype: metadata.dtype,
1484            shape,
1485            extent_identity: InputExtentIdentity::Concrete,
1486            default_tensor: None,
1487            scalar_identity: metadata.scalar_identity(),
1488        });
1489    }
1490    Ok(InputDescriptor {
1491        dtype: metadata.dtype,
1492        shape: vec![0; metadata.rank()],
1493        extent_identity: InputExtentIdentity::Symbolic,
1494        default_tensor: None,
1495        scalar_identity: metadata.scalar_identity(),
1496    })
1497}
1498
1499fn concrete_shape_from_sym_dims(shape: &[SymDim]) -> Option<Vec<usize>> {
1500    shape.iter().map(SymDim::constant_value).collect()
1501}
1502
1503fn default_input_extent_identity(
1504    key: &TensorInputKey,
1505    tensor: &RetainedValue,
1506) -> Result<InputExtentIdentity> {
1507    let metadata = registered_meta(&ValueKey::Input(key.clone()))?;
1508    let exact_shape = metadata.exact_shape();
1509    if metadata.dtype == tensor.dtype()
1510        && exact_shape_matches_tensor_shape(exact_shape.as_deref(), tensor.shape())
1511    {
1512        Ok(InputExtentIdentity::Concrete)
1513    } else {
1514        Ok(InputExtentIdentity::Symbolic)
1515    }
1516}
1517
1518fn exact_shape_matches_tensor_shape(
1519    shape: Option<&[tenferro_ops::SymDim]>,
1520    tensor_shape: &[usize],
1521) -> bool {
1522    let Some(shape) = shape else {
1523        return false;
1524    };
1525    shape.len() == tensor_shape.len()
1526        && shape
1527            .iter()
1528            .zip(tensor_shape)
1529            .all(|(dim, &extent)| dim.constant_value() == Some(extent))
1530}
1531
1532fn prune_compiled_extension_outputs(prog: &mut CompiledProgram<StdTensorOp>) -> Result<()> {
1533    let mut live_slots = vec![false; prog.n_slots];
1534    for &slot in &prog.output_slots {
1535        let Some(live) = live_slots.get_mut(slot) else {
1536            return Err(invalid_compiled_graph(format!(
1537                "program output slot {slot} is outside slot table of length {}",
1538                prog.n_slots
1539            )));
1540        };
1541        *live = true;
1542    }
1543
1544    for instr in prog.instructions.iter_mut().rev() {
1545        let live_outputs = instr
1546            .outputs
1547            .iter()
1548            .map(|&slot| {
1549                live_slots.get(slot).copied().ok_or_else(|| {
1550                    invalid_compiled_graph(format!(
1551                        "instruction output slot {slot} is outside slot table of length {}",
1552                        prog.n_slots
1553                    ))
1554                })
1555            })
1556            .collect::<Result<Vec<_>>>()?;
1557
1558        if let StdTensorOp::Extension(ext) = &instr.operation
1559            && let Some(pruned) = ext.prune_outputs(&live_outputs)
1560        {
1561            let kept_outputs = instr
1562                .outputs
1563                .iter()
1564                .zip(live_outputs.iter())
1565                .filter_map(|(&slot, &live)| live.then_some(slot))
1566                .collect::<Vec<_>>();
1567            if pruned.output_count() != kept_outputs.len() {
1568                return Err(invalid_compiled_graph(format!(
1569                    "extension family_id={:?} pruned to {} outputs for {} live slots",
1570                    ext.family_id(),
1571                    pruned.output_count(),
1572                    kept_outputs.len()
1573                )));
1574            }
1575            instr.operation = StdTensorOp::Extension(pruned);
1576            instr.outputs = kept_outputs;
1577        }
1578
1579        if live_outputs.iter().any(|&live| live) {
1580            for &slot in &instr.inputs {
1581                let Some(live) = live_slots.get_mut(slot) else {
1582                    return Err(invalid_compiled_graph(format!(
1583                        "instruction input slot {slot} is outside slot table of length {}",
1584                        prog.n_slots
1585                    )));
1586                };
1587                *live = true;
1588            }
1589        }
1590    }
1591
1592    Ok(())
1593}
1594
1595fn invalid_compiled_graph(message: impl Into<String>) -> Error {
1596    Error::Internal(message.into())
1597}
1598
1599fn default_tensors_equivalent(lhs: &Arc<RetainedValue>, rhs: &Arc<RetainedValue>) -> bool {
1600    if Arc::ptr_eq(lhs, rhs) {
1601        return true;
1602    }
1603    if lhs.dtype() != rhs.dtype() || lhs.shape() != rhs.shape() {
1604        return false;
1605    }
1606    match lhs.dtype() {
1607        DType::F32 => default_slices_equivalent::<f32>(lhs, rhs),
1608        DType::F64 => default_slices_equivalent::<f64>(lhs, rhs),
1609        DType::I32 => default_slices_equivalent::<i32>(lhs, rhs),
1610        DType::I64 => default_slices_equivalent::<i64>(lhs, rhs),
1611        DType::Bool => default_slices_equivalent::<bool>(lhs, rhs),
1612        DType::C32 => default_slices_equivalent::<Complex32>(lhs, rhs),
1613        DType::C64 => default_slices_equivalent::<Complex64>(lhs, rhs),
1614        // An externally defined payload is opaque here, so two distinct values of
1615        // that kind are reported as not equivalent rather than compared by bytes.
1616        DType::External(_) => false,
1617    }
1618}
1619
1620fn default_slices_equivalent<T: TensorScalar + PartialEq>(
1621    lhs: &RetainedValue,
1622    rhs: &RetainedValue,
1623) -> bool {
1624    let (Ok(lhs), Ok(rhs)) = (lhs.tensor_read(), rhs.tensor_read()) else {
1625        return false;
1626    };
1627    let lhs = lhs.tensor_view();
1628    let rhs = rhs.tensor_view();
1629    match (lhs.as_slice::<T>(), rhs.as_slice::<T>()) {
1630        (Ok(lhs), Ok(rhs)) => lhs == rhs,
1631        // Backend-resident defaults cannot be inspected here; only the same
1632        // value handle is considered equivalent by `default_tensors_equivalent`.
1633        _ => false,
1634    }
1635}
1636
1637#[cfg(test)]
1638mod constraint_scope_tests;
1639
1640#[cfg(test)]
1641mod test_support;
1642
1643#[cfg(test)]
1644mod tests {
1645    use super::*;
1646    use std::any::Any;
1647    use std::hash::Hasher;
1648    use std::sync::Arc;
1649    use tenferro_ops::{
1650        ext_op::{ExtensionAliasDeclaration, ExtensionEffectDeclaration, ExtensionOp},
1651        SymDim,
1652    };
1653    use tenferro_tensor::{
1654        BackendStorageHandle, DeviceId, DeviceKind, GpuBackendKind, MemoryKind, Placement,
1655        StorageBuffer, TypedTensor,
1656    };
1657
1658    #[test]
1659    fn compile_publishes_semantic_program_and_separate_default_bindings() {
1660        let input = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1661        let output = input.neg().unwrap();
1662
1663        let program = GraphCompiler::new().compile(&output).unwrap();
1664
1665        assert_eq!(program.program().inputs().len(), 1);
1666        assert_eq!(program.program().outputs().len(), 1);
1667        assert_eq!(program.program().operations().count(), 1);
1668        assert_eq!(program.bindings().len(), 1);
1669        assert_eq!(
1670            program
1671                .program()
1672                .value_metadata(program.program().inputs()[0])
1673                .unwrap()
1674                .shape(),
1675            &[ShapeExtent::Exact(DimExpr::Const(2))]
1676        );
1677        assert_eq!(
1678            program
1679                .program()
1680                .value_metadata(program.program().outputs()[0])
1681                .unwrap()
1682                .dtype(),
1683            DType::F64
1684        );
1685    }
1686
1687    #[test]
1688    fn compile_preserves_symbolic_default_input_extent_identity() {
1689        let tensor = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1690        let input = TracedTensor::from_tensor_symbolic_shape(tensor).unwrap();
1691        let output = input.neg().unwrap();
1692
1693        let program = GraphCompiler::new().compile(&output).unwrap();
1694
1695        assert!(matches!(
1696            program
1697                .program()
1698                .value_metadata(program.program().inputs()[0])
1699                .unwrap()
1700                .shape(),
1701            [ShapeExtent::Exact(DimExpr::InputDim {
1702                input_idx: 0,
1703                axis: 0
1704            })]
1705        ));
1706    }
1707
1708    #[test]
1709    fn compile_many_rejects_conflicting_default_inputs_for_same_key() {
1710        let x = TracedTensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap();
1711        let y1 = x.neg().unwrap();
1712        let mut y2 = x.neg().unwrap();
1713        let key = x.input_key().expect("concrete traced tensor has input key");
1714        let replacement = Arc::new(RetainedValue::from_tensor(
1715            Tensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap(),
1716        ));
1717        let mut inputs = (*y2.inputs_map).clone();
1718        inputs.insert(key.clone(), replacement);
1719        y2.inputs_map = Arc::new(inputs);
1720
1721        let err = GraphCompiler::new().compile_many(&[&y1, &y2]).unwrap_err();
1722
1723        assert!(matches!(
1724            err,
1725            Error::DuplicateBinding { ref input_key } if input_key.contains(&format!("{key:?}"))
1726        ));
1727    }
1728
1729    #[test]
1730    fn default_tensors_equivalent_rejects_distinct_backend_buffers() {
1731        let placement = Placement {
1732            memory_kind: MemoryKind::Device,
1733            device: Some(DeviceId {
1734                kind: DeviceKind::Gpu(GpuBackendKind::Cuda),
1735                ordinal: 0,
1736            }),
1737            cpu_affinity: None,
1738        };
1739        let lhs = Arc::new(RetainedValue::from_tensor(Tensor::from_typed::<f64>(
1740            TypedTensor::from_buffer_col_major(
1741                vec![2],
1742                StorageBuffer::Backend(Box::new(BackendStorageHandle::<f64>::new_with_len(1, 2))),
1743                placement.clone(),
1744            )
1745            .unwrap(),
1746        )));
1747        let rhs = Arc::new(RetainedValue::from_tensor(Tensor::from_typed::<f64>(
1748            TypedTensor::from_buffer_col_major(
1749                vec![2],
1750                StorageBuffer::Backend(Box::new(BackendStorageHandle::<f64>::new_with_len(2, 2))),
1751                placement,
1752            )
1753            .unwrap(),
1754        )));
1755
1756        assert!(
1757            !default_tensors_equivalent(&lhs, &rhs),
1758            "distinct backend-resident default tensors must not compare equal just because both are unreadable on host"
1759        );
1760        assert!(default_tensors_equivalent(&lhs, &lhs));
1761    }
1762
1763    #[test]
1764    fn compile_frozen_program_rejects_static_unbound_shape_guard_mismatch() {
1765        let mut builder = SemanticProgramBuilder::new();
1766        let lhs = builder
1767            .input(ProgramInputSpec::new(DType::F64, [DimExpr::Const(2)]))
1768            .unwrap();
1769        let rhs = builder
1770            .input(ProgramInputSpec::new(DType::F64, [DimExpr::Const(3)]))
1771            .unwrap();
1772        let output = builder.add_op(CoreSemanticOp::Neg, &[lhs]).unwrap()[0];
1773        builder
1774            .add_shape_guards_to_output(
1775                output,
1776                [ProgramShapeGuard::new(
1777                    ProgramShapeRelation::Equal,
1778                    DimExpr::InputDim {
1779                        input_idx: 0,
1780                        axis: 0,
1781                    },
1782                    DimExpr::InputDim {
1783                        input_idx: 1,
1784                        axis: 0,
1785                    },
1786                )],
1787            )
1788            .unwrap();
1789        let frozen = builder.finish(&[output]).unwrap();
1790
1791        let err = GraphCompiler::new()
1792            .compile_frozen_program(&frozen)
1793            .unwrap_err();
1794
1795        assert!(matches!(
1796            err,
1797            Error::ShapeConstraintViolation {
1798                lhs_value: 2,
1799                rhs_value: 3,
1800                ..
1801            }
1802        ));
1803        let _ = rhs;
1804    }
1805
1806    #[test]
1807    fn compile_frozen_program_defers_dynamic_unbound_shape_guard() {
1808        let mut builder = SemanticProgramBuilder::new();
1809        let lhs = builder
1810            .input(ProgramInputSpec::new(DType::F64, [DimExpr::Const(2)]))
1811            .unwrap();
1812        let rhs = builder
1813            .input(ProgramInputSpec::from_metadata(
1814                ProgramValueMetadata::from_extents(DType::F64, [ShapeExtent::Unknown]),
1815            ))
1816            .unwrap();
1817        let output = builder.add_op(CoreSemanticOp::Neg, &[lhs]).unwrap()[0];
1818        builder
1819            .add_shape_guards_to_output(
1820                output,
1821                [ProgramShapeGuard::new(
1822                    ProgramShapeRelation::Equal,
1823                    DimExpr::InputDim {
1824                        input_idx: 0,
1825                        axis: 0,
1826                    },
1827                    DimExpr::InputDim {
1828                        input_idx: 1,
1829                        axis: 0,
1830                    },
1831                )],
1832            )
1833            .unwrap();
1834        let frozen = builder.finish(&[output]).unwrap();
1835
1836        GraphCompiler::new()
1837            .compile_frozen_program(&frozen)
1838            .unwrap();
1839        let _ = rhs;
1840    }
1841
1842    #[derive(Clone, Debug, PartialEq, Eq)]
1843    struct PrunableTestOp {
1844        pruned: bool,
1845    }
1846
1847    impl ExtensionOp for PrunableTestOp {
1848        fn family_id(&self) -> &'static str {
1849            "tenferro-runtime.test-prunable.v1"
1850        }
1851
1852        fn payload_hash(&self, hasher: &mut dyn Hasher) {
1853            hasher.write_u8(u8::from(self.pruned));
1854        }
1855
1856        fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
1857            other
1858                .as_any()
1859                .downcast_ref::<Self>()
1860                .is_some_and(|that| self == that)
1861        }
1862
1863        fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
1864            Arc::new(self.clone())
1865        }
1866
1867        fn as_any(&self) -> &dyn Any {
1868            self
1869        }
1870
1871        fn input_count(&self) -> usize {
1872            1
1873        }
1874
1875        fn output_count(&self) -> usize {
1876            if self.pruned {
1877                1
1878            } else {
1879                3
1880            }
1881        }
1882
1883        fn semantic_effects(&self) -> ExtensionEffectDeclaration<'_> {
1884            ExtensionEffectDeclaration::Declared(&[])
1885        }
1886
1887        fn semantic_aliases(&self) -> ExtensionAliasDeclaration<'_> {
1888            ExtensionAliasDeclaration::AllFresh
1889        }
1890
1891        fn infer_output_meta(
1892            &self,
1893            ctx: &mut tenferro_ops::ExtensionShapeContext<'_>,
1894        ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1895            let dtype = ctx.input_dtype(0)?;
1896            let shape = ctx.input_shape(0)?.to_vec();
1897            Ok((0..self.output_count())
1898                .map(|_| (dtype, shape.clone()))
1899                .collect())
1900        }
1901
1902        fn prune_outputs(&self, live_outputs: &[bool]) -> Option<Arc<dyn ExtensionOp>> {
1903            (!self.pruned && live_outputs == [false, true, false])
1904                .then(|| Arc::new(Self { pruned: true }) as Arc<dyn ExtensionOp>)
1905        }
1906    }
1907
1908    #[test]
1909    fn compile_prunes_extension_outputs_with_replacement_op() {
1910        let input = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1911        let outputs =
1912            crate::extension::apply(Arc::new(PrunableTestOp { pruned: false }), &[&input]).unwrap();
1913
1914        let program = GraphCompiler::new().compile(&outputs[1]).unwrap();
1915        let staging =
1916            stage_semantic_program(program.program(), program.compiler_options()).unwrap();
1917        let pruned_instruction = staging
1918            .instructions
1919            .iter()
1920            .find_map(|inst| match &inst.op {
1921                crate::exec::ExecOp::Extension(op)
1922                    if op.family_id() == "tenferro-runtime.test-prunable.v1" =>
1923                {
1924                    Some((
1925                        inst.output_slots.clone(),
1926                        format!("{op:?}"),
1927                        op.output_count(),
1928                    ))
1929                }
1930                _ => None,
1931            })
1932            .expect("compiled program should contain the test extension");
1933
1934        assert_eq!(pruned_instruction.0.len(), 1);
1935        assert!(pruned_instruction.1.contains("pruned: true"));
1936        assert_eq!(pruned_instruction.2, 1);
1937    }
1938
1939    #[test]
1940    fn compiled_graph_input_keys_preserve_order_for_binary_graph() {
1941        let a = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1942        let b = TracedTensor::from_vec_col_major(vec![2], vec![3.0_f64, 4.0]).unwrap();
1943        let a_key = a.input_key().expect("concrete traced tensor has input key");
1944        let b_key = b.input_key().expect("concrete traced tensor has input key");
1945        let c = a.mul(&b).unwrap();
1946
1947        let program = GraphCompiler::new().compile_many(&[&c]).unwrap();
1948
1949        assert_eq!(program.input_count(), 2);
1950        assert_eq!(program.input_keys().len(), 2);
1951        let keys: Vec<_> = program.input_keys().to_vec();
1952        assert!(keys.contains(&a_key), "input_keys must contain a's key");
1953        assert!(keys.contains(&b_key), "input_keys must contain b's key");
1954        // input_key_index maps each key to its position
1955        assert_eq!(
1956            program.input_key_index(&a_key),
1957            keys.iter().position(|k| k == &a_key)
1958        );
1959        assert_eq!(
1960            program.input_key_index(&b_key),
1961            keys.iter().position(|k| k == &b_key)
1962        );
1963    }
1964
1965    #[test]
1966    fn compile_with_input_specs_reorders_inputs_by_explicit_order() {
1967        let a = TracedTensor::input_symbolic_shape(DType::F64, 1).unwrap();
1968        let b = TracedTensor::input_symbolic_shape(DType::F64, 1).unwrap();
1969        let a_key = a.input_key().expect("symbolic traced tensor has input key");
1970        let b_key = b.input_key().expect("symbolic traced tensor has input key");
1971        let c = a.mul(&b).unwrap();
1972
1973        // Request b first, then a
1974        let program = GraphCompiler::new()
1975            .compile_with_input_specs(&c, &[(&b, DType::F64, &[2]), (&a, DType::F64, &[2])])
1976            .unwrap();
1977
1978        assert_eq!(program.input_count(), 2);
1979        assert_eq!(program.input_keys().len(), 2);
1980        // b must be at position 0 per the explicit order
1981        assert_eq!(
1982            program.input_key_index(&b_key),
1983            Some(0),
1984            "b must be first in explicit order"
1985        );
1986        assert_eq!(
1987            program.input_key_index(&a_key),
1988            Some(1),
1989            "a must be second in explicit order"
1990        );
1991    }
1992}