1use std::collections::HashMap;
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_tensor, inputs_map as tensor_inputs_map, leaf_input_key,
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, 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_optional_impl(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 grad_optional_with_rules_and_cache(
59 output: &TracedTensor,
60 wrt: &TracedTensor,
61 rules: &SemanticExtensionRuleSet,
62 ad_transform_cache: Option<&AdTransformCache>,
63) -> Result<Option<TracedTensor>> {
64 if output.rank != 0 {
65 return Err(Error::NonScalarGrad {
66 shape: error_shape_hint(output),
67 });
68 }
69
70 let ones = ones_tensor(output.dtype, vec![])?;
71 let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
72 vjp_optional_impl(output, wrt, &seed, rules, "grad", ad_transform_cache)
73}
74
75pub(crate) fn jvp_optional_with_rules_and_cache(
76 output: &TracedTensor,
77 wrt: &TracedTensor,
78 tangent: &TracedTensor,
79 rules: &SemanticExtensionRuleSet,
80 ad_transform_cache: Option<&AdTransformCache>,
81) -> Result<Option<TracedTensor>> {
82 jvp_optional_impl(output, wrt, tangent, rules, ad_transform_cache)
83}
84
85pub(crate) fn vjp_with_rules_and_cache(
86 output: &TracedTensor,
87 wrt: &TracedTensor,
88 cotangent: &TracedTensor,
89 rules: &SemanticExtensionRuleSet,
90 ad_transform_cache: Option<&AdTransformCache>,
91) -> Result<TracedTensor> {
92 let wrt_input_key = leaf_input_key(wrt)?;
93 vjp_optional_impl(output, wrt, cotangent, rules, "vjp", ad_transform_cache)?
94 .ok_or_else(|| Error::Internal(format!("vjp output is inactive for {:?}", wrt_input_key)))
95}
96
97pub(crate) fn vjp_optional_with_rules_and_cache(
98 output: &TracedTensor,
99 wrt: &TracedTensor,
100 cotangent: &TracedTensor,
101 rules: &SemanticExtensionRuleSet,
102 ad_transform_cache: Option<&AdTransformCache>,
103) -> Result<Option<TracedTensor>> {
104 vjp_optional_impl(output, wrt, cotangent, rules, "vjp", ad_transform_cache)
105}
106
107fn grad_with_optional_rules(
108 output: &TracedTensor,
109 wrt: &TracedTensor,
110 rules: &SemanticExtensionRuleSet,
111 ad_transform_cache: Option<&AdTransformCache>,
112) -> Result<TracedTensor> {
113 if output.rank != 0 {
114 return Err(Error::NonScalarGrad {
115 shape: error_shape_hint(output),
116 });
117 }
118
119 let ones = ones_tensor(output.dtype, vec![])?;
120 let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
121 let wrt_input_key = leaf_input_key(wrt)?;
122 vjp_optional_impl(output, wrt, &seed, rules, "grad", ad_transform_cache)?
123 .ok_or_else(|| Error::Internal(format!("grad output is inactive for {:?}", wrt_input_key)))
124}
125
126fn single_runtime_output(mut outputs: Vec<Tensor>, op: &'static str) -> Result<Tensor> {
127 let actual = outputs.len();
128 if actual != 1 {
129 return Err(Error::runtime_state(
130 op,
131 ErrorPhase::Execution,
132 format!("expected one runtime output, got {actual}"),
133 ));
134 }
135 outputs.pop().ok_or_else(|| {
136 Error::runtime_state(
137 op,
138 ErrorPhase::Execution,
139 "runtime returned no output after successful output-count validation",
140 )
141 })
142}
143
144pub trait TracedTensorAdExt {
158 fn grad(&self, wrt: &TracedTensor) -> Result<TracedTensor>;
204
205 fn grad_optional(&self, wrt: &TracedTensor) -> Result<Option<TracedTensor>>;
233
234 fn checkpoint(&mut self, compiler: &mut GraphCompiler, runtime: &Runtime) -> Result<()>;
266
267 fn jvp(&self, wrt: &TracedTensor, tangent: &TracedTensor) -> Result<TracedTensor>;
309
310 fn jvp_optional(
339 &self,
340 wrt: &TracedTensor,
341 tangent: &TracedTensor,
342 ) -> Result<Option<TracedTensor>>;
343
344 fn vjp(&self, wrt: &TracedTensor, cotangent: &TracedTensor) -> Result<TracedTensor>;
391
392 fn vjp_optional(
421 &self,
422 wrt: &TracedTensor,
423 cotangent: &TracedTensor,
424 ) -> Result<Option<TracedTensor>>;
425}
426
427impl TracedTensorAdExt for TracedTensor {
428 fn grad(&self, wrt: &TracedTensor) -> Result<TracedTensor> {
429 let rules = SemanticExtensionRuleSet::default();
430 grad_with_optional_rules(self, wrt, &rules, None)
431 }
432
433 fn grad_optional(&self, wrt: &TracedTensor) -> Result<Option<TracedTensor>> {
434 if self.rank != 0 {
435 return Err(Error::NonScalarGrad {
436 shape: error_shape_hint(self),
437 });
438 }
439
440 let ones = ones_tensor(self.dtype, vec![])?;
441 let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
442 let rules = SemanticExtensionRuleSet::default();
443 vjp_optional_impl(self, wrt, &seed, &rules, "grad", None)
444 }
445
446 fn checkpoint(&mut self, compiler: &mut GraphCompiler, runtime: &Runtime) -> Result<()> {
447 let data = if let Some(data) = self.attached_data() {
448 Arc::clone(data)
449 } else {
450 let program = compiler.compile(self)?;
451 Arc::new(single_runtime_output(
452 runtime.run_compiled(&program, &[])?,
453 "TracedTensorAdExt::checkpoint",
454 )?)
455 };
456 checkpoint_tensor(self, data)?;
457 Ok(())
458 }
459
460 fn jvp(&self, wrt: &TracedTensor, tangent: &TracedTensor) -> Result<TracedTensor> {
461 let wrt_input_key = leaf_input_key(wrt)?;
462 self.jvp_optional(wrt, tangent)?.ok_or_else(|| {
463 Error::Internal(format!("jvp output is inactive for {:?}", wrt_input_key))
464 })
465 }
466
467 fn jvp_optional(
468 &self,
469 wrt: &TracedTensor,
470 tangent: &TracedTensor,
471 ) -> Result<Option<TracedTensor>> {
472 let rules = SemanticExtensionRuleSet::default();
473 jvp_optional_impl(self, wrt, tangent, &rules, None)
474 }
475
476 fn vjp(&self, wrt: &TracedTensor, cotangent: &TracedTensor) -> Result<TracedTensor> {
477 let wrt_input_key = leaf_input_key(wrt)?;
478 self.vjp_optional(wrt, cotangent)?.ok_or_else(|| {
479 Error::Internal(format!("vjp output is inactive for {:?}", wrt_input_key))
480 })
481 }
482
483 fn vjp_optional(
484 &self,
485 wrt: &TracedTensor,
486 cotangent: &TracedTensor,
487 ) -> Result<Option<TracedTensor>> {
488 let rules = SemanticExtensionRuleSet::default();
489 vjp_optional_impl(self, wrt, cotangent, &rules, "vjp", None)
490 }
491}
492
493fn jvp_optional_impl(
494 output: &TracedTensor,
495 wrt: &TracedTensor,
496 tangent: &TracedTensor,
497 rules: &SemanticExtensionRuleSet,
498 ad_transform_cache: Option<&AdTransformCache>,
499) -> Result<Option<TracedTensor>> {
500 let wrt_input_key = leaf_input_key(wrt)?;
501 let tangent_data = tangent.attached_data().cloned().ok_or_else(|| {
502 Error::invalid_argument(
503 "jvp",
504 ErrorPhase::GraphBuild,
505 "tangent",
506 "jvp tangent must have concrete tensor data",
507 )
508 })?;
509 let mut compiler = GraphCompiler::new();
510 let source = compile_ad_source(&mut compiler, output)?;
511 let Some(wrt_input_index) = source.input_key_index(&wrt_input_key) else {
512 return Ok(None);
513 };
514
515 let mut active_inputs = vec![false; source.input_count()];
516 active_inputs[wrt_input_index] = true;
517 let derivative = semantic_jvp_with_cache(
518 source.frozen_program(),
519 &active_inputs,
520 rules,
521 ad_transform_cache,
522 )?;
523 let Some(seed_input_index) = derivative
524 .derivative_input_indices()
525 .get(wrt_input_index)
526 .copied()
527 .flatten()
528 else {
529 return Ok(None);
530 };
531 let Some(derivative_output_index) = derivative
532 .derivative_output_indices()
533 .first()
534 .copied()
535 .flatten()
536 else {
537 return Ok(None);
538 };
539
540 derivative_tensor_from_program(
541 &source,
542 &derivative,
543 derivative_output_index,
544 &[(seed_input_index, tangent_data)],
545 [output, wrt, tangent],
546 tensor_shape_hint(output),
547 "jvp",
548 )
549 .map(Some)
550}
551
552fn vjp_optional_impl(
553 output: &TracedTensor,
554 wrt: &TracedTensor,
555 cotangent: &TracedTensor,
556 rules: &SemanticExtensionRuleSet,
557 transform: &'static str,
558 ad_transform_cache: Option<&AdTransformCache>,
559) -> Result<Option<TracedTensor>> {
560 let wrt_input_key = leaf_input_key(wrt)?;
561 let cotangent_data = cotangent.attached_data().cloned().ok_or_else(|| {
562 Error::invalid_argument(
563 transform,
564 ErrorPhase::GraphBuild,
565 "cotangent",
566 "vjp cotangent must have concrete tensor data",
567 )
568 })?;
569 let mut compiler = GraphCompiler::new();
570 let source = compile_ad_source(&mut compiler, output)?;
571 let Some(wrt_input_index) = source.input_key_index(&wrt_input_key) else {
572 return Ok(None);
573 };
574
575 let mut active_inputs = vec![false; source.input_count()];
576 active_inputs[wrt_input_index] = true;
577 let active_outputs = vec![true; source.output_count()];
578 let derivative = semantic_vjp_with_cache(
579 source.frozen_program(),
580 &active_inputs,
581 &active_outputs,
582 rules,
583 ad_transform_cache,
584 )?;
585 let Some(seed_input_index) = derivative
586 .derivative_input_indices()
587 .first()
588 .copied()
589 .flatten()
590 else {
591 return Ok(None);
592 };
593 let Some(derivative_output_index) = derivative
594 .derivative_output_indices()
595 .get(wrt_input_index)
596 .copied()
597 .flatten()
598 else {
599 return Ok(None);
600 };
601
602 derivative_tensor_from_program(
603 &source,
604 &derivative,
605 derivative_output_index,
606 &[(seed_input_index, cotangent_data)],
607 [output, wrt, cotangent],
608 tensor_shape_hint(wrt),
609 transform,
610 )
611 .map(Some)
612}
613
614fn semantic_jvp_with_cache(
615 source: &FrozenProgram,
616 active_inputs: &[bool],
617 rules: &SemanticExtensionRuleSet,
618 ad_transform_cache: Option<&AdTransformCache>,
619) -> Result<SemanticAdProgram> {
620 let key = SemanticAdTransformCacheKey::jvp(source, active_inputs);
621 if let Some(cache) = ad_transform_cache {
622 if let Some(cached) = cache.get_semantic(&key, source)? {
623 return cached
624 .as_ref()
625 .with_input_prefix_bindings_from(source)
626 .map_err(|source| {
627 Error::runtime_state_source(
628 "semantic traced jvp cache",
629 ErrorPhase::GraphBuild,
630 source,
631 )
632 });
633 }
634 }
635 let derivative =
636 semantic_jvp(source, active_inputs, rules).map_err(semantic_transform_error("jvp"))?;
637 if let Some(cache) = ad_transform_cache {
638 cache.put_semantic(key, source, Arc::new(derivative.clone()))?;
639 }
640 Ok(derivative)
641}
642
643fn semantic_vjp_with_cache(
644 source: &FrozenProgram,
645 active_inputs: &[bool],
646 active_outputs: &[bool],
647 rules: &SemanticExtensionRuleSet,
648 ad_transform_cache: Option<&AdTransformCache>,
649) -> Result<SemanticAdProgram> {
650 let key = SemanticAdTransformCacheKey::vjp(source, active_inputs, active_outputs);
651 if let Some(cache) = ad_transform_cache {
652 if let Some(cached) = cache.get_semantic(&key, source)? {
653 return cached
654 .as_ref()
655 .with_input_prefix_bindings_from(source)
656 .map_err(|source| {
657 Error::runtime_state_source(
658 "semantic traced vjp cache",
659 ErrorPhase::GraphBuild,
660 source,
661 )
662 });
663 }
664 }
665 let derivative = semantic_vjp(source, active_inputs, active_outputs, rules)
666 .map_err(semantic_transform_error("vjp"))?;
667 if let Some(cache) = ad_transform_cache {
668 cache.put_semantic(key, source, Arc::new(derivative.clone()))?;
669 }
670 Ok(derivative)
671}
672
673fn semantic_transform_error(
674 transform: &'static str,
675) -> impl FnOnce(SemanticAdTransformError) -> Error {
676 move |source| {
677 semantic_transform_validation_error(transform, &source).unwrap_or_else(|| {
678 Error::runtime_state_source(transform, ErrorPhase::GraphBuild, source)
679 })
680 }
681}
682
683fn semantic_transform_validation_error(
684 transform: &'static str,
685 source: &SemanticAdTransformError,
686) -> Option<Error> {
687 if let SemanticAdTransformError::Extension(
688 SemanticAdError::Unsupported { family_id, .. }
689 | SemanticAdError::MissingRule { family_id, .. },
690 ) = source
691 {
692 return Some(Error::UnsupportedAdRule {
693 transform,
694 op: (*family_id).to_owned(),
695 });
696 }
697
698 let SemanticAdTransformError::Extension(SemanticAdError::Rule { source, .. }) = source else {
699 return None;
700 };
701 let tenferro_ops::ad::ADRuleError::InvalidInput { op, message, .. } =
702 source.downcast_ref::<tenferro_ops::ad::ADRuleError>()?
703 else {
704 return None;
705 };
706 Some(Error::invalid_argument(
707 transform,
708 ErrorPhase::GraphBuild,
709 "semantic_ad_rule",
710 format!("{op}: {message}"),
711 ))
712}
713
714fn derivative_tensor_from_program(
715 source: &CompiledGraph,
716 derivative: &SemanticAdProgram,
717 derivative_output_index: usize,
718 seed_tensors: &[(usize, Arc<Tensor>)],
719 inherited_tensors: [&TracedTensor; 3],
720 fallback_shape_hint: Option<Vec<SymDim>>,
721 transform: &'static str,
722) -> Result<TracedTensor> {
723 derivative_trace_from_frozen_program(
724 source,
725 derivative.frozen(),
726 derivative_output_index,
727 seed_tensors,
728 &inherited_tensors,
729 fallback_shape_hint,
730 transform,
731 )
732}
733
734pub(crate) fn derivative_trace_from_frozen_program(
735 source: &CompiledGraph,
736 frozen: &FrozenProgram,
737 derivative_output_index: usize,
738 seed_tensors: &[(usize, Arc<Tensor>)],
739 inherited_tensors: &[&TracedTensor],
740 fallback_shape_hint: Option<Vec<SymDim>>,
741 transform: &'static str,
742) -> Result<TracedTensor> {
743 let input_shapes = symbolic_input_shapes(frozen)?;
744 let input_shape_refs: Vec<_> = input_shapes.iter().map(Vec::as_slice).collect();
745 let input_metas = frozen
746 .program
747 .inputs()
748 .iter()
749 .copied()
750 .map(|value| tensor_meta_for_value(frozen, value, &input_shape_refs, transform))
751 .collect::<Result<Vec<_>>>()?;
752
753 let output_value = *frozen
754 .program
755 .outputs()
756 .get(derivative_output_index)
757 .ok_or_else(|| {
758 Error::runtime_state(
759 transform,
760 ErrorPhase::GraphBuild,
761 format!(
762 "derivative output index {derivative_output_index} is outside {} outputs",
763 frozen.program.outputs().len()
764 ),
765 )
766 })?;
767 let output_meta = tensor_meta_for_value(frozen, output_value, &input_shape_refs, transform)?;
768
769 let mut builder = GraphBuilder::<StdTensorOp>::new();
770 let mut value_map = HashMap::<ProgramValue, LocalValueId>::new();
771 let mut input_keys = Vec::with_capacity(frozen.program.inputs().len());
772 for (input_index, input) in frozen.program.inputs().iter().copied().enumerate() {
773 let key = if input_index < source.input_keys().len() {
774 source.input_keys()[input_index].clone()
775 } else {
776 allocate_input_key()
777 };
778 let local = builder.add_input(key.clone());
779 value_map.insert(input, local);
780 input_keys.push(key);
781 }
782
783 for operation in frozen.program.operations() {
784 let inputs = operation
785 .inputs()
786 .iter()
787 .copied()
788 .map(|value| {
789 value_map
790 .get(&value)
791 .copied()
792 .map(ValueRef::Local)
793 .ok_or_else(|| missing_program_value(transform, "operation input"))
794 })
795 .collect::<Result<Vec<_>>>()?;
796 let op = match operation.op() {
797 SemanticOpRef::Core(op) => StdTensorOp::from(op),
798 SemanticOpRef::Extension(op) => StdTensorOp::Extension(op.clone_arc()),
799 _ => {
800 return Err(Error::runtime_state(
801 transform,
802 ErrorPhase::GraphBuild,
803 "unsupported semantic operation variant in derivative graph",
804 ));
805 }
806 };
807 let outputs = builder.add_operation(op, inputs, OperationRole::Primary);
808 if outputs.len() != operation.outputs().len() {
809 return Err(Error::runtime_state(
810 transform,
811 ErrorPhase::GraphBuild,
812 format!(
813 "semantic operation expected {} outputs, graph builder produced {}",
814 operation.outputs().len(),
815 outputs.len()
816 ),
817 ));
818 }
819 for (value, local) in operation.outputs().iter().copied().zip(outputs) {
820 value_map.insert(value, local);
821 }
822 }
823
824 let graph_outputs = frozen
825 .program
826 .outputs()
827 .iter()
828 .copied()
829 .map(|value| {
830 value_map
831 .get(&value)
832 .copied()
833 .ok_or_else(|| missing_program_value(transform, "program output"))
834 })
835 .collect::<Result<Vec<_>>>()?;
836 let val = *graph_outputs.get(derivative_output_index).ok_or_else(|| {
837 Error::runtime_state(
838 transform,
839 ErrorPhase::GraphBuild,
840 "derivative output index missing after graph conversion",
841 )
842 })?;
843 builder.set_outputs(graph_outputs);
844 let graph = Arc::new(builder.build());
845
846 let Some(primary_tensor) = inherited_tensors.first() else {
847 return Err(Error::runtime_state(
848 transform,
849 ErrorPhase::GraphBuild,
850 "derivative trace construction requires inherited source tensors",
851 ));
852 };
853 let mut inputs_map = (*tensor_inputs_map(primary_tensor)).clone();
854 for (input_index, key) in input_keys.iter().enumerate() {
855 if let Some(tensor) = frozen_input_tensor(frozen, input_index) {
856 inputs_map.insert(key.clone(), tensor);
857 }
858 }
859 for (seed_input_index, tensor) in seed_tensors {
860 let meta = input_metas.get(*seed_input_index).ok_or_else(|| {
861 Error::runtime_state(
862 transform,
863 ErrorPhase::GraphBuild,
864 format!("seed input index {seed_input_index} is outside derivative inputs"),
865 )
866 })?;
867 validate_seed_tensor(transform, *seed_input_index, tensor.as_ref(), meta)?;
868 let key = input_keys.get(*seed_input_index).ok_or_else(|| {
869 Error::runtime_state(
870 transform,
871 ErrorPhase::GraphBuild,
872 format!("seed input key {seed_input_index} is outside derivative inputs"),
873 )
874 })?;
875 inputs_map.insert(key.clone(), Arc::clone(tensor));
876 }
877
878 let source_input_count = source.input_keys().len();
879 let graph_input_metadata = graph
880 .inputs()
881 .iter()
882 .copied()
883 .zip(input_metas.iter().cloned())
884 .enumerate()
885 .filter_map(|(input_index, (input, meta))| {
886 if input_index < source_input_count {
892 None
893 } else {
894 Some((graph.values()[input].key.clone(), meta))
895 }
896 });
897 let analysis = register_scoped_graph_analysis(graph.as_ref(), graph_input_metadata)?;
898 let inherited_constraint_scopes = inherited_tensors
899 .iter()
900 .map(|tensor| ConstraintScopeTransfer::from_tensor(tensor))
901 .collect::<Vec<_>>();
902
903 Ok(tensor_from_parts(TracedTensorParts {
904 rank: output_meta.rank(),
905 dtype: output_meta.dtype,
906 graph,
907 val,
908 data: None,
909 shape_hint: output_meta.exact_shape().or(fallback_shape_hint),
910 inputs_map: Arc::new(inputs_map),
911 extra_roots: Vec::new(),
912 checkpoint_chain: None,
913 metadata_scopes: metadata_scopes_with_new(
914 analysis.metadata,
915 inherited_tensors
916 .iter()
917 .map(|tensor| tensor_metadata_scopes(tensor)),
918 ),
919 constraint_scope_transfer: ConstraintScopeTransfer::with_new(
920 analysis.constraints,
921 inherited_constraint_scopes.iter(),
922 ),
923 }))
924}
925
926fn missing_program_value(transform: &'static str, role: &'static str) -> Error {
927 Error::runtime_state(
928 transform,
929 ErrorPhase::GraphBuild,
930 format!("semantic derivative graph references missing {role}"),
931 )
932}
933
934fn symbolic_input_shapes(frozen: &FrozenProgram) -> Result<Vec<Vec<SymDim>>> {
935 frozen
936 .program
937 .inputs()
938 .iter()
939 .copied()
940 .map(|value| {
941 let meta = frozen.program.value_metadata(value).map_err(|source| {
942 Error::runtime_state_source(
943 "semantic traced AD input metadata",
944 ErrorPhase::GraphBuild,
945 source,
946 )
947 })?;
948 let tensor_id = allocate_shape_tensor_id();
949 Ok((0..meta.shape().len())
950 .map(|axis| SymDim::tensor_axis(tensor_id, axis))
951 .collect())
952 })
953 .collect()
954}
955
956fn tensor_meta_for_value(
957 frozen: &FrozenProgram,
958 value: ProgramValue,
959 input_shapes: &[&[SymDim]],
960 transform: &'static str,
961) -> Result<TensorMeta> {
962 let meta = frozen
963 .program
964 .value_metadata(value)
965 .map_err(|source| Error::runtime_state_source(transform, ErrorPhase::GraphBuild, source))?;
966 Ok(program_metadata_to_tensor_meta(meta, input_shapes))
967}
968
969fn program_metadata_to_tensor_meta(
970 metadata: &ProgramValueMetadata,
971 input_shapes: &[&[SymDim]],
972) -> TensorMeta {
973 let extents = metadata
974 .shape()
975 .iter()
976 .cloned()
977 .map(|extent| extent.map(|dim| SymDim::from_dim_expr(&dim, input_shapes)))
978 .collect();
979 TensorMeta::with_extents(metadata.dtype(), extents)
980}
981
982fn validate_seed_tensor(
983 transform: &'static str,
984 input_index: usize,
985 tensor: &Tensor,
986 expected: &TensorMeta,
987) -> Result<()> {
988 let actual_dtype = tensor.dtype();
989 if actual_dtype != expected.dtype {
990 return Err(Error::invalid_argument(
991 transform,
992 ErrorPhase::GraphBuild,
993 "seed",
994 format!(
995 "seed input {input_index} dtype mismatch: expected {:?}, got {:?}",
996 expected.dtype, actual_dtype
997 ),
998 ));
999 }
1000 let actual_shape = tensor.shape();
1001 if actual_shape.len() != expected.rank() {
1002 return Err(Error::invalid_argument(
1003 transform,
1004 ErrorPhase::GraphBuild,
1005 "seed",
1006 format!(
1007 "seed input {input_index} rank mismatch: expected {}, got {}",
1008 expected.rank(),
1009 actual_shape.len()
1010 ),
1011 ));
1012 }
1013 if let Some(expected_shape) = expected
1014 .exact_shape()
1015 .filter(|shape| shape.iter().all(|dim| dim.constant_value().is_some()))
1016 .map(|shape| {
1017 shape
1018 .into_iter()
1019 .map(|dim| dim.constant_value().expect("filtered constant shape"))
1020 .collect::<Vec<_>>()
1021 })
1022 {
1023 if expected_shape != actual_shape {
1024 return Err(Error::invalid_argument(
1025 transform,
1026 ErrorPhase::GraphBuild,
1027 "seed",
1028 format!(
1029 "seed input {input_index} shape mismatch: expected {:?}, got {:?}",
1030 expected_shape, actual_shape
1031 ),
1032 ));
1033 }
1034 }
1035 Ok(())
1036}
1037
1038#[cfg(test)]
1039mod semantic_transform_error_tests {
1040 use super::*;
1041 use crate::semantic_extension::SemanticAdRuleRole;
1042
1043 #[test]
1044 fn unsupported_semantic_rule_maps_to_public_transform_error() {
1045 let source = SemanticAdTransformError::Extension(SemanticAdError::Unsupported {
1046 family_id: "tenferro-tests.unsupported.v1",
1047 role: SemanticAdRuleRole::LinearTranspose,
1048 message: "unsupported test payload".into(),
1049 });
1050
1051 let error = semantic_transform_validation_error("vjp", &source)
1052 .expect("semantic rejection must map to a public unsupported-rule error");
1053
1054 assert!(matches!(
1055 error,
1056 Error::UnsupportedAdRule { transform: "vjp", ref op }
1057 if op == "tenferro-tests.unsupported.v1"
1058 ));
1059 }
1060
1061 #[test]
1062 fn missing_semantic_rule_maps_to_public_transform_error() {
1063 let source = SemanticAdTransformError::Extension(SemanticAdError::MissingRule {
1064 family_id: "tenferro-tests.missing.v1",
1065 role: SemanticAdRuleRole::Linearize,
1066 });
1067
1068 let error = semantic_transform_validation_error("jvp", &source)
1069 .expect("missing semantic rule must map to a public unsupported-rule error");
1070
1071 assert!(matches!(
1072 error,
1073 Error::UnsupportedAdRule { transform: "jvp", ref op }
1074 if op == "tenferro-tests.missing.v1"
1075 ));
1076 }
1077}