Skip to main content

tenferro_ad/
traced.rs

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
144/// Automatic differentiation helpers for [`TracedTensor`].
145///
146/// # Examples
147///
148/// ```rust
149/// use tenferro_ad::TracedTensorAdExt;
150/// use tenferro_runtime::TracedTensor;
151///
152/// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
153/// let loss = x.scale_real(2.0).unwrap();
154/// let maybe_dx = loss.grad_optional(&x).unwrap();
155/// assert!(maybe_dx.is_some());
156/// ```
157pub trait TracedTensorAdExt {
158    /// Gradient of a scalar output with respect to a traced input.
159    ///
160    /// For complex scalar outputs, tenferro returns the Hermitian-adjoint
161    /// cotangent. To compare seed-`1` scalar gradients with JAX's public
162    /// `grad` values, use the complex conjugate of this result. See
163    /// <https://tensor4all.org/tenferro-rs/guides/complex-ad.html>.
164    ///
165    /// # Examples
166    ///
167    /// ```rust
168    /// use tenferro_ad::TracedTensorAdExt;
169    /// use tenferro_cpu::CpuBackend;
170    /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
171    ///
172    /// fn eval(tensor: &TracedTensor) -> tenferro_runtime::Tensor {
173    ///     let mut compiler = GraphCompiler::new();
174    ///     let program = compiler.compile(tensor).unwrap();
175    ///     let backend = CpuBackend::new();
176    ///     let mut builder = Runtime::builder();
177    ///     builder
178    ///         .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
179    ///         .unwrap();
180    ///     let runtime = builder.build().unwrap();
181    ///     runtime.run_compiled(&program, &[]).unwrap().pop().unwrap()
182    /// }
183    ///
184    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
185    /// let loss = (&x * &x).unwrap();
186    /// let dx = loss.grad(&x).unwrap();
187    ///
188    /// assert_eq!(eval(&dx).as_slice::<f64>().unwrap(), &[6.0]);
189    /// ```
190    ///
191    /// # Errors
192    ///
193    /// Returns [`tenferro_runtime::Error::NonScalarGrad`] for a non-scalar
194    /// output, [`tenferro_runtime::Error::UnsupportedAdRule`] when an AD rule
195    /// is unavailable, or a typed validation/backend/runtime-state error.
196    ///
197    /// # Deferred errors
198    ///
199    /// Symbolic shape constraints can later produce
200    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
201    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
202    /// or execution.
203    fn grad(&self, wrt: &TracedTensor) -> Result<TracedTensor>;
204
205    /// Like [`grad`](Self::grad), but returns `None` when `wrt` is inactive.
206    ///
207    /// # Examples
208    ///
209    /// ```rust
210    /// use tenferro_ad::TracedTensorAdExt;
211    /// use tenferro_runtime::TracedTensor;
212    ///
213    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
214    /// let y = TracedTensor::from_vec_col_major(vec![], vec![4.0_f64]).unwrap();
215    /// let loss = (&y * &y).unwrap();
216    ///
217    /// assert!(loss.grad_optional(&x).unwrap().is_none());
218    /// ```
219    ///
220    /// # Errors
221    ///
222    /// Returns [`tenferro_runtime::Error::NonScalarGrad`] for a non-scalar
223    /// output, [`Error::UnsupportedAdRule`] when an AD rule is unavailable, or
224    /// a typed validation/backend/runtime-state error.
225    ///
226    /// # Deferred errors
227    ///
228    /// Symbolic shape constraints can later produce
229    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
230    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
231    /// or execution.
232    fn grad_optional(&self, wrt: &TracedTensor) -> Result<Option<TracedTensor>>;
233
234    /// Evaluate this tensor and replace its graph with a concrete leaf while
235    /// preserving the previous graph for AD replay.
236    ///
237    /// # Examples
238    ///
239    /// ```rust
240    /// use tenferro_ad::TracedTensorAdExt;
241    /// use tenferro_cpu::CpuBackend;
242    /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
243    ///
244    /// let mut compiler = GraphCompiler::new();
245    /// let backend = CpuBackend::new();
246    /// let mut builder = Runtime::builder();
247    /// builder
248    ///     .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
249    ///     .unwrap();
250    /// let runtime = builder.build().unwrap();
251    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
252    /// let mut y = (&x * &x).unwrap();
253    ///
254    /// y.checkpoint(&mut compiler, &runtime).unwrap();
255    ///
256    /// let value = y.attached_data().unwrap();
257    /// assert_eq!(value.as_slice::<f64>().unwrap(), &[9.0]);
258    /// ```
259    ///
260    /// # Errors
261    ///
262    /// Returns [`tenferro_runtime::Error::Validation`] when checkpoint metadata
263    /// is invalid, [`Error::RuntimeState`] when graph metadata or runtime
264    /// state is unavailable, or a typed backend error from evaluation.
265    fn checkpoint(&mut self, compiler: &mut GraphCompiler, runtime: &Runtime) -> Result<()>;
266
267    /// Forward-mode Jacobian-vector product.
268    ///
269    /// # Examples
270    ///
271    /// ```rust
272    /// use tenferro_ad::TracedTensorAdExt;
273    /// use tenferro_cpu::CpuBackend;
274    /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
275    ///
276    /// fn eval(tensor: &TracedTensor) -> tenferro_runtime::Tensor {
277    ///     let mut compiler = GraphCompiler::new();
278    ///     let program = compiler.compile(tensor).unwrap();
279    ///     let backend = CpuBackend::new();
280    ///     let mut builder = Runtime::builder();
281    ///     builder
282    ///         .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
283    ///         .unwrap();
284    ///     let runtime = builder.build().unwrap();
285    ///     runtime.run_compiled(&program, &[]).unwrap().pop().unwrap()
286    /// }
287    ///
288    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
289    /// let tangent = TracedTensor::from_vec_col_major(vec![], vec![2.0_f64]).unwrap();
290    /// let y = (&x * &x).unwrap();
291    /// let dy = y.jvp(&x, &tangent).unwrap();
292    ///
293    /// assert_eq!(eval(&dy).as_slice::<f64>().unwrap(), &[12.0]);
294    /// ```
295    ///
296    /// # Errors
297    ///
298    /// Returns [`tenferro_runtime::Error::UnsupportedAdRule`] when a JVP rule
299    /// is unavailable, [`Error::Validation`] for incompatible tangent metadata,
300    /// or a typed backend/runtime-state error.
301    ///
302    /// # Deferred errors
303    ///
304    /// Symbolic shape constraints can later produce
305    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
306    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
307    /// or execution.
308    fn jvp(&self, wrt: &TracedTensor, tangent: &TracedTensor) -> Result<TracedTensor>;
309
310    /// Like [`jvp`](Self::jvp), but returns `None` when `wrt` is inactive.
311    ///
312    /// # Examples
313    ///
314    /// ```rust
315    /// use tenferro_ad::TracedTensorAdExt;
316    /// use tenferro_runtime::TracedTensor;
317    ///
318    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
319    /// let y = TracedTensor::from_vec_col_major(vec![], vec![4.0_f64]).unwrap();
320    /// let tangent = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
321    /// let loss = (&y * &y).unwrap();
322    ///
323    /// assert!(loss.jvp_optional(&x, &tangent).unwrap().is_none());
324    /// ```
325    ///
326    /// # Errors
327    ///
328    /// Returns [`tenferro_runtime::Error::UnsupportedAdRule`] when a JVP rule
329    /// is unavailable, [`Error::Validation`] for incompatible tangent metadata,
330    /// or a typed backend/runtime-state error.
331    ///
332    /// # Deferred errors
333    ///
334    /// Symbolic shape constraints can later produce
335    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
336    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
337    /// or execution.
338    fn jvp_optional(
339        &self,
340        wrt: &TracedTensor,
341        tangent: &TracedTensor,
342    ) -> Result<Option<TracedTensor>>;
343
344    /// Reverse-mode vector-Jacobian product.
345    ///
346    /// Complex cotangents use tenferro's Hermitian real-inner-product
347    /// convention. Non-real complex cotangent seeds therefore need an explicit
348    /// seed-convention comparison when matching JAX. See
349    /// <https://tensor4all.org/tenferro-rs/guides/complex-ad.html>.
350    ///
351    /// # Examples
352    ///
353    /// ```rust
354    /// use tenferro_ad::TracedTensorAdExt;
355    /// use tenferro_cpu::CpuBackend;
356    /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
357    ///
358    /// fn eval(tensor: &TracedTensor) -> tenferro_runtime::Tensor {
359    ///     let mut compiler = GraphCompiler::new();
360    ///     let program = compiler.compile(tensor).unwrap();
361    ///     let backend = CpuBackend::new();
362    ///     let mut builder = Runtime::builder();
363    ///     builder
364    ///         .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
365    ///         .unwrap();
366    ///     let runtime = builder.build().unwrap();
367    ///     runtime.run_compiled(&program, &[]).unwrap().pop().unwrap()
368    /// }
369    ///
370    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
371    /// let cotangent = TracedTensor::from_vec_col_major(vec![], vec![0.5_f64]).unwrap();
372    /// let y = (&x * &x).unwrap();
373    /// let dx = y.vjp(&x, &cotangent).unwrap();
374    ///
375    /// assert_eq!(eval(&dx).as_slice::<f64>().unwrap(), &[3.0]);
376    /// ```
377    ///
378    /// # Errors
379    ///
380    /// Returns [`tenferro_runtime::Error::UnsupportedAdRule`] when a VJP rule
381    /// is unavailable, [`Error::Validation`] for incompatible cotangent
382    /// metadata, or a typed backend/runtime-state error.
383    ///
384    /// # Deferred errors
385    ///
386    /// Symbolic shape constraints can later produce
387    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
388    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
389    /// or execution.
390    fn vjp(&self, wrt: &TracedTensor, cotangent: &TracedTensor) -> Result<TracedTensor>;
391
392    /// Like [`vjp`](Self::vjp), but returns `None` when `wrt` is inactive.
393    ///
394    /// # Examples
395    ///
396    /// ```rust
397    /// use tenferro_ad::TracedTensorAdExt;
398    /// use tenferro_runtime::TracedTensor;
399    ///
400    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
401    /// let y = TracedTensor::from_vec_col_major(vec![], vec![4.0_f64]).unwrap();
402    /// let cotangent = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
403    /// let loss = (&y * &y).unwrap();
404    ///
405    /// assert!(loss.vjp_optional(&x, &cotangent).unwrap().is_none());
406    /// ```
407    ///
408    /// # Errors
409    ///
410    /// Returns [`tenferro_runtime::Error::UnsupportedAdRule`] when a VJP rule
411    /// is unavailable, [`Error::Validation`] for incompatible cotangent
412    /// metadata, or a typed backend/runtime-state error.
413    ///
414    /// # Deferred errors
415    ///
416    /// Symbolic shape constraints can later produce
417    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
418    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
419    /// or execution.
420    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            // Source input keys are reused so derivative traces compose with the
887            // original eager tensor. Keep those keys owned by the source tensor's
888            // metadata scopes; derivative input metadata may be shape-specialized
889            // for the current run and must not shadow the source while the VJP/JVP
890            // result stays alive.
891            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}