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