1use std::collections::{HashMap, HashSet};
2use std::sync::Arc;
3
4use computegraph::graph::GraphBuilder;
5use computegraph::{LocalValueId, OperationRole, ValueRef};
6use tenferro_ops::input_key::TensorInputKey;
7use tenferro_ops::std_tensor_op::StdTensorOp;
8use tenferro_ops::{SymDim, TensorMeta};
9use tenferro_runtime::ad_support::{
10 allocate_input_key, allocate_shape_tensor_id, checkpoint_tensor, compile_ad_source,
11 frozen_input_value, inputs_map as tensor_inputs_map, leaf_input_key, merge_traced_leaf_metas,
12 metadata_scopes as tensor_metadata_scopes, metadata_scopes_with_new, ones_tensor,
13 register_scoped_graph_analysis, shape_hint as tensor_shape_hint, tensor_from_parts,
14 ConstraintScopeTransfer, RetainedValue, TracedTensorParts,
15};
16use tenferro_runtime::program::{FrozenProgram, ProgramValue, ProgramValueMetadata, SemanticOpRef};
17use tenferro_runtime::{
18 CompiledGraph, Error, ErrorPhase, GraphCompiler, Result, Runtime, Tensor, TracedTensor,
19};
20
21use crate::semantic_extension::{SemanticAdError, SemanticExtensionRuleSet};
22use crate::semantic_transform::{
23 semantic_jvp, semantic_vjp, SemanticAdProgram, SemanticAdTransformError,
24};
25use crate::transform_cache::{AdTransformCache, SemanticAdTransformCacheKey};
26
27pub(crate) fn next_input_key() -> TensorInputKey {
28 tenferro_runtime::ad_support::allocate_input_key()
29}
30
31fn error_shape_hint(tensor: &TracedTensor) -> Vec<usize> {
32 tensor
33 .try_concrete_shape()
34 .unwrap_or_else(|| vec![0; tensor.rank])
35}
36
37pub(crate) fn inactive_wrt_error(op: &'static str, wrt_key: impl std::fmt::Debug) -> Error {
42 Error::invalid_argument(
43 op,
44 ErrorPhase::GraphBuild,
45 "wrt",
46 format!(
47 "{op} is inactive for {wrt_key:?}; use {op}_optional to observe the inactive state"
48 ),
49 )
50}
51
52pub(crate) fn grad_with_rules_and_cache(
53 output: &TracedTensor,
54 wrt: &TracedTensor,
55 rules: &SemanticExtensionRuleSet,
56 ad_transform_cache: Option<&AdTransformCache>,
57) -> Result<TracedTensor> {
58 grad_with_optional_rules(output, wrt, rules, ad_transform_cache)
59}
60
61pub(crate) fn jvp_with_rules_and_cache(
62 output: &TracedTensor,
63 wrt: &TracedTensor,
64 tangent: &TracedTensor,
65 rules: &SemanticExtensionRuleSet,
66 ad_transform_cache: Option<&AdTransformCache>,
67) -> Result<TracedTensor> {
68 let wrt_input_key = leaf_input_key(wrt)?;
69 jvp_many_with_rules_and_cache(output, &[(wrt, tangent)], rules, ad_transform_cache)?
70 .ok_or_else(|| inactive_wrt_error("jvp", &wrt_input_key))
71}
72
73pub(crate) fn jvp_many_with_rules_and_cache(
74 output: &TracedTensor,
75 wrt_tangents: &[(&TracedTensor, &TracedTensor)],
76 rules: &SemanticExtensionRuleSet,
77 ad_transform_cache: Option<&AdTransformCache>,
78) -> Result<Option<TracedTensor>> {
79 jvp_many_optional_impl(output, wrt_tangents, rules, ad_transform_cache)
80}
81
82pub(crate) fn grad_optional_with_rules_and_cache(
83 output: &TracedTensor,
84 wrt: &TracedTensor,
85 rules: &SemanticExtensionRuleSet,
86 ad_transform_cache: Option<&AdTransformCache>,
87) -> Result<Option<TracedTensor>> {
88 if output.rank != 0 {
89 return Err(Error::NonScalarGrad {
90 shape: error_shape_hint(output),
91 });
92 }
93
94 let ones = ones_tensor(output.dtype, vec![])?;
95 let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
96 vjp_many_with_transform_and_cache(output, &[wrt], &seed, rules, "grad", ad_transform_cache)
97 .map(|mut results| results.pop().flatten())
98}
99
100pub(crate) fn jvp_optional_with_rules_and_cache(
101 output: &TracedTensor,
102 wrt: &TracedTensor,
103 tangent: &TracedTensor,
104 rules: &SemanticExtensionRuleSet,
105 ad_transform_cache: Option<&AdTransformCache>,
106) -> Result<Option<TracedTensor>> {
107 jvp_many_with_rules_and_cache(output, &[(wrt, tangent)], rules, ad_transform_cache)
108}
109
110pub(crate) fn vjp_with_rules_and_cache(
111 output: &TracedTensor,
112 wrt: &TracedTensor,
113 cotangent: &TracedTensor,
114 rules: &SemanticExtensionRuleSet,
115 ad_transform_cache: Option<&AdTransformCache>,
116) -> Result<TracedTensor> {
117 let wrt_input_key = leaf_input_key(wrt)?;
118 vjp_many_with_rules_and_cache(output, &[wrt], cotangent, rules, ad_transform_cache)?
119 .into_iter()
120 .next()
121 .flatten()
122 .ok_or_else(|| inactive_wrt_error("vjp", &wrt_input_key))
123}
124
125pub(crate) fn vjp_many_with_rules_and_cache(
126 output: &TracedTensor,
127 wrts: &[&TracedTensor],
128 cotangent: &TracedTensor,
129 rules: &SemanticExtensionRuleSet,
130 ad_transform_cache: Option<&AdTransformCache>,
131) -> Result<Vec<Option<TracedTensor>>> {
132 vjp_many_with_transform_and_cache(output, wrts, cotangent, rules, "vjp", ad_transform_cache)
133}
134
135fn vjp_many_with_transform_and_cache(
136 output: &TracedTensor,
137 wrts: &[&TracedTensor],
138 cotangent: &TracedTensor,
139 rules: &SemanticExtensionRuleSet,
140 transform: &'static str,
141 ad_transform_cache: Option<&AdTransformCache>,
142) -> Result<Vec<Option<TracedTensor>>> {
143 vjp_many_optional_impl(
144 output,
145 wrts,
146 cotangent,
147 rules,
148 transform,
149 ad_transform_cache,
150 )
151}
152
153pub(crate) fn vjp_optional_with_rules_and_cache(
154 output: &TracedTensor,
155 wrt: &TracedTensor,
156 cotangent: &TracedTensor,
157 rules: &SemanticExtensionRuleSet,
158 ad_transform_cache: Option<&AdTransformCache>,
159) -> Result<Option<TracedTensor>> {
160 vjp_many_with_rules_and_cache(output, &[wrt], cotangent, rules, ad_transform_cache)
161 .map(|mut results| results.pop().flatten())
162}
163
164fn grad_with_optional_rules(
165 output: &TracedTensor,
166 wrt: &TracedTensor,
167 rules: &SemanticExtensionRuleSet,
168 ad_transform_cache: Option<&AdTransformCache>,
169) -> Result<TracedTensor> {
170 if output.rank != 0 {
171 return Err(Error::NonScalarGrad {
172 shape: error_shape_hint(output),
173 });
174 }
175
176 let ones = ones_tensor(output.dtype, vec![])?;
177 let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
178 let wrt_input_key = leaf_input_key(wrt)?;
179 vjp_many_with_transform_and_cache(output, &[wrt], &seed, rules, "grad", ad_transform_cache)?
180 .pop()
181 .flatten()
182 .ok_or_else(|| inactive_wrt_error("grad", &wrt_input_key))
183}
184
185fn single_runtime_output(mut outputs: Vec<Tensor>, op: &'static str) -> Result<Tensor> {
186 let actual = outputs.len();
187 if actual != 1 {
188 return Err(Error::runtime_state(
189 op,
190 ErrorPhase::Execution,
191 format!("expected one runtime output, got {actual}"),
192 ));
193 }
194 outputs.pop().ok_or_else(|| {
195 Error::runtime_state(
196 op,
197 ErrorPhase::Execution,
198 "runtime returned no output after successful output-count validation",
199 )
200 })
201}
202
203pub trait TracedTensorAdExt {
217 fn grad(&self, wrt: &TracedTensor) -> Result<TracedTensor>;
265
266 fn grad_optional(&self, wrt: &TracedTensor) -> Result<Option<TracedTensor>>;
294
295 fn checkpoint(&mut self, compiler: &mut GraphCompiler, runtime: &Runtime) -> Result<()>;
328
329 fn jvp(&self, wrt: &TracedTensor, tangent: &TracedTensor) -> Result<TracedTensor>;
373
374 fn jvp_optional(
403 &self,
404 wrt: &TracedTensor,
405 tangent: &TracedTensor,
406 ) -> Result<Option<TracedTensor>>;
407
408 fn vjp(&self, wrt: &TracedTensor, cotangent: &TracedTensor) -> Result<TracedTensor>;
457
458 fn vjp_optional(
487 &self,
488 wrt: &TracedTensor,
489 cotangent: &TracedTensor,
490 ) -> Result<Option<TracedTensor>>;
491}
492
493impl TracedTensorAdExt for TracedTensor {
494 fn grad(&self, wrt: &TracedTensor) -> Result<TracedTensor> {
495 let rules = SemanticExtensionRuleSet::default();
496 grad_with_optional_rules(self, wrt, &rules, None)
497 }
498
499 fn grad_optional(&self, wrt: &TracedTensor) -> Result<Option<TracedTensor>> {
500 if self.rank != 0 {
501 return Err(Error::NonScalarGrad {
502 shape: error_shape_hint(self),
503 });
504 }
505
506 let ones = ones_tensor(self.dtype, vec![])?;
507 let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
508 let rules = SemanticExtensionRuleSet::default();
509 vjp_many_with_transform_and_cache(self, &[wrt], &seed, &rules, "grad", None)
510 .map(|mut results| results.pop().flatten())
511 }
512
513 fn checkpoint(&mut self, compiler: &mut GraphCompiler, runtime: &Runtime) -> Result<()> {
514 let data = if let Some(data) = self.attached_value() {
515 Arc::clone(data)
516 } else {
517 let program = compiler.compile(self)?;
518 Arc::new(RetainedValue::from_tensor(single_runtime_output(
519 runtime.run_compiled(&program, &[])?,
520 "TracedTensorAdExt::checkpoint",
521 )?))
522 };
523 checkpoint_tensor(self, data)?;
524 Ok(())
525 }
526
527 fn jvp(&self, wrt: &TracedTensor, tangent: &TracedTensor) -> Result<TracedTensor> {
528 let wrt_input_key = leaf_input_key(wrt)?;
529 self.jvp_optional(wrt, tangent)?
530 .ok_or_else(|| inactive_wrt_error("jvp", &wrt_input_key))
531 }
532
533 fn jvp_optional(
534 &self,
535 wrt: &TracedTensor,
536 tangent: &TracedTensor,
537 ) -> Result<Option<TracedTensor>> {
538 let rules = SemanticExtensionRuleSet::default();
539 jvp_many_optional_impl(self, &[(wrt, tangent)], &rules, None)
540 }
541
542 fn vjp(&self, wrt: &TracedTensor, cotangent: &TracedTensor) -> Result<TracedTensor> {
543 let wrt_input_key = leaf_input_key(wrt)?;
544 self.vjp_optional(wrt, cotangent)?
545 .ok_or_else(|| inactive_wrt_error("vjp", &wrt_input_key))
546 }
547
548 fn vjp_optional(
549 &self,
550 wrt: &TracedTensor,
551 cotangent: &TracedTensor,
552 ) -> Result<Option<TracedTensor>> {
553 let rules = SemanticExtensionRuleSet::default();
554 vjp_many_with_transform_and_cache(self, &[wrt], cotangent, &rules, "vjp", None)
555 .map(|mut results| results.pop().flatten())
556 }
557}
558
559fn jvp_many_optional_impl(
560 output: &TracedTensor,
561 wrt_tangents: &[(&TracedTensor, &TracedTensor)],
562 rules: &SemanticExtensionRuleSet,
563 ad_transform_cache: Option<&AdTransformCache>,
564) -> Result<Option<TracedTensor>> {
565 if wrt_tangents.is_empty() {
566 return Ok(None);
567 }
568
569 let mut seen = HashSet::with_capacity(wrt_tangents.len());
570 let mut requested = Vec::with_capacity(wrt_tangents.len());
571 for (wrt, tangent) in wrt_tangents {
572 let key = leaf_input_key(wrt)?;
573 if !seen.insert(key.clone()) {
574 return Err(Error::invalid_argument(
575 "jvp",
576 ErrorPhase::GraphBuild,
577 "wrt_tangents",
578 "wrt_tangents contains duplicate wrt leaves",
579 ));
580 }
581 let tangent_data = tangent.attached_value().cloned().ok_or_else(|| {
582 Error::invalid_argument(
583 "jvp",
584 ErrorPhase::GraphBuild,
585 "tangent",
586 "jvp tangent must have concrete tensor data",
587 )
588 })?;
589 requested.push((key, tangent_data));
590 }
591
592 let mut compiler = GraphCompiler::new();
593 let source = compile_ad_source(&mut compiler, output)?;
594 let mut active_inputs = vec![false; source.input_count()];
595 let mut source_indices = Vec::with_capacity(requested.len());
596 for (key, _) in &requested {
597 let source_index = source.input_key_index(key);
598 if let Some(index) = source_index {
599 active_inputs[index] = true;
600 }
601 source_indices.push(source_index);
602 }
603 if !active_inputs.iter().any(|active| *active) {
604 return Ok(None);
605 }
606
607 let derivative = semantic_jvp_with_cache(
608 source.frozen_program(),
609 &active_inputs,
610 rules,
611 ad_transform_cache,
612 )?;
613 let Some(derivative_output_index) = derivative
614 .derivative_output_indices()
615 .first()
616 .copied()
617 .flatten()
618 else {
619 return Ok(None);
620 };
621
622 let mut seed_tensors = Vec::with_capacity(requested.len());
623 for ((_, tangent_data), source_index) in requested.iter().zip(source_indices) {
624 let Some(source_index) = source_index else {
625 continue;
626 };
627 let Some(seed_input_index) = derivative
628 .derivative_input_indices()
629 .get(source_index)
630 .copied()
631 .flatten()
632 else {
633 continue;
634 };
635 seed_tensors.push((seed_input_index, Arc::clone(tangent_data)));
636 }
637 let mut inherited_tensors = Vec::with_capacity(1 + 2 * wrt_tangents.len());
638 inherited_tensors.push(output);
639 for (wrt, tangent) in wrt_tangents {
640 inherited_tensors.push(*wrt);
641 inherited_tensors.push(*tangent);
642 }
643 let mut traces = derivative_tensors_from_program(
644 &source,
645 &derivative,
646 &[derivative_output_index],
647 &seed_tensors,
648 &inherited_tensors,
649 &[tensor_shape_hint(output)],
650 "jvp",
651 )?;
652 traces.pop().map(Some).ok_or_else(|| {
653 Error::runtime_state(
654 "jvp",
655 ErrorPhase::GraphBuild,
656 "derivative program returned no requested output",
657 )
658 })
659}
660
661fn vjp_many_optional_impl(
662 output: &TracedTensor,
663 wrts: &[&TracedTensor],
664 cotangent: &TracedTensor,
665 rules: &SemanticExtensionRuleSet,
666 transform: &'static str,
667 ad_transform_cache: Option<&AdTransformCache>,
668) -> Result<Vec<Option<TracedTensor>>> {
669 let cotangent_data = cotangent.attached_value().cloned().ok_or_else(|| {
670 Error::invalid_argument(
671 transform,
672 ErrorPhase::GraphBuild,
673 "cotangent",
674 "vjp cotangent must have concrete tensor data",
675 )
676 })?;
677 if wrts.is_empty() {
678 return Ok(Vec::new());
679 }
680
681 let requested_keys = wrts
682 .iter()
683 .map(|wrt| leaf_input_key(wrt))
684 .collect::<Result<Vec<_>>>()?;
685 let mut compiler = GraphCompiler::new();
686 let source = compile_ad_source(&mut compiler, output)?;
687 let source_indices: Vec<_> = requested_keys
688 .iter()
689 .map(|key| source.input_key_index(key))
690 .collect();
691 let mut active_inputs = vec![false; source.input_count()];
692 for source_index in source_indices.iter().flatten().copied() {
693 active_inputs[source_index] = true;
694 }
695 if !active_inputs.iter().any(|active| *active) {
696 return Ok(vec![None; wrts.len()]);
697 }
698
699 let active_outputs = vec![true; source.output_count()];
700 let derivative = semantic_vjp_with_cache(
701 source.frozen_program(),
702 &active_inputs,
703 &active_outputs,
704 rules,
705 ad_transform_cache,
706 )?;
707 let Some(seed_input_index) = derivative
708 .derivative_input_indices()
709 .first()
710 .copied()
711 .flatten()
712 else {
713 return Ok(vec![None; wrts.len()]);
714 };
715
716 let mut output_slots = vec![None; source.input_count()];
717 let mut derivative_output_indices = Vec::new();
718 let mut fallback_shape_hints = Vec::new();
719 for (request_index, source_index) in source_indices.iter().enumerate() {
720 let Some(source_index) = source_index else {
721 continue;
722 };
723 if output_slots[*source_index].is_some() {
724 continue;
725 }
726 let Some(derivative_output_index) = derivative
727 .derivative_output_indices()
728 .get(*source_index)
729 .copied()
730 .flatten()
731 else {
732 continue;
733 };
734 output_slots[*source_index] = Some(derivative_output_indices.len());
735 derivative_output_indices.push(derivative_output_index);
736 fallback_shape_hints.push(tensor_shape_hint(wrts[request_index]));
737 }
738 if derivative_output_indices.is_empty() {
739 return Ok(vec![None; wrts.len()]);
740 }
741
742 let inherited_tensors: Vec<&TracedTensor> = std::iter::once(output)
743 .chain(wrts.iter().copied())
744 .chain(std::iter::once(cotangent))
745 .collect();
746 let traces = derivative_tensors_from_program(
747 &source,
748 &derivative,
749 &derivative_output_indices,
750 &[(seed_input_index, cotangent_data)],
751 &inherited_tensors,
752 &fallback_shape_hints,
753 transform,
754 )?;
755 Ok(source_indices
756 .into_iter()
757 .map(|source_index| {
758 source_index
759 .and_then(|index| output_slots[index])
760 .map(|slot| traces[slot].clone())
761 })
762 .collect())
763}
764
765fn semantic_jvp_with_cache(
766 source: &FrozenProgram,
767 active_inputs: &[bool],
768 rules: &SemanticExtensionRuleSet,
769 ad_transform_cache: Option<&AdTransformCache>,
770) -> Result<SemanticAdProgram> {
771 let key = SemanticAdTransformCacheKey::jvp(source, active_inputs);
772 if let Some(cache) = ad_transform_cache
773 && let Some(cached) = cache.get_semantic(&key, source)?
774 {
775 return cached
776 .as_ref()
777 .with_input_prefix_bindings_from(source)
778 .map_err(|source| {
779 Error::runtime_state_source(
780 "semantic traced jvp cache",
781 ErrorPhase::GraphBuild,
782 source,
783 )
784 });
785 }
786 let derivative =
787 semantic_jvp(source, active_inputs, rules).map_err(semantic_transform_error("jvp"))?;
788 if let Some(cache) = ad_transform_cache {
789 cache.put_semantic(key, source, Arc::new(derivative.clone()))?;
790 }
791 Ok(derivative)
792}
793
794fn semantic_vjp_with_cache(
795 source: &FrozenProgram,
796 active_inputs: &[bool],
797 active_outputs: &[bool],
798 rules: &SemanticExtensionRuleSet,
799 ad_transform_cache: Option<&AdTransformCache>,
800) -> Result<SemanticAdProgram> {
801 let key = SemanticAdTransformCacheKey::vjp(source, active_inputs, active_outputs);
802 if let Some(cache) = ad_transform_cache
803 && let Some(cached) = cache.get_semantic(&key, source)?
804 {
805 return cached
806 .as_ref()
807 .with_input_prefix_bindings_from(source)
808 .map_err(|source| {
809 Error::runtime_state_source(
810 "semantic traced vjp cache",
811 ErrorPhase::GraphBuild,
812 source,
813 )
814 });
815 }
816 let derivative = semantic_vjp(source, active_inputs, active_outputs, rules)
817 .map_err(semantic_transform_error("vjp"))?;
818 if let Some(cache) = ad_transform_cache {
819 cache.put_semantic(key, source, Arc::new(derivative.clone()))?;
820 }
821 Ok(derivative)
822}
823
824fn semantic_transform_error(
825 transform: &'static str,
826) -> impl FnOnce(SemanticAdTransformError) -> Error {
827 move |source| {
828 semantic_transform_validation_error(transform, &source).unwrap_or_else(|| {
829 Error::runtime_state_source(transform, ErrorPhase::GraphBuild, source)
830 })
831 }
832}
833
834fn semantic_transform_validation_error(
835 transform: &'static str,
836 source: &SemanticAdTransformError,
837) -> Option<Error> {
838 if let SemanticAdTransformError::Extension(
839 SemanticAdError::Unsupported { family_id, .. }
840 | SemanticAdError::MissingRule { family_id, .. },
841 ) = source
842 {
843 return Some(Error::UnsupportedAdRule {
844 transform,
845 op: (*family_id).to_owned(),
846 });
847 }
848
849 let SemanticAdTransformError::Extension(SemanticAdError::Rule { source, .. }) = source else {
850 return None;
851 };
852 let tenferro_ops::ad::ADRuleError::InvalidInput { op, message, .. } =
853 source.downcast_ref::<tenferro_ops::ad::ADRuleError>()?
854 else {
855 return None;
856 };
857 Some(Error::invalid_argument(
858 transform,
859 ErrorPhase::GraphBuild,
860 "semantic_ad_rule",
861 format!("{op}: {message}"),
862 ))
863}
864
865fn derivative_tensors_from_program(
866 source: &CompiledGraph,
867 derivative: &SemanticAdProgram,
868 derivative_output_indices: &[usize],
869 seed_tensors: &[(usize, Arc<RetainedValue>)],
870 inherited_tensors: &[&TracedTensor],
871 fallback_shape_hints: &[Option<Vec<SymDim>>],
872 transform: &'static str,
873) -> Result<Vec<TracedTensor>> {
874 derivative_traces_from_frozen_program(
875 source,
876 derivative.frozen(),
877 derivative_output_indices,
878 seed_tensors,
879 inherited_tensors,
880 fallback_shape_hints,
881 transform,
882 )
883}
884
885pub(crate) fn derivative_trace_from_frozen_program(
886 source: &CompiledGraph,
887 frozen: &FrozenProgram,
888 derivative_output_index: usize,
889 seed_tensors: &[(usize, Arc<RetainedValue>)],
890 inherited_tensors: &[&TracedTensor],
891 fallback_shape_hint: Option<Vec<SymDim>>,
892 transform: &'static str,
893) -> Result<TracedTensor> {
894 let mut outputs = derivative_traces_from_frozen_program(
895 source,
896 frozen,
897 &[derivative_output_index],
898 seed_tensors,
899 inherited_tensors,
900 &[fallback_shape_hint],
901 transform,
902 )?;
903 outputs.pop().ok_or_else(|| {
904 Error::runtime_state(
905 transform,
906 ErrorPhase::GraphBuild,
907 "derivative program returned no requested output",
908 )
909 })
910}
911
912fn derivative_traces_from_frozen_program(
913 source: &CompiledGraph,
914 frozen: &FrozenProgram,
915 derivative_output_indices: &[usize],
916 seed_tensors: &[(usize, Arc<RetainedValue>)],
917 inherited_tensors: &[&TracedTensor],
918 fallback_shape_hints: &[Option<Vec<SymDim>>],
919 transform: &'static str,
920) -> Result<Vec<TracedTensor>> {
921 let input_shapes = symbolic_input_shapes(frozen)?;
922 let input_shape_refs: Vec<_> = input_shapes.iter().map(Vec::as_slice).collect();
923 let input_metas = frozen
924 .program
925 .inputs()
926 .iter()
927 .copied()
928 .map(|value| tensor_meta_for_value(frozen, value, &input_shape_refs, transform))
929 .collect::<Result<Vec<_>>>()?;
930
931 if derivative_output_indices.len() != fallback_shape_hints.len() {
932 return Err(Error::runtime_state(
933 transform,
934 ErrorPhase::GraphBuild,
935 "derivative output indices and fallback shape hints differ in length",
936 ));
937 }
938 let output_metas = derivative_output_indices
939 .iter()
940 .map(|&derivative_output_index| {
941 let output_value = *frozen
942 .program
943 .outputs()
944 .get(derivative_output_index)
945 .ok_or_else(|| {
946 Error::runtime_state(
947 transform,
948 ErrorPhase::GraphBuild,
949 format!(
950 "derivative output index {derivative_output_index} is outside {} outputs",
951 frozen.program.outputs().len()
952 ),
953 )
954 })?;
955 tensor_meta_for_value(frozen, output_value, &input_shape_refs, transform)
956 })
957 .collect::<Result<Vec<_>>>()?;
958
959 let mut builder = GraphBuilder::<StdTensorOp>::new();
960 let mut value_map = HashMap::<ProgramValue, LocalValueId>::new();
961 let mut input_keys = Vec::with_capacity(frozen.program.inputs().len());
962 for (input_index, input) in frozen.program.inputs().iter().copied().enumerate() {
963 let key = if input_index < source.input_keys().len() {
964 source.input_keys()[input_index].clone()
965 } else {
966 allocate_input_key()
967 };
968 let local = builder.add_input(key.clone());
969 value_map.insert(input, local);
970 input_keys.push(key);
971 }
972
973 for operation in frozen.program.operations() {
974 let inputs = operation
975 .inputs()
976 .iter()
977 .copied()
978 .map(|value| {
979 value_map
980 .get(&value)
981 .copied()
982 .map(ValueRef::Local)
983 .ok_or_else(|| missing_program_value(transform, "operation input"))
984 })
985 .collect::<Result<Vec<_>>>()?;
986 let op = match operation.op() {
987 SemanticOpRef::Core(op) => StdTensorOp::from(op),
988 SemanticOpRef::Extension(op) => StdTensorOp::Extension(op.clone_arc()),
989 _ => {
990 return Err(Error::runtime_state(
991 transform,
992 ErrorPhase::GraphBuild,
993 "unsupported semantic operation variant in derivative graph",
994 ));
995 }
996 };
997 let outputs = builder.add_operation(op, inputs, OperationRole::Primary);
998 if outputs.len() != operation.outputs().len() {
999 return Err(Error::runtime_state(
1000 transform,
1001 ErrorPhase::GraphBuild,
1002 format!(
1003 "semantic operation expected {} outputs, graph builder produced {}",
1004 operation.outputs().len(),
1005 outputs.len()
1006 ),
1007 ));
1008 }
1009 for (value, local) in operation.outputs().iter().copied().zip(outputs) {
1010 value_map.insert(value, local);
1011 }
1012 }
1013
1014 let graph_outputs = frozen
1015 .program
1016 .outputs()
1017 .iter()
1018 .copied()
1019 .map(|value| {
1020 value_map
1021 .get(&value)
1022 .copied()
1023 .ok_or_else(|| missing_program_value(transform, "program output"))
1024 })
1025 .collect::<Result<Vec<_>>>()?;
1026 let vals = derivative_output_indices
1027 .iter()
1028 .map(|&index| {
1029 graph_outputs.get(index).copied().ok_or_else(|| {
1030 Error::runtime_state(
1031 transform,
1032 ErrorPhase::GraphBuild,
1033 "derivative output index missing after graph conversion",
1034 )
1035 })
1036 })
1037 .collect::<Result<Vec<_>>>()?;
1038 builder.set_outputs(graph_outputs);
1039 let graph = Arc::new(builder.build());
1040
1041 let Some(primary_tensor) = inherited_tensors.first() else {
1042 return Err(Error::runtime_state(
1043 transform,
1044 ErrorPhase::GraphBuild,
1045 "derivative trace construction requires inherited source tensors",
1046 ));
1047 };
1048 let mut inputs_map = (*tensor_inputs_map(primary_tensor)).clone();
1049 for (input_index, key) in input_keys.iter().enumerate() {
1050 if let Some(tensor) = frozen_input_value(frozen, input_index) {
1051 inputs_map.insert(key.clone(), tensor);
1052 }
1053 }
1054 for (seed_input_index, tensor) in seed_tensors {
1055 let meta = input_metas.get(*seed_input_index).ok_or_else(|| {
1056 Error::runtime_state(
1057 transform,
1058 ErrorPhase::GraphBuild,
1059 format!("seed input index {seed_input_index} is outside derivative inputs"),
1060 )
1061 })?;
1062 validate_seed_tensor(transform, *seed_input_index, tensor.as_ref(), meta)?;
1063 let key = input_keys.get(*seed_input_index).ok_or_else(|| {
1064 Error::runtime_state(
1065 transform,
1066 ErrorPhase::GraphBuild,
1067 format!("seed input key {seed_input_index} is outside derivative inputs"),
1068 )
1069 })?;
1070 inputs_map.insert(key.clone(), Arc::clone(tensor));
1071 }
1072
1073 let source_input_count = source.input_keys().len();
1074 let graph_input_metadata = graph
1075 .inputs()
1076 .iter()
1077 .copied()
1078 .zip(input_metas.iter().cloned())
1079 .enumerate()
1080 .filter_map(|(input_index, (input, meta))| {
1081 if input_index < source_input_count {
1087 None
1088 } else {
1089 Some((graph.values()[input].key.clone(), meta))
1090 }
1091 });
1092 let analysis = register_scoped_graph_analysis(graph.as_ref(), graph_input_metadata)?;
1093 let inherited_constraint_scopes = inherited_tensors
1094 .iter()
1095 .map(|tensor| ConstraintScopeTransfer::from_tensor(tensor))
1096 .collect::<Vec<_>>();
1097
1098 let inputs_map = Arc::new(inputs_map);
1099 let leaf_metas = merge_traced_leaf_metas(inherited_tensors.iter().copied());
1103 let metadata_scopes = metadata_scopes_with_new(
1104 analysis.metadata,
1105 inherited_tensors
1106 .iter()
1107 .map(|tensor| tensor_metadata_scopes(tensor)),
1108 );
1109 let constraint_scope_transfer =
1110 ConstraintScopeTransfer::with_new(analysis.constraints, inherited_constraint_scopes.iter());
1111
1112 Ok(vals
1113 .into_iter()
1114 .zip(output_metas)
1115 .zip(fallback_shape_hints)
1116 .map(|((val, output_meta), fallback_shape_hint)| {
1117 tensor_from_parts(TracedTensorParts {
1118 rank: output_meta.rank(),
1119 dtype: output_meta.dtype,
1120 graph: Arc::clone(&graph),
1121 val,
1122 data: None,
1123 shape_hint: output_meta
1124 .exact_shape()
1125 .or_else(|| fallback_shape_hint.clone()),
1126 inputs_map: Arc::clone(&inputs_map),
1127 leaf_metas: Arc::clone(&leaf_metas),
1128 extra_roots: Vec::new(),
1129 checkpoint_chain: None,
1130 metadata_scopes: metadata_scopes.clone(),
1131 constraint_scope_transfer: constraint_scope_transfer.clone(),
1132 })
1133 })
1134 .collect())
1135}
1136
1137fn missing_program_value(transform: &'static str, role: &'static str) -> Error {
1138 Error::runtime_state(
1139 transform,
1140 ErrorPhase::GraphBuild,
1141 format!("semantic derivative graph references missing {role}"),
1142 )
1143}
1144
1145fn symbolic_input_shapes(frozen: &FrozenProgram) -> Result<Vec<Vec<SymDim>>> {
1146 frozen
1147 .program
1148 .inputs()
1149 .iter()
1150 .copied()
1151 .map(|value| {
1152 let meta = frozen.program.value_metadata(value).map_err(|source| {
1153 Error::runtime_state_source(
1154 "semantic traced AD input metadata",
1155 ErrorPhase::GraphBuild,
1156 source,
1157 )
1158 })?;
1159 let tensor_id = allocate_shape_tensor_id();
1160 Ok((0..meta.shape().len())
1161 .map(|axis| SymDim::tensor_axis(tensor_id, axis))
1162 .collect())
1163 })
1164 .collect()
1165}
1166
1167fn tensor_meta_for_value(
1168 frozen: &FrozenProgram,
1169 value: ProgramValue,
1170 input_shapes: &[&[SymDim]],
1171 transform: &'static str,
1172) -> Result<TensorMeta> {
1173 let meta = frozen
1174 .program
1175 .value_metadata(value)
1176 .map_err(|source| Error::runtime_state_source(transform, ErrorPhase::GraphBuild, source))?;
1177 Ok(program_metadata_to_tensor_meta(meta, input_shapes))
1178}
1179
1180fn program_metadata_to_tensor_meta(
1181 metadata: &ProgramValueMetadata,
1182 input_shapes: &[&[SymDim]],
1183) -> TensorMeta {
1184 let extents = metadata
1185 .shape()
1186 .iter()
1187 .cloned()
1188 .map(|extent| extent.map(|dim| SymDim::from_dim_expr(&dim, input_shapes)))
1189 .collect();
1190 let meta = TensorMeta::with_extents(metadata.dtype(), extents);
1193 match metadata.scalar_identity() {
1194 Some(identity) => meta.with_scalar_identity(identity),
1195 None => meta,
1196 }
1197}
1198
1199fn validate_seed_tensor(
1200 transform: &'static str,
1201 input_index: usize,
1202 tensor: &RetainedValue,
1203 expected: &TensorMeta,
1204) -> Result<()> {
1205 let actual_dtype = tensor.dtype();
1206 if actual_dtype != expected.dtype {
1207 return Err(Error::invalid_argument(
1208 transform,
1209 ErrorPhase::GraphBuild,
1210 "seed",
1211 format!(
1212 "seed input {input_index} dtype mismatch: expected {:?}, got {:?}",
1213 expected.dtype, actual_dtype
1214 ),
1215 ));
1216 }
1217 let actual_shape = tensor.shape();
1218 if actual_shape.len() != expected.rank() {
1219 return Err(Error::invalid_argument(
1220 transform,
1221 ErrorPhase::GraphBuild,
1222 "seed",
1223 format!(
1224 "seed input {input_index} rank mismatch: expected {}, got {}",
1225 expected.rank(),
1226 actual_shape.len()
1227 ),
1228 ));
1229 }
1230 if let Some(expected_shape) = expected
1231 .exact_shape()
1232 .filter(|shape| shape.iter().all(|dim| dim.constant_value().is_some()))
1233 .map(|shape| {
1234 shape
1235 .into_iter()
1236 .map(|dim| dim.constant_value().expect("filtered constant shape"))
1237 .collect::<Vec<_>>()
1238 })
1239 && expected_shape != actual_shape
1240 {
1241 return Err(Error::invalid_argument(
1242 transform,
1243 ErrorPhase::GraphBuild,
1244 "seed",
1245 format!(
1246 "seed input {input_index} shape mismatch: expected {:?}, got {:?}",
1247 expected_shape, actual_shape
1248 ),
1249 ));
1250 }
1251 Ok(())
1252}
1253
1254#[cfg(test)]
1255mod semantic_transform_error_tests {
1256 use super::*;
1257 use crate::semantic_extension::SemanticAdRuleRole;
1258
1259 #[test]
1260 fn unsupported_semantic_rule_maps_to_public_transform_error() {
1261 let source = SemanticAdTransformError::Extension(SemanticAdError::Unsupported {
1262 family_id: "tenferro-tests.unsupported.v1",
1263 role: SemanticAdRuleRole::LinearTranspose,
1264 message: "unsupported test payload".into(),
1265 });
1266
1267 let error = semantic_transform_validation_error("vjp", &source)
1268 .expect("semantic rejection must map to a public unsupported-rule error");
1269
1270 assert!(matches!(
1271 error,
1272 Error::UnsupportedAdRule { transform: "vjp", ref op }
1273 if op == "tenferro-tests.unsupported.v1"
1274 ));
1275 }
1276
1277 #[test]
1278 fn missing_semantic_rule_maps_to_public_transform_error() {
1279 let source = SemanticAdTransformError::Extension(SemanticAdError::MissingRule {
1280 family_id: "tenferro-tests.missing.v1",
1281 role: SemanticAdRuleRole::Linearize,
1282 });
1283
1284 let error = semantic_transform_validation_error("jvp", &source)
1285 .expect("missing semantic rule must map to a public unsupported-rule error");
1286
1287 assert!(matches!(
1288 error,
1289 Error::UnsupportedAdRule { transform: "jvp", ref op }
1290 if op == "tenferro-tests.missing.v1"
1291 ));
1292 }
1293}