Skip to main content

tenferro_ad/
traced.rs

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
37/// Typed error for the documented, user-reachable inactive-`wrt` state.
38///
39/// The `_optional` siblings expose this state as `Ok(None)`; the strict wrapper
40/// must not report a reachable input condition as an internal invariant.
41pub(crate) fn inactive_wrt_error(op: &'static str, wrt_key: impl std::fmt::Debug) -> Error {
42    Error::invalid_argument(
43        op,
44        ErrorPhase::GraphBuild,
45        "wrt",
46        format!(
47            "{op} is inactive for {wrt_key:?}; use {op}_optional to observe the inactive state"
48        ),
49    )
50}
51
52pub(crate) fn grad_with_rules_and_cache(
53    output: &TracedTensor,
54    wrt: &TracedTensor,
55    rules: &SemanticExtensionRuleSet,
56    ad_transform_cache: Option<&AdTransformCache>,
57) -> Result<TracedTensor> {
58    grad_with_optional_rules(output, wrt, rules, ad_transform_cache)
59}
60
61pub(crate) fn jvp_with_rules_and_cache(
62    output: &TracedTensor,
63    wrt: &TracedTensor,
64    tangent: &TracedTensor,
65    rules: &SemanticExtensionRuleSet,
66    ad_transform_cache: Option<&AdTransformCache>,
67) -> Result<TracedTensor> {
68    let wrt_input_key = leaf_input_key(wrt)?;
69    jvp_many_with_rules_and_cache(output, &[(wrt, tangent)], rules, ad_transform_cache)?
70        .ok_or_else(|| inactive_wrt_error("jvp", &wrt_input_key))
71}
72
73pub(crate) fn jvp_many_with_rules_and_cache(
74    output: &TracedTensor,
75    wrt_tangents: &[(&TracedTensor, &TracedTensor)],
76    rules: &SemanticExtensionRuleSet,
77    ad_transform_cache: Option<&AdTransformCache>,
78) -> Result<Option<TracedTensor>> {
79    jvp_many_optional_impl(output, wrt_tangents, rules, ad_transform_cache)
80}
81
82pub(crate) fn grad_optional_with_rules_and_cache(
83    output: &TracedTensor,
84    wrt: &TracedTensor,
85    rules: &SemanticExtensionRuleSet,
86    ad_transform_cache: Option<&AdTransformCache>,
87) -> Result<Option<TracedTensor>> {
88    if output.rank != 0 {
89        return Err(Error::NonScalarGrad {
90            shape: error_shape_hint(output),
91        });
92    }
93
94    let ones = ones_tensor(output.dtype, vec![])?;
95    let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
96    vjp_many_with_transform_and_cache(output, &[wrt], &seed, rules, "grad", ad_transform_cache)
97        .map(|mut results| results.pop().flatten())
98}
99
100pub(crate) fn jvp_optional_with_rules_and_cache(
101    output: &TracedTensor,
102    wrt: &TracedTensor,
103    tangent: &TracedTensor,
104    rules: &SemanticExtensionRuleSet,
105    ad_transform_cache: Option<&AdTransformCache>,
106) -> Result<Option<TracedTensor>> {
107    jvp_many_with_rules_and_cache(output, &[(wrt, tangent)], rules, ad_transform_cache)
108}
109
110pub(crate) fn vjp_with_rules_and_cache(
111    output: &TracedTensor,
112    wrt: &TracedTensor,
113    cotangent: &TracedTensor,
114    rules: &SemanticExtensionRuleSet,
115    ad_transform_cache: Option<&AdTransformCache>,
116) -> Result<TracedTensor> {
117    let wrt_input_key = leaf_input_key(wrt)?;
118    vjp_many_with_rules_and_cache(output, &[wrt], cotangent, rules, ad_transform_cache)?
119        .into_iter()
120        .next()
121        .flatten()
122        .ok_or_else(|| inactive_wrt_error("vjp", &wrt_input_key))
123}
124
125pub(crate) fn vjp_many_with_rules_and_cache(
126    output: &TracedTensor,
127    wrts: &[&TracedTensor],
128    cotangent: &TracedTensor,
129    rules: &SemanticExtensionRuleSet,
130    ad_transform_cache: Option<&AdTransformCache>,
131) -> Result<Vec<Option<TracedTensor>>> {
132    vjp_many_with_transform_and_cache(output, wrts, cotangent, rules, "vjp", ad_transform_cache)
133}
134
135fn vjp_many_with_transform_and_cache(
136    output: &TracedTensor,
137    wrts: &[&TracedTensor],
138    cotangent: &TracedTensor,
139    rules: &SemanticExtensionRuleSet,
140    transform: &'static str,
141    ad_transform_cache: Option<&AdTransformCache>,
142) -> Result<Vec<Option<TracedTensor>>> {
143    vjp_many_optional_impl(
144        output,
145        wrts,
146        cotangent,
147        rules,
148        transform,
149        ad_transform_cache,
150    )
151}
152
153pub(crate) fn vjp_optional_with_rules_and_cache(
154    output: &TracedTensor,
155    wrt: &TracedTensor,
156    cotangent: &TracedTensor,
157    rules: &SemanticExtensionRuleSet,
158    ad_transform_cache: Option<&AdTransformCache>,
159) -> Result<Option<TracedTensor>> {
160    vjp_many_with_rules_and_cache(output, &[wrt], cotangent, rules, ad_transform_cache)
161        .map(|mut results| results.pop().flatten())
162}
163
164fn grad_with_optional_rules(
165    output: &TracedTensor,
166    wrt: &TracedTensor,
167    rules: &SemanticExtensionRuleSet,
168    ad_transform_cache: Option<&AdTransformCache>,
169) -> Result<TracedTensor> {
170    if output.rank != 0 {
171        return Err(Error::NonScalarGrad {
172            shape: error_shape_hint(output),
173        });
174    }
175
176    let ones = ones_tensor(output.dtype, vec![])?;
177    let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
178    let wrt_input_key = leaf_input_key(wrt)?;
179    vjp_many_with_transform_and_cache(output, &[wrt], &seed, rules, "grad", ad_transform_cache)?
180        .pop()
181        .flatten()
182        .ok_or_else(|| inactive_wrt_error("grad", &wrt_input_key))
183}
184
185fn single_runtime_output(mut outputs: Vec<Tensor>, op: &'static str) -> Result<Tensor> {
186    let actual = outputs.len();
187    if actual != 1 {
188        return Err(Error::runtime_state(
189            op,
190            ErrorPhase::Execution,
191            format!("expected one runtime output, got {actual}"),
192        ));
193    }
194    outputs.pop().ok_or_else(|| {
195        Error::runtime_state(
196            op,
197            ErrorPhase::Execution,
198            "runtime returned no output after successful output-count validation",
199        )
200    })
201}
202
203/// Automatic differentiation helpers for [`TracedTensor`].
204///
205/// # Examples
206///
207/// ```rust
208/// use tenferro_ad::TracedTensorAdExt;
209/// use tenferro_runtime::TracedTensor;
210///
211/// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
212/// let loss = x.scale_real(2.0).unwrap();
213/// let maybe_dx = loss.grad_optional(&x).unwrap();
214/// assert!(maybe_dx.is_some());
215/// ```
216pub trait TracedTensorAdExt {
217    /// Gradient of a scalar output with respect to a traced input.
218    ///
219    /// For complex scalar outputs, tenferro returns the Hermitian-adjoint
220    /// cotangent. To compare seed-`1` scalar gradients with JAX's public
221    /// `grad` values, use the complex conjugate of this result. See
222    /// <https://tensor4all.org/tenferro-rs/guides/complex-ad.html>.
223    ///
224    /// # Examples
225    ///
226    /// ```rust
227    /// use tenferro_ad::TracedTensorAdExt;
228    /// use tenferro_cpu::CpuBackend;
229    /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
230    ///
231    /// fn eval(tensor: &TracedTensor) -> tenferro_runtime::Tensor {
232    ///     let mut compiler = GraphCompiler::new();
233    ///     let program = compiler.compile(tensor).unwrap();
234    ///     let backend = CpuBackend::new();
235    ///     let mut builder = Runtime::builder();
236    ///     builder
237    ///         .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
238    ///         .unwrap();
239    ///     let runtime = builder.build().unwrap();
240    ///     runtime.run_compiled(&program, &[]).unwrap().pop().unwrap()
241    /// }
242    ///
243    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
244    /// let loss = (&x * &x).unwrap();
245    /// let dx = loss.grad(&x).unwrap();
246    ///
247    /// assert_eq!(eval(&dx).as_slice::<f64>().unwrap(), &[6.0]);
248    /// ```
249    ///
250    /// # Errors
251    ///
252    /// Returns [`tenferro_runtime::Error::NonScalarGrad`] for a non-scalar
253    /// output, [`tenferro_runtime::Error::UnsupportedAdRule`] when an AD rule
254    /// is unavailable, or a typed validation/backend/runtime-state error. An
255    /// inactive `wrt` returns [`Error::Validation`] with `argument: "wrt"`; use
256    /// [`grad_optional`](Self::grad_optional) to observe that state.
257    ///
258    /// # Deferred errors
259    ///
260    /// Symbolic shape constraints can later produce
261    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
262    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
263    /// or execution.
264    fn grad(&self, wrt: &TracedTensor) -> Result<TracedTensor>;
265
266    /// Like [`grad`](Self::grad), but returns `None` when `wrt` is inactive.
267    ///
268    /// # Examples
269    ///
270    /// ```rust
271    /// use tenferro_ad::TracedTensorAdExt;
272    /// use tenferro_runtime::TracedTensor;
273    ///
274    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
275    /// let y = TracedTensor::from_vec_col_major(vec![], vec![4.0_f64]).unwrap();
276    /// let loss = (&y * &y).unwrap();
277    ///
278    /// assert!(loss.grad_optional(&x).unwrap().is_none());
279    /// ```
280    ///
281    /// # Errors
282    ///
283    /// Returns [`tenferro_runtime::Error::NonScalarGrad`] for a non-scalar
284    /// output, [`Error::UnsupportedAdRule`] when an AD rule is unavailable, or
285    /// a typed validation/backend/runtime-state error.
286    ///
287    /// # Deferred errors
288    ///
289    /// Symbolic shape constraints can later produce
290    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
291    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
292    /// or execution.
293    fn grad_optional(&self, wrt: &TracedTensor) -> Result<Option<TracedTensor>>;
294
295    /// Evaluate this tensor and replace its graph with a concrete leaf while
296    /// preserving the previous graph for AD replay.
297    ///
298    /// # Examples
299    ///
300    /// ```rust
301    /// use tenferro_ad::TracedTensorAdExt;
302    /// use tenferro_cpu::CpuBackend;
303    /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
304    ///
305    /// let mut compiler = GraphCompiler::new();
306    /// let backend = CpuBackend::new();
307    /// let mut builder = Runtime::builder();
308    /// builder
309    ///     .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
310    ///     .unwrap();
311    /// let runtime = builder.build().unwrap();
312    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
313    /// let mut y = (&x * &x).unwrap();
314    ///
315    /// y.checkpoint(&mut compiler, &runtime).unwrap();
316    ///
317    /// let value = y.attached_value().unwrap();
318    /// let materialized = value.tensor_read().unwrap().tensor_view().duplicate().unwrap();
319    /// assert_eq!(materialized.as_slice::<f64>().unwrap(), &[9.0]);
320    /// ```
321    ///
322    /// # Errors
323    ///
324    /// Returns [`tenferro_runtime::Error::Validation`] when checkpoint metadata
325    /// is invalid, [`Error::RuntimeState`] when graph metadata or runtime
326    /// state is unavailable, or a typed backend error from evaluation.
327    fn checkpoint(&mut self, compiler: &mut GraphCompiler, runtime: &Runtime) -> Result<()>;
328
329    /// Forward-mode Jacobian-vector product.
330    ///
331    /// # Examples
332    ///
333    /// ```rust
334    /// use tenferro_ad::TracedTensorAdExt;
335    /// use tenferro_cpu::CpuBackend;
336    /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
337    ///
338    /// fn eval(tensor: &TracedTensor) -> tenferro_runtime::Tensor {
339    ///     let mut compiler = GraphCompiler::new();
340    ///     let program = compiler.compile(tensor).unwrap();
341    ///     let backend = CpuBackend::new();
342    ///     let mut builder = Runtime::builder();
343    ///     builder
344    ///         .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
345    ///         .unwrap();
346    ///     let runtime = builder.build().unwrap();
347    ///     runtime.run_compiled(&program, &[]).unwrap().pop().unwrap()
348    /// }
349    ///
350    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
351    /// let tangent = TracedTensor::from_vec_col_major(vec![], vec![2.0_f64]).unwrap();
352    /// let y = (&x * &x).unwrap();
353    /// let dy = y.jvp(&x, &tangent).unwrap();
354    ///
355    /// assert_eq!(eval(&dy).as_slice::<f64>().unwrap(), &[12.0]);
356    /// ```
357    ///
358    /// # Errors
359    ///
360    /// Returns [`tenferro_runtime::Error::UnsupportedAdRule`] when a JVP rule
361    /// is unavailable, [`Error::Validation`] for incompatible tangent metadata,
362    /// or a typed backend/runtime-state error. An inactive `wrt` returns
363    /// [`Error::Validation`] with `argument: "wrt"`; use
364    /// [`jvp_optional`](Self::jvp_optional) to observe that state.
365    ///
366    /// # Deferred errors
367    ///
368    /// Symbolic shape constraints can later produce
369    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
370    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
371    /// or execution.
372    fn jvp(&self, wrt: &TracedTensor, tangent: &TracedTensor) -> Result<TracedTensor>;
373
374    /// Like [`jvp`](Self::jvp), but returns `None` when `wrt` is inactive.
375    ///
376    /// # Examples
377    ///
378    /// ```rust
379    /// use tenferro_ad::TracedTensorAdExt;
380    /// use tenferro_runtime::TracedTensor;
381    ///
382    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
383    /// let y = TracedTensor::from_vec_col_major(vec![], vec![4.0_f64]).unwrap();
384    /// let tangent = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
385    /// let loss = (&y * &y).unwrap();
386    ///
387    /// assert!(loss.jvp_optional(&x, &tangent).unwrap().is_none());
388    /// ```
389    ///
390    /// # Errors
391    ///
392    /// Returns [`tenferro_runtime::Error::UnsupportedAdRule`] when a JVP rule
393    /// is unavailable, [`Error::Validation`] for incompatible tangent metadata,
394    /// or a typed backend/runtime-state error.
395    ///
396    /// # Deferred errors
397    ///
398    /// Symbolic shape constraints can later produce
399    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
400    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
401    /// or execution.
402    fn jvp_optional(
403        &self,
404        wrt: &TracedTensor,
405        tangent: &TracedTensor,
406    ) -> Result<Option<TracedTensor>>;
407
408    /// Reverse-mode vector-Jacobian product.
409    ///
410    /// Complex cotangents use tenferro's Hermitian real-inner-product
411    /// convention. Non-real complex cotangent seeds therefore need an explicit
412    /// seed-convention comparison when matching JAX. See
413    /// <https://tensor4all.org/tenferro-rs/guides/complex-ad.html>.
414    ///
415    /// # Examples
416    ///
417    /// ```rust
418    /// use tenferro_ad::TracedTensorAdExt;
419    /// use tenferro_cpu::CpuBackend;
420    /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
421    ///
422    /// fn eval(tensor: &TracedTensor) -> tenferro_runtime::Tensor {
423    ///     let mut compiler = GraphCompiler::new();
424    ///     let program = compiler.compile(tensor).unwrap();
425    ///     let backend = CpuBackend::new();
426    ///     let mut builder = Runtime::builder();
427    ///     builder
428    ///         .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
429    ///         .unwrap();
430    ///     let runtime = builder.build().unwrap();
431    ///     runtime.run_compiled(&program, &[]).unwrap().pop().unwrap()
432    /// }
433    ///
434    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
435    /// let cotangent = TracedTensor::from_vec_col_major(vec![], vec![0.5_f64]).unwrap();
436    /// let y = (&x * &x).unwrap();
437    /// let dx = y.vjp(&x, &cotangent).unwrap();
438    ///
439    /// assert_eq!(eval(&dx).as_slice::<f64>().unwrap(), &[3.0]);
440    /// ```
441    ///
442    /// # Errors
443    ///
444    /// Returns [`tenferro_runtime::Error::UnsupportedAdRule`] when a VJP rule
445    /// is unavailable, [`Error::Validation`] for incompatible cotangent
446    /// metadata, or a typed backend/runtime-state error. An inactive `wrt`
447    /// returns [`Error::Validation`] with `argument: "wrt"`; use
448    /// [`vjp_optional`](Self::vjp_optional) to observe that state.
449    ///
450    /// # Deferred errors
451    ///
452    /// Symbolic shape constraints can later produce
453    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
454    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
455    /// or execution.
456    fn vjp(&self, wrt: &TracedTensor, cotangent: &TracedTensor) -> Result<TracedTensor>;
457
458    /// Like [`vjp`](Self::vjp), but returns `None` when `wrt` is inactive.
459    ///
460    /// # Examples
461    ///
462    /// ```rust
463    /// use tenferro_ad::TracedTensorAdExt;
464    /// use tenferro_runtime::TracedTensor;
465    ///
466    /// let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
467    /// let y = TracedTensor::from_vec_col_major(vec![], vec![4.0_f64]).unwrap();
468    /// let cotangent = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
469    /// let loss = (&y * &y).unwrap();
470    ///
471    /// assert!(loss.vjp_optional(&x, &cotangent).unwrap().is_none());
472    /// ```
473    ///
474    /// # Errors
475    ///
476    /// Returns [`tenferro_runtime::Error::UnsupportedAdRule`] when a VJP rule
477    /// is unavailable, [`Error::Validation`] for incompatible cotangent
478    /// metadata, or a typed backend/runtime-state error.
479    ///
480    /// # Deferred errors
481    ///
482    /// Symbolic shape constraints can later produce
483    /// [`tenferro_runtime::Error::ShapeConstraintViolation`] or
484    /// [`tenferro_runtime::Error::ShapeConstraintEvaluation`] during compile
485    /// or execution.
486    fn vjp_optional(
487        &self,
488        wrt: &TracedTensor,
489        cotangent: &TracedTensor,
490    ) -> Result<Option<TracedTensor>>;
491}
492
493impl TracedTensorAdExt for TracedTensor {
494    fn grad(&self, wrt: &TracedTensor) -> Result<TracedTensor> {
495        let rules = SemanticExtensionRuleSet::default();
496        grad_with_optional_rules(self, wrt, &rules, None)
497    }
498
499    fn grad_optional(&self, wrt: &TracedTensor) -> Result<Option<TracedTensor>> {
500        if self.rank != 0 {
501            return Err(Error::NonScalarGrad {
502                shape: error_shape_hint(self),
503            });
504        }
505
506        let ones = ones_tensor(self.dtype, vec![])?;
507        let seed = TracedTensor::from_tensor_concrete_shape(ones)?;
508        let rules = SemanticExtensionRuleSet::default();
509        vjp_many_with_transform_and_cache(self, &[wrt], &seed, &rules, "grad", None)
510            .map(|mut results| results.pop().flatten())
511    }
512
513    fn checkpoint(&mut self, compiler: &mut GraphCompiler, runtime: &Runtime) -> Result<()> {
514        let data = if let Some(data) = self.attached_value() {
515            Arc::clone(data)
516        } else {
517            let program = compiler.compile(self)?;
518            Arc::new(RetainedValue::from_tensor(single_runtime_output(
519                runtime.run_compiled(&program, &[])?,
520                "TracedTensorAdExt::checkpoint",
521            )?))
522        };
523        checkpoint_tensor(self, data)?;
524        Ok(())
525    }
526
527    fn jvp(&self, wrt: &TracedTensor, tangent: &TracedTensor) -> Result<TracedTensor> {
528        let wrt_input_key = leaf_input_key(wrt)?;
529        self.jvp_optional(wrt, tangent)?
530            .ok_or_else(|| inactive_wrt_error("jvp", &wrt_input_key))
531    }
532
533    fn jvp_optional(
534        &self,
535        wrt: &TracedTensor,
536        tangent: &TracedTensor,
537    ) -> Result<Option<TracedTensor>> {
538        let rules = SemanticExtensionRuleSet::default();
539        jvp_many_optional_impl(self, &[(wrt, tangent)], &rules, None)
540    }
541
542    fn vjp(&self, wrt: &TracedTensor, cotangent: &TracedTensor) -> Result<TracedTensor> {
543        let wrt_input_key = leaf_input_key(wrt)?;
544        self.vjp_optional(wrt, cotangent)?
545            .ok_or_else(|| inactive_wrt_error("vjp", &wrt_input_key))
546    }
547
548    fn vjp_optional(
549        &self,
550        wrt: &TracedTensor,
551        cotangent: &TracedTensor,
552    ) -> Result<Option<TracedTensor>> {
553        let rules = SemanticExtensionRuleSet::default();
554        vjp_many_with_transform_and_cache(self, &[wrt], cotangent, &rules, "vjp", None)
555            .map(|mut results| results.pop().flatten())
556    }
557}
558
559fn jvp_many_optional_impl(
560    output: &TracedTensor,
561    wrt_tangents: &[(&TracedTensor, &TracedTensor)],
562    rules: &SemanticExtensionRuleSet,
563    ad_transform_cache: Option<&AdTransformCache>,
564) -> Result<Option<TracedTensor>> {
565    if wrt_tangents.is_empty() {
566        return Ok(None);
567    }
568
569    let mut seen = HashSet::with_capacity(wrt_tangents.len());
570    let mut requested = Vec::with_capacity(wrt_tangents.len());
571    for (wrt, tangent) in wrt_tangents {
572        let key = leaf_input_key(wrt)?;
573        if !seen.insert(key.clone()) {
574            return Err(Error::invalid_argument(
575                "jvp",
576                ErrorPhase::GraphBuild,
577                "wrt_tangents",
578                "wrt_tangents contains duplicate wrt leaves",
579            ));
580        }
581        let tangent_data = tangent.attached_value().cloned().ok_or_else(|| {
582            Error::invalid_argument(
583                "jvp",
584                ErrorPhase::GraphBuild,
585                "tangent",
586                "jvp tangent must have concrete tensor data",
587            )
588        })?;
589        requested.push((key, tangent_data));
590    }
591
592    let mut compiler = GraphCompiler::new();
593    let source = compile_ad_source(&mut compiler, output)?;
594    let mut active_inputs = vec![false; source.input_count()];
595    let mut source_indices = Vec::with_capacity(requested.len());
596    for (key, _) in &requested {
597        let source_index = source.input_key_index(key);
598        if let Some(index) = source_index {
599            active_inputs[index] = true;
600        }
601        source_indices.push(source_index);
602    }
603    if !active_inputs.iter().any(|active| *active) {
604        return Ok(None);
605    }
606
607    let derivative = semantic_jvp_with_cache(
608        source.frozen_program(),
609        &active_inputs,
610        rules,
611        ad_transform_cache,
612    )?;
613    let Some(derivative_output_index) = derivative
614        .derivative_output_indices()
615        .first()
616        .copied()
617        .flatten()
618    else {
619        return Ok(None);
620    };
621
622    let mut seed_tensors = Vec::with_capacity(requested.len());
623    for ((_, tangent_data), source_index) in requested.iter().zip(source_indices) {
624        let Some(source_index) = source_index else {
625            continue;
626        };
627        let Some(seed_input_index) = derivative
628            .derivative_input_indices()
629            .get(source_index)
630            .copied()
631            .flatten()
632        else {
633            continue;
634        };
635        seed_tensors.push((seed_input_index, Arc::clone(tangent_data)));
636    }
637    let mut inherited_tensors = Vec::with_capacity(1 + 2 * wrt_tangents.len());
638    inherited_tensors.push(output);
639    for (wrt, tangent) in wrt_tangents {
640        inherited_tensors.push(*wrt);
641        inherited_tensors.push(*tangent);
642    }
643    let mut traces = derivative_tensors_from_program(
644        &source,
645        &derivative,
646        &[derivative_output_index],
647        &seed_tensors,
648        &inherited_tensors,
649        &[tensor_shape_hint(output)],
650        "jvp",
651    )?;
652    traces.pop().map(Some).ok_or_else(|| {
653        Error::runtime_state(
654            "jvp",
655            ErrorPhase::GraphBuild,
656            "derivative program returned no requested output",
657        )
658    })
659}
660
661fn vjp_many_optional_impl(
662    output: &TracedTensor,
663    wrts: &[&TracedTensor],
664    cotangent: &TracedTensor,
665    rules: &SemanticExtensionRuleSet,
666    transform: &'static str,
667    ad_transform_cache: Option<&AdTransformCache>,
668) -> Result<Vec<Option<TracedTensor>>> {
669    let cotangent_data = cotangent.attached_value().cloned().ok_or_else(|| {
670        Error::invalid_argument(
671            transform,
672            ErrorPhase::GraphBuild,
673            "cotangent",
674            "vjp cotangent must have concrete tensor data",
675        )
676    })?;
677    if wrts.is_empty() {
678        return Ok(Vec::new());
679    }
680
681    let requested_keys = wrts
682        .iter()
683        .map(|wrt| leaf_input_key(wrt))
684        .collect::<Result<Vec<_>>>()?;
685    let mut compiler = GraphCompiler::new();
686    let source = compile_ad_source(&mut compiler, output)?;
687    let source_indices: Vec<_> = requested_keys
688        .iter()
689        .map(|key| source.input_key_index(key))
690        .collect();
691    let mut active_inputs = vec![false; source.input_count()];
692    for source_index in source_indices.iter().flatten().copied() {
693        active_inputs[source_index] = true;
694    }
695    if !active_inputs.iter().any(|active| *active) {
696        return Ok(vec![None; wrts.len()]);
697    }
698
699    let active_outputs = vec![true; source.output_count()];
700    let derivative = semantic_vjp_with_cache(
701        source.frozen_program(),
702        &active_inputs,
703        &active_outputs,
704        rules,
705        ad_transform_cache,
706    )?;
707    let Some(seed_input_index) = derivative
708        .derivative_input_indices()
709        .first()
710        .copied()
711        .flatten()
712    else {
713        return Ok(vec![None; wrts.len()]);
714    };
715
716    let mut output_slots = vec![None; source.input_count()];
717    let mut derivative_output_indices = Vec::new();
718    let mut fallback_shape_hints = Vec::new();
719    for (request_index, source_index) in source_indices.iter().enumerate() {
720        let Some(source_index) = source_index else {
721            continue;
722        };
723        if output_slots[*source_index].is_some() {
724            continue;
725        }
726        let Some(derivative_output_index) = derivative
727            .derivative_output_indices()
728            .get(*source_index)
729            .copied()
730            .flatten()
731        else {
732            continue;
733        };
734        output_slots[*source_index] = Some(derivative_output_indices.len());
735        derivative_output_indices.push(derivative_output_index);
736        fallback_shape_hints.push(tensor_shape_hint(wrts[request_index]));
737    }
738    if derivative_output_indices.is_empty() {
739        return Ok(vec![None; wrts.len()]);
740    }
741
742    let inherited_tensors: Vec<&TracedTensor> = std::iter::once(output)
743        .chain(wrts.iter().copied())
744        .chain(std::iter::once(cotangent))
745        .collect();
746    let traces = derivative_tensors_from_program(
747        &source,
748        &derivative,
749        &derivative_output_indices,
750        &[(seed_input_index, cotangent_data)],
751        &inherited_tensors,
752        &fallback_shape_hints,
753        transform,
754    )?;
755    Ok(source_indices
756        .into_iter()
757        .map(|source_index| {
758            source_index
759                .and_then(|index| output_slots[index])
760                .map(|slot| traces[slot].clone())
761        })
762        .collect())
763}
764
765fn semantic_jvp_with_cache(
766    source: &FrozenProgram,
767    active_inputs: &[bool],
768    rules: &SemanticExtensionRuleSet,
769    ad_transform_cache: Option<&AdTransformCache>,
770) -> Result<SemanticAdProgram> {
771    let key = SemanticAdTransformCacheKey::jvp(source, active_inputs);
772    if let Some(cache) = ad_transform_cache
773        && let Some(cached) = cache.get_semantic(&key, source)?
774    {
775        return cached
776            .as_ref()
777            .with_input_prefix_bindings_from(source)
778            .map_err(|source| {
779                Error::runtime_state_source(
780                    "semantic traced jvp cache",
781                    ErrorPhase::GraphBuild,
782                    source,
783                )
784            });
785    }
786    let derivative =
787        semantic_jvp(source, active_inputs, rules).map_err(semantic_transform_error("jvp"))?;
788    if let Some(cache) = ad_transform_cache {
789        cache.put_semantic(key, source, Arc::new(derivative.clone()))?;
790    }
791    Ok(derivative)
792}
793
794fn semantic_vjp_with_cache(
795    source: &FrozenProgram,
796    active_inputs: &[bool],
797    active_outputs: &[bool],
798    rules: &SemanticExtensionRuleSet,
799    ad_transform_cache: Option<&AdTransformCache>,
800) -> Result<SemanticAdProgram> {
801    let key = SemanticAdTransformCacheKey::vjp(source, active_inputs, active_outputs);
802    if let Some(cache) = ad_transform_cache
803        && let Some(cached) = cache.get_semantic(&key, source)?
804    {
805        return cached
806            .as_ref()
807            .with_input_prefix_bindings_from(source)
808            .map_err(|source| {
809                Error::runtime_state_source(
810                    "semantic traced vjp cache",
811                    ErrorPhase::GraphBuild,
812                    source,
813                )
814            });
815    }
816    let derivative = semantic_vjp(source, active_inputs, active_outputs, rules)
817        .map_err(semantic_transform_error("vjp"))?;
818    if let Some(cache) = ad_transform_cache {
819        cache.put_semantic(key, source, Arc::new(derivative.clone()))?;
820    }
821    Ok(derivative)
822}
823
824fn semantic_transform_error(
825    transform: &'static str,
826) -> impl FnOnce(SemanticAdTransformError) -> Error {
827    move |source| {
828        semantic_transform_validation_error(transform, &source).unwrap_or_else(|| {
829            Error::runtime_state_source(transform, ErrorPhase::GraphBuild, source)
830        })
831    }
832}
833
834fn semantic_transform_validation_error(
835    transform: &'static str,
836    source: &SemanticAdTransformError,
837) -> Option<Error> {
838    if let SemanticAdTransformError::Extension(
839        SemanticAdError::Unsupported { family_id, .. }
840        | SemanticAdError::MissingRule { family_id, .. },
841    ) = source
842    {
843        return Some(Error::UnsupportedAdRule {
844            transform,
845            op: (*family_id).to_owned(),
846        });
847    }
848
849    let SemanticAdTransformError::Extension(SemanticAdError::Rule { source, .. }) = source else {
850        return None;
851    };
852    let tenferro_ops::ad::ADRuleError::InvalidInput { op, message, .. } =
853        source.downcast_ref::<tenferro_ops::ad::ADRuleError>()?
854    else {
855        return None;
856    };
857    Some(Error::invalid_argument(
858        transform,
859        ErrorPhase::GraphBuild,
860        "semantic_ad_rule",
861        format!("{op}: {message}"),
862    ))
863}
864
865fn derivative_tensors_from_program(
866    source: &CompiledGraph,
867    derivative: &SemanticAdProgram,
868    derivative_output_indices: &[usize],
869    seed_tensors: &[(usize, Arc<RetainedValue>)],
870    inherited_tensors: &[&TracedTensor],
871    fallback_shape_hints: &[Option<Vec<SymDim>>],
872    transform: &'static str,
873) -> Result<Vec<TracedTensor>> {
874    derivative_traces_from_frozen_program(
875        source,
876        derivative.frozen(),
877        derivative_output_indices,
878        seed_tensors,
879        inherited_tensors,
880        fallback_shape_hints,
881        transform,
882    )
883}
884
885pub(crate) fn derivative_trace_from_frozen_program(
886    source: &CompiledGraph,
887    frozen: &FrozenProgram,
888    derivative_output_index: usize,
889    seed_tensors: &[(usize, Arc<RetainedValue>)],
890    inherited_tensors: &[&TracedTensor],
891    fallback_shape_hint: Option<Vec<SymDim>>,
892    transform: &'static str,
893) -> Result<TracedTensor> {
894    let mut outputs = derivative_traces_from_frozen_program(
895        source,
896        frozen,
897        &[derivative_output_index],
898        seed_tensors,
899        inherited_tensors,
900        &[fallback_shape_hint],
901        transform,
902    )?;
903    outputs.pop().ok_or_else(|| {
904        Error::runtime_state(
905            transform,
906            ErrorPhase::GraphBuild,
907            "derivative program returned no requested output",
908        )
909    })
910}
911
912fn derivative_traces_from_frozen_program(
913    source: &CompiledGraph,
914    frozen: &FrozenProgram,
915    derivative_output_indices: &[usize],
916    seed_tensors: &[(usize, Arc<RetainedValue>)],
917    inherited_tensors: &[&TracedTensor],
918    fallback_shape_hints: &[Option<Vec<SymDim>>],
919    transform: &'static str,
920) -> Result<Vec<TracedTensor>> {
921    let input_shapes = symbolic_input_shapes(frozen)?;
922    let input_shape_refs: Vec<_> = input_shapes.iter().map(Vec::as_slice).collect();
923    let input_metas = frozen
924        .program
925        .inputs()
926        .iter()
927        .copied()
928        .map(|value| tensor_meta_for_value(frozen, value, &input_shape_refs, transform))
929        .collect::<Result<Vec<_>>>()?;
930
931    if derivative_output_indices.len() != fallback_shape_hints.len() {
932        return Err(Error::runtime_state(
933            transform,
934            ErrorPhase::GraphBuild,
935            "derivative output indices and fallback shape hints differ in length",
936        ));
937    }
938    let output_metas = derivative_output_indices
939        .iter()
940        .map(|&derivative_output_index| {
941            let output_value = *frozen
942                .program
943                .outputs()
944                .get(derivative_output_index)
945                .ok_or_else(|| {
946                    Error::runtime_state(
947                        transform,
948                        ErrorPhase::GraphBuild,
949                        format!(
950                            "derivative output index {derivative_output_index} is outside {} outputs",
951                            frozen.program.outputs().len()
952                        ),
953                    )
954                })?;
955            tensor_meta_for_value(frozen, output_value, &input_shape_refs, transform)
956        })
957        .collect::<Result<Vec<_>>>()?;
958
959    let mut builder = GraphBuilder::<StdTensorOp>::new();
960    let mut value_map = HashMap::<ProgramValue, LocalValueId>::new();
961    let mut input_keys = Vec::with_capacity(frozen.program.inputs().len());
962    for (input_index, input) in frozen.program.inputs().iter().copied().enumerate() {
963        let key = if input_index < source.input_keys().len() {
964            source.input_keys()[input_index].clone()
965        } else {
966            allocate_input_key()
967        };
968        let local = builder.add_input(key.clone());
969        value_map.insert(input, local);
970        input_keys.push(key);
971    }
972
973    for operation in frozen.program.operations() {
974        let inputs = operation
975            .inputs()
976            .iter()
977            .copied()
978            .map(|value| {
979                value_map
980                    .get(&value)
981                    .copied()
982                    .map(ValueRef::Local)
983                    .ok_or_else(|| missing_program_value(transform, "operation input"))
984            })
985            .collect::<Result<Vec<_>>>()?;
986        let op = match operation.op() {
987            SemanticOpRef::Core(op) => StdTensorOp::from(op),
988            SemanticOpRef::Extension(op) => StdTensorOp::Extension(op.clone_arc()),
989            _ => {
990                return Err(Error::runtime_state(
991                    transform,
992                    ErrorPhase::GraphBuild,
993                    "unsupported semantic operation variant in derivative graph",
994                ));
995            }
996        };
997        let outputs = builder.add_operation(op, inputs, OperationRole::Primary);
998        if outputs.len() != operation.outputs().len() {
999            return Err(Error::runtime_state(
1000                transform,
1001                ErrorPhase::GraphBuild,
1002                format!(
1003                    "semantic operation expected {} outputs, graph builder produced {}",
1004                    operation.outputs().len(),
1005                    outputs.len()
1006                ),
1007            ));
1008        }
1009        for (value, local) in operation.outputs().iter().copied().zip(outputs) {
1010            value_map.insert(value, local);
1011        }
1012    }
1013
1014    let graph_outputs = frozen
1015        .program
1016        .outputs()
1017        .iter()
1018        .copied()
1019        .map(|value| {
1020            value_map
1021                .get(&value)
1022                .copied()
1023                .ok_or_else(|| missing_program_value(transform, "program output"))
1024        })
1025        .collect::<Result<Vec<_>>>()?;
1026    let vals = derivative_output_indices
1027        .iter()
1028        .map(|&index| {
1029            graph_outputs.get(index).copied().ok_or_else(|| {
1030                Error::runtime_state(
1031                    transform,
1032                    ErrorPhase::GraphBuild,
1033                    "derivative output index missing after graph conversion",
1034                )
1035            })
1036        })
1037        .collect::<Result<Vec<_>>>()?;
1038    builder.set_outputs(graph_outputs);
1039    let graph = Arc::new(builder.build());
1040
1041    let Some(primary_tensor) = inherited_tensors.first() else {
1042        return Err(Error::runtime_state(
1043            transform,
1044            ErrorPhase::GraphBuild,
1045            "derivative trace construction requires inherited source tensors",
1046        ));
1047    };
1048    let mut inputs_map = (*tensor_inputs_map(primary_tensor)).clone();
1049    for (input_index, key) in input_keys.iter().enumerate() {
1050        if let Some(tensor) = frozen_input_value(frozen, input_index) {
1051            inputs_map.insert(key.clone(), tensor);
1052        }
1053    }
1054    for (seed_input_index, tensor) in seed_tensors {
1055        let meta = input_metas.get(*seed_input_index).ok_or_else(|| {
1056            Error::runtime_state(
1057                transform,
1058                ErrorPhase::GraphBuild,
1059                format!("seed input index {seed_input_index} is outside derivative inputs"),
1060            )
1061        })?;
1062        validate_seed_tensor(transform, *seed_input_index, tensor.as_ref(), meta)?;
1063        let key = input_keys.get(*seed_input_index).ok_or_else(|| {
1064            Error::runtime_state(
1065                transform,
1066                ErrorPhase::GraphBuild,
1067                format!("seed input key {seed_input_index} is outside derivative inputs"),
1068            )
1069        })?;
1070        inputs_map.insert(key.clone(), Arc::clone(tensor));
1071    }
1072
1073    let source_input_count = source.input_keys().len();
1074    let graph_input_metadata = graph
1075        .inputs()
1076        .iter()
1077        .copied()
1078        .zip(input_metas.iter().cloned())
1079        .enumerate()
1080        .filter_map(|(input_index, (input, meta))| {
1081            // Source input keys are reused so derivative traces compose with the
1082            // original eager tensor. Keep those keys owned by the source tensor's
1083            // metadata scopes; derivative input metadata may be shape-specialized
1084            // for the current run and must not shadow the source while the VJP/JVP
1085            // result stays alive.
1086            if input_index < source_input_count {
1087                None
1088            } else {
1089                Some((graph.values()[input].key.clone(), meta))
1090            }
1091        });
1092    let analysis = register_scoped_graph_analysis(graph.as_ref(), graph_input_metadata)?;
1093    let inherited_constraint_scopes = inherited_tensors
1094        .iter()
1095        .map(|tensor| ConstraintScopeTransfer::from_tensor(tensor))
1096        .collect::<Vec<_>>();
1097
1098    let inputs_map = Arc::new(inputs_map);
1099    // The derivative graph is standalone (its input keys reference the source
1100    // leaves directly), so inherited leaf metadata covers every leaf key the
1101    // deferred eager analysis may re-seed from these traces.
1102    let leaf_metas = merge_traced_leaf_metas(inherited_tensors.iter().copied());
1103    let metadata_scopes = metadata_scopes_with_new(
1104        analysis.metadata,
1105        inherited_tensors
1106            .iter()
1107            .map(|tensor| tensor_metadata_scopes(tensor)),
1108    );
1109    let constraint_scope_transfer =
1110        ConstraintScopeTransfer::with_new(analysis.constraints, inherited_constraint_scopes.iter());
1111
1112    Ok(vals
1113        .into_iter()
1114        .zip(output_metas)
1115        .zip(fallback_shape_hints)
1116        .map(|((val, output_meta), fallback_shape_hint)| {
1117            tensor_from_parts(TracedTensorParts {
1118                rank: output_meta.rank(),
1119                dtype: output_meta.dtype,
1120                graph: Arc::clone(&graph),
1121                val,
1122                data: None,
1123                shape_hint: output_meta
1124                    .exact_shape()
1125                    .or_else(|| fallback_shape_hint.clone()),
1126                inputs_map: Arc::clone(&inputs_map),
1127                leaf_metas: Arc::clone(&leaf_metas),
1128                extra_roots: Vec::new(),
1129                checkpoint_chain: None,
1130                metadata_scopes: metadata_scopes.clone(),
1131                constraint_scope_transfer: constraint_scope_transfer.clone(),
1132            })
1133        })
1134        .collect())
1135}
1136
1137fn missing_program_value(transform: &'static str, role: &'static str) -> Error {
1138    Error::runtime_state(
1139        transform,
1140        ErrorPhase::GraphBuild,
1141        format!("semantic derivative graph references missing {role}"),
1142    )
1143}
1144
1145fn symbolic_input_shapes(frozen: &FrozenProgram) -> Result<Vec<Vec<SymDim>>> {
1146    frozen
1147        .program
1148        .inputs()
1149        .iter()
1150        .copied()
1151        .map(|value| {
1152            let meta = frozen.program.value_metadata(value).map_err(|source| {
1153                Error::runtime_state_source(
1154                    "semantic traced AD input metadata",
1155                    ErrorPhase::GraphBuild,
1156                    source,
1157                )
1158            })?;
1159            let tensor_id = allocate_shape_tensor_id();
1160            Ok((0..meta.shape().len())
1161                .map(|axis| SymDim::tensor_axis(tensor_id, axis))
1162                .collect())
1163        })
1164        .collect()
1165}
1166
1167fn tensor_meta_for_value(
1168    frozen: &FrozenProgram,
1169    value: ProgramValue,
1170    input_shapes: &[&[SymDim]],
1171    transform: &'static str,
1172) -> Result<TensorMeta> {
1173    let meta = frozen
1174        .program
1175        .value_metadata(value)
1176        .map_err(|source| Error::runtime_state_source(transform, ErrorPhase::GraphBuild, source))?;
1177    Ok(program_metadata_to_tensor_meta(meta, input_shapes))
1178}
1179
1180fn program_metadata_to_tensor_meta(
1181    metadata: &ProgramValueMetadata,
1182    input_shapes: &[&[SymDim]],
1183) -> TensorMeta {
1184    let extents = metadata
1185        .shape()
1186        .iter()
1187        .cloned()
1188        .map(|extent| extent.map(|dim| SymDim::from_dim_expr(&dim, input_shapes)))
1189        .collect();
1190    // An externally defined scalar keeps the canonical identity its program declared,
1191    // so a differentiated value can still be given a reproducible program identity.
1192    let meta = TensorMeta::with_extents(metadata.dtype(), extents);
1193    match metadata.scalar_identity() {
1194        Some(identity) => meta.with_scalar_identity(identity),
1195        None => meta,
1196    }
1197}
1198
1199fn validate_seed_tensor(
1200    transform: &'static str,
1201    input_index: usize,
1202    tensor: &RetainedValue,
1203    expected: &TensorMeta,
1204) -> Result<()> {
1205    let actual_dtype = tensor.dtype();
1206    if actual_dtype != expected.dtype {
1207        return Err(Error::invalid_argument(
1208            transform,
1209            ErrorPhase::GraphBuild,
1210            "seed",
1211            format!(
1212                "seed input {input_index} dtype mismatch: expected {:?}, got {:?}",
1213                expected.dtype, actual_dtype
1214            ),
1215        ));
1216    }
1217    let actual_shape = tensor.shape();
1218    if actual_shape.len() != expected.rank() {
1219        return Err(Error::invalid_argument(
1220            transform,
1221            ErrorPhase::GraphBuild,
1222            "seed",
1223            format!(
1224                "seed input {input_index} rank mismatch: expected {}, got {}",
1225                expected.rank(),
1226                actual_shape.len()
1227            ),
1228        ));
1229    }
1230    if let Some(expected_shape) = expected
1231        .exact_shape()
1232        .filter(|shape| shape.iter().all(|dim| dim.constant_value().is_some()))
1233        .map(|shape| {
1234            shape
1235                .into_iter()
1236                .map(|dim| dim.constant_value().expect("filtered constant shape"))
1237                .collect::<Vec<_>>()
1238        })
1239        && expected_shape != actual_shape
1240    {
1241        return Err(Error::invalid_argument(
1242            transform,
1243            ErrorPhase::GraphBuild,
1244            "seed",
1245            format!(
1246                "seed input {input_index} shape mismatch: expected {:?}, got {:?}",
1247                expected_shape, actual_shape
1248            ),
1249        ));
1250    }
1251    Ok(())
1252}
1253
1254#[cfg(test)]
1255mod semantic_transform_error_tests {
1256    use super::*;
1257    use crate::semantic_extension::SemanticAdRuleRole;
1258
1259    #[test]
1260    fn unsupported_semantic_rule_maps_to_public_transform_error() {
1261        let source = SemanticAdTransformError::Extension(SemanticAdError::Unsupported {
1262            family_id: "tenferro-tests.unsupported.v1",
1263            role: SemanticAdRuleRole::LinearTranspose,
1264            message: "unsupported test payload".into(),
1265        });
1266
1267        let error = semantic_transform_validation_error("vjp", &source)
1268            .expect("semantic rejection must map to a public unsupported-rule error");
1269
1270        assert!(matches!(
1271            error,
1272            Error::UnsupportedAdRule { transform: "vjp", ref op }
1273                if op == "tenferro-tests.unsupported.v1"
1274        ));
1275    }
1276
1277    #[test]
1278    fn missing_semantic_rule_maps_to_public_transform_error() {
1279        let source = SemanticAdTransformError::Extension(SemanticAdError::MissingRule {
1280            family_id: "tenferro-tests.missing.v1",
1281            role: SemanticAdRuleRole::Linearize,
1282        });
1283
1284        let error = semantic_transform_validation_error("jvp", &source)
1285            .expect("missing semantic rule must map to a public unsupported-rule error");
1286
1287        assert!(matches!(
1288            error,
1289            Error::UnsupportedAdRule { transform: "jvp", ref op }
1290                if op == "tenferro-tests.missing.v1"
1291        ));
1292    }
1293}