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 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
78pub 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 pub fn new() -> Self {
121 Self {
122 extension_cache: ExtensionCacheStore::new(),
123 compiler_options: CompilerOptions::default(),
124 }
125 }
126
127 pub fn with_compiler_options(compiler_options: CompilerOptions) -> Self {
144 Self {
145 extension_cache: ExtensionCacheStore::new(),
146 compiler_options,
147 }
148 }
149
150 pub fn compile(&mut self, output: &TracedTensor) -> Result<CompiledGraph> {
173 self.compile_many(&[output])
174 }
175
176 pub fn compile_traced_graph(&mut self, graph: &TracedGraph) -> Result<CompiledGraph> {
189 self.compile_frozen(graph.frozen())
190 }
191
192 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 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 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 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 pub fn compiler_options(&self) -> CompilerOptions {
397 self.compiler_options
398 }
399
400 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 pub fn clear_extension_caches(&mut self) {
439 self.extension_cache.clear();
440 }
441
442 pub fn clear_caches(&mut self) {
454 self.clear_extension_caches();
455 }
456
457 pub fn cache_stats(&self) -> CacheStats {
469 self.extension_cache.stats(ExtensionCacheSelector::All)
470 }
471
472 pub fn extension_caches(&self) -> &ExtensionCacheStore {
483 &self.extension_cache
484 }
485
486 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 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 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 _ => 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 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 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 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}