Skip to main content

tenferro_df64_proof/
extension.rs

1//! A minimal extension-owned operation for the external scalar.
2//!
3//! The operation is an `ExtensionOp` family whose payload, numerical body, and
4//! output construction all live in this crate. The runtime reaches it through a
5//! registered `ExtensionModule` and prepared execution, and the input travels as
6//! a runtime `Tensor` that carries a caller-owned payload.
7
8use std::any::Any;
9use std::hash::Hasher;
10use std::sync::Arc;
11
12use tenferro_ad::extension::{apply_eager_with_extension_session, ExtensionOp};
13use tenferro_ad::EagerTensor;
14use tenferro_cpu::{scalar_fold, CpuBackend};
15use tenferro_ops::{ExtensionShapeContext, SymDim};
16use tenferro_runtime::{
17    EngineId, ErasedExecutionContext, ExecutionContextIdentity, ExtensionCacheKey,
18    ExtensionCacheStore, ExtensionEngine, ExtensionModule, ExtensionModuleError, ExtensionModuleId,
19    ExtensionModuleRegistrar, ExtensionPlanningConfig, ExtensionPrepareRequest, PrepareCapability,
20    PrepareError, PreparedOperation, PreparedOperationBinding, PreparedOperationExecutor,
21    PreparedOperationExecutorHandle, PreparedOperationHandle, PreparedOperationPlan,
22    SpecializationProjection,
23};
24use tenferro_tensor::{DType, Tensor, TensorRead, TensorView};
25use tenferro_tensor::{DynRank, ErasedHostTensor, Host, TypedTensor};
26use tenferro_tensor_core::Scalar;
27
28use crate::{Df64, Df64Add};
29
30/// Family identifier of the contribution's externally defined operations.
31///
32/// One family holds both the total sum and its adjoint broadcast, because the runtime
33/// keys one planning config per engine and a contribution owns one numerical engine.
34pub const DF64_OPS_FAMILY: &str = "tenferro-df64-proof.df64_ops.v1";
35
36/// Canonical identity of the externally defined `Df64` scalar.
37///
38/// A semantic program's identity must be reproducible across processes, so the
39/// contribution that owns a scalar declares its stable name. This is the name a
40/// program carrying `Df64` values reports instead of a process-local `TypeId`.
41pub const DF64_SCALAR_IDENTITY: &str = "tenferro-df64-proof.df64.v1";
42
43/// Implement the parts of `ExtensionOp` that every payload-free operation shares.
44///
45/// The operation supplies its arity and its own output metadata; the family identity,
46/// the contribution's scalar identity, the empty payload, and the pure, fresh-output
47/// declarations are the same for all of them, so they are declared once here.
48///
49/// # Examples
50///
51/// ```rust
52/// use tenferro_ad::extension::ExtensionOp;
53/// use tenferro_df64_proof::extension::{Df64Total, DF64_OPS_FAMILY, DF64_SCALAR_IDENTITY};
54///
55/// assert_eq!(<Df64Total as ExtensionOp>::family_id(&Df64Total), DF64_OPS_FAMILY);
56/// assert_eq!(<Df64Total as ExtensionOp>::input_count(&Df64Total), 1);
57/// assert_eq!(<Df64Total as ExtensionOp>::output_count(&Df64Total), 1);
58/// assert_eq!(
59///     <Df64Total as ExtensionOp>::scalar_identity(&Df64Total),
60///     Some(DF64_SCALAR_IDENTITY)
61/// );
62/// ```
63macro_rules! df64_operation {
64    ($operation:ty, inputs = $inputs:expr, outputs = $outputs:expr, infer = |$ctx:ident| $infer:block) => {
65        impl ExtensionOp for $operation {
66            fn family_id(&self) -> &'static str {
67                DF64_OPS_FAMILY
68            }
69
70            fn payload_hash(&self, _hasher: &mut dyn Hasher) {}
71
72            fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
73                other.as_any().downcast_ref::<Self>().is_some()
74            }
75
76            fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
77                Arc::new(self.clone())
78            }
79
80            fn as_any(&self) -> &dyn Any {
81                self
82            }
83
84            fn input_count(&self) -> usize {
85                $inputs
86            }
87
88            fn output_count(&self) -> usize {
89                $outputs
90            }
91
92            fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
93                tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
94            }
95
96            fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
97                tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
98            }
99
100            fn scalar_identity(&self) -> Option<&'static str> {
101                Some(DF64_SCALAR_IDENTITY)
102            }
103
104            fn infer_output_meta(
105                &self,
106                $ctx: &mut ExtensionShapeContext<'_>,
107            ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
108                $infer
109            }
110        }
111    };
112}
113
114/// Family identifier of the extension-owned scalar broadcast.
115/// Total sum of an externally defined scalar tensor.
116///
117/// The payload carries no parameters, so every instance is equal to every other.
118///
119/// # Examples
120///
121/// ```rust
122/// use tenferro_df64_proof::extension::Df64Total;
123/// use tenferro_ad::extension::ExtensionOp;
124///
125/// assert_eq!(<Df64Total as ExtensionOp>::input_count(&Df64Total), 1);
126/// ```
127#[derive(Clone, Copy, Debug, Default)]
128pub struct Df64Total;
129
130df64_operation!(
131    Df64Total,
132    inputs = 1,
133    outputs = 1,
134    infer = |ctx| {
135        let dtype = ctx.input_dtype(0)?;
136        if !matches!(dtype, DType::External(_)) {
137            // The body is only defined for the external scalar, so anything else
138            // fails explicitly instead of being coerced.
139            return Err(tenferro_tensor::Error::unsupported_dtype(
140                "df64_total",
141                dtype,
142                "df64_total takes an externally defined scalar",
143            ));
144        }
145        // A total sum has rank zero.
146        Ok(vec![(dtype, Vec::new())])
147    }
148);
149
150/// Broadcast a scalar external value to a declared shape.
151///
152/// The total sum's adjoint needs to place the output cotangent back into the input's
153/// shape, and a preset broadcast is not available for a scalar tenferro does not
154/// declare, so the contribution owns this operation too.
155///
156/// # Examples
157///
158/// ```rust
159/// use tenferro_df64_proof::extension::Df64Expand;
160/// use tenferro_ad::extension::ExtensionOp;
161///
162/// let expand = Df64Expand::new(vec![2, 3]);
163/// assert_eq!(<Df64Expand as ExtensionOp>::family_id(&expand), "tenferro-df64-proof.df64_ops.v1");
164/// assert_eq!(&*expand.shape, &[2, 3]);
165/// ```
166#[derive(Clone, Debug)]
167pub struct Df64Expand {
168    /// Shape the scalar is broadcast to.
169    pub shape: Box<[usize]>,
170}
171
172impl Df64Expand {
173    /// Construct the operation for one output shape.
174    ///
175    /// # Examples
176    ///
177    /// ```rust
178    /// use tenferro_df64_proof::extension::Df64Expand;
179    ///
180    /// let expand = Df64Expand::new(vec![2, 3]);
181    /// assert_eq!(&*expand.shape, &[2, 3]);
182    /// ```
183    #[must_use]
184    pub fn new(shape: Vec<usize>) -> Self {
185        Self {
186            shape: shape.into_boxed_slice(),
187        }
188    }
189}
190
191impl ExtensionOp for Df64Expand {
192    fn family_id(&self) -> &'static str {
193        DF64_OPS_FAMILY
194    }
195
196    fn payload_hash(&self, hasher: &mut dyn Hasher) {
197        for extent in self.shape.iter() {
198            hasher.write_usize(*extent);
199        }
200    }
201
202    fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
203        other
204            .as_any()
205            .downcast_ref::<Self>()
206            .is_some_and(|other| other.shape == self.shape)
207    }
208
209    fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
210        Arc::new(self.clone())
211    }
212
213    fn as_any(&self) -> &dyn Any {
214        self
215    }
216
217    fn input_count(&self) -> usize {
218        1
219    }
220
221    fn output_count(&self) -> usize {
222        1
223    }
224
225    fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
226        tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
227    }
228
229    fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
230        tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
231    }
232
233    fn scalar_identity(&self) -> Option<&'static str> {
234        Some(DF64_SCALAR_IDENTITY)
235    }
236
237    fn infer_output_meta(
238        &self,
239        ctx: &mut ExtensionShapeContext<'_>,
240    ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
241        let dtype = ctx.input_dtype(0)?;
242        if !matches!(dtype, DType::External(_)) {
243            return Err(tenferro_tensor::Error::unsupported_dtype(
244                "df64_expand",
245                dtype,
246                "df64_expand takes an externally defined scalar",
247            ));
248        }
249        Ok(vec![(
250            dtype,
251            self.shape
252                .iter()
253                .map(|extent| SymDim::from(*extent))
254                .collect(),
255        )])
256    }
257}
258
259/// Widen a preset `f64` tensor into the externally defined scalar.
260///
261/// The widening is exact, so the result's low component is zero. A conversion is a
262/// separate operation from the factorization, which is what lets a connected program
263/// start from ordinary `f64` values.
264///
265/// # Examples
266///
267/// ```rust
268/// use tenferro_df64_proof::extension::Df64FromF64;
269/// use tenferro_ad::extension::ExtensionOp;
270///
271/// assert_eq!(<Df64FromF64 as ExtensionOp>::input_count(&Df64FromF64), 1);
272/// ```
273#[derive(Clone, Copy, Debug, Default)]
274pub struct Df64FromF64;
275
276df64_operation!(
277    Df64FromF64,
278    inputs = 1,
279    outputs = 1,
280    infer = |ctx| {
281        let dtype = ctx.input_dtype(0)?;
282        if dtype != DType::F64 {
283            return Err(tenferro_tensor::Error::unsupported_dtype(
284                "df64_from_f64",
285                dtype,
286                "df64_from_f64 takes a preset f64 tensor",
287            ));
288        }
289        Ok(vec![(
290            DType::External(DF64_SCALAR),
291            ctx.input_shape(0)?.to_vec(),
292        )])
293    }
294);
295
296/// Narrow the externally defined scalar into a preset `f64` tensor.
297///
298/// The low component participates and the result is rounded to nearest with ties to
299/// even, so this is not the truncation to the high component that reinterpreting the
300/// payload would give. Information the narrowing discards is not recovered later.
301///
302/// # Examples
303///
304/// ```rust
305/// use tenferro_df64_proof::extension::Df64ToF64;
306/// use tenferro_ad::extension::ExtensionOp;
307///
308/// assert_eq!(<Df64ToF64 as ExtensionOp>::output_count(&Df64ToF64), 1);
309/// ```
310#[derive(Clone, Copy, Debug, Default)]
311pub struct Df64ToF64;
312
313df64_operation!(
314    Df64ToF64,
315    inputs = 1,
316    outputs = 1,
317    infer = |ctx| {
318        let dtype = ctx.input_dtype(0)?;
319        if !matches!(dtype, DType::External(_)) {
320            return Err(tenferro_tensor::Error::unsupported_dtype(
321                "df64_to_f64",
322                dtype,
323                "df64_to_f64 takes an externally defined scalar",
324            ));
325        }
326        Ok(vec![(DType::F64, ctx.input_shape(0)?.to_vec())])
327    }
328);
329
330/// Reverse-mode adjoint of the reduced QR factorization.
331///
332/// The adjoint is a numerical body of its own: it needs a triangular solve against the
333/// primal factor, so it is an operation rather than a graph of preset operations, which
334/// a scalar tenferro does not declare could not execute anyway.
335///
336/// # Examples
337///
338/// ```rust
339/// use tenferro_df64_proof::extension::Df64QrVjp;
340/// use tenferro_ad::extension::ExtensionOp;
341///
342/// // A loss that depends on the triangular factor alone supplies one cotangent.
343/// let adjoint = Df64QrVjp::of(false, true);
344/// assert_eq!(<Df64QrVjp as ExtensionOp>::input_count(&adjoint), 3);
345/// assert_eq!(<Df64QrVjp as ExtensionOp>::output_count(&adjoint), 1);
346/// assert_eq!(
347///     <Df64QrVjp as ExtensionOp>::input_count(&Df64QrVjp::of(true, true)),
348///     4
349/// );
350/// ```
351#[derive(Clone, Copy, Debug, Default)]
352pub struct Df64QrVjp {
353    /// Whether the caller supplied the factor's cotangent.
354    pub has_q: bool,
355    /// Whether the caller supplied the triangular factor's cotangent.
356    pub has_r: bool,
357}
358
359impl Df64QrVjp {
360    /// Construct the adjoint for one cotangent availability.
361    ///
362    /// # Examples
363    ///
364    /// ```rust
365    /// use tenferro_df64_proof::extension::Df64QrVjp;
366    ///
367    /// // Only the triangular factor carries a cotangent here.
368    /// let adjoint = Df64QrVjp::of(false, true);
369    /// assert!(!adjoint.has_q);
370    /// assert!(adjoint.has_r);
371    /// ```
372    #[must_use]
373    pub const fn of(has_q: bool, has_r: bool) -> Self {
374        Self { has_q, has_r }
375    }
376}
377
378impl ExtensionOp for Df64QrVjp {
379    fn family_id(&self) -> &'static str {
380        DF64_OPS_FAMILY
381    }
382
383    fn payload_hash(&self, hasher: &mut dyn Hasher) {
384        hasher.write_u8(u8::from(self.has_q) | (u8::from(self.has_r) << 1));
385    }
386
387    fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
388        other
389            .as_any()
390            .downcast_ref::<Self>()
391            .is_some_and(|other| other.has_q == self.has_q && other.has_r == self.has_r)
392    }
393
394    fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
395        Arc::new(*self)
396    }
397
398    fn as_any(&self) -> &dyn Any {
399        self
400    }
401
402    fn input_count(&self) -> usize {
403        // The primal factors, plus one cotangent per available output.
404        2 + usize::from(self.has_q) + usize::from(self.has_r)
405    }
406
407    fn output_count(&self) -> usize {
408        1
409    }
410
411    fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
412        tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
413    }
414
415    fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
416        tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
417    }
418
419    fn scalar_identity(&self) -> Option<&'static str> {
420        Some(DF64_SCALAR_IDENTITY)
421    }
422
423    fn infer_output_meta(
424        &self,
425        ctx: &mut ExtensionShapeContext<'_>,
426    ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
427        let dtype = ctx.input_dtype(0)?;
428        if !matches!(dtype, DType::External(_)) {
429            return Err(tenferro_tensor::Error::unsupported_dtype(
430                "df64_qr_vjp",
431                dtype,
432                "df64_qr_vjp takes an externally defined scalar",
433            ));
434        }
435        Ok(vec![(dtype, ctx.input_shape(0)?.to_vec())])
436    }
437}
438
439/// A matrix contraction in the external scalar, written in ordinary einsum notation.
440///
441/// #1793's example is `einsum("ik,kj->ij", A, B)` evaluated in the external scalar, where the
442/// contraction of `[1, 1]` with `[1, 2^-80]` has to keep the low component an `f64` accumulator
443/// would drop. The operation accepts exactly that pattern: two rank-2 inputs that share one
444/// contracted label, and an output of the two free labels. Any other pattern is refused with a
445/// typed error rather than approximated, because the general label cases need the diagonal,
446/// reduction, and permutation stages the ordinary lowering plans and this body does not execute.
447///
448/// # Examples
449///
450/// ```rust
451/// use tenferro_ad::extension::ExtensionOp;
452/// use tenferro_df64_proof::extension::Df64Einsum;
453///
454/// let op = Df64Einsum::new(&[0, 1], &[1, 2], &[0, 2]).expect("a matrix contraction");
455/// assert_eq!(<Df64Einsum as ExtensionOp>::input_count(&op), 2);
456/// assert_eq!(<Df64Einsum as ExtensionOp>::output_count(&op), 1);
457/// ```
458#[derive(Clone, Debug, PartialEq, Eq)]
459pub struct Df64Einsum {
460    inputs: Vec<Vec<u32>>,
461    out: Vec<u32>,
462}
463
464impl Df64Einsum {
465    /// Build the matrix-contraction pattern `lhs,rhs->out`.
466    ///
467    /// The labels are the ones an ordinary einsum subscript string names, in order, so
468    /// `"ik,kj->ij"` is `(&[0, 1], &[1, 2], &[0, 2])`.
469    ///
470    /// # Errors
471    ///
472    /// Returns [`tenferro_tensor::Error::InvalidArgument`] when either input is not rank two, when
473    /// a label repeats within one input, when the inputs do not share exactly one contracted label
474    /// as the second and first label respectively, or when the output is not the two free labels in
475    /// that order.
476    ///
477    /// # Examples
478    ///
479    /// ```rust
480    /// use tenferro_df64_proof::extension::Df64Einsum;
481    ///
482    /// assert!(Df64Einsum::new(&[0, 1], &[1, 2], &[0, 2]).is_ok());
483    /// // A repeated label inside one input is a trace, which the body evaluates.
484    /// assert!(Df64Einsum::new(&[0, 0], &[0, 2], &[0, 2]).is_ok());
485    /// ```
486    pub fn new(lhs: &[u32], rhs: &[u32], out: &[u32]) -> tenferro_runtime::Result<Self> {
487        Self::new_nary(&[lhs, rhs], out)
488    }
489
490    /// Build the pattern for any number of operands.
491    ///
492    /// A label that two operands share and the output omits is contracted; a label the output omits
493    /// is summed; a label that repeats inside one operand is a trace or a diagonal extraction. The
494    /// operands are contracted from the left in the order given, and an intermediate keeps exactly
495    /// the labels the remaining operands or the output still need.
496    ///
497    /// # Errors
498    ///
499    /// Returns an error when fewer than two operands are given, when one carries no label, or when
500    /// an output label appears in no operand.
501    ///
502    /// # Examples
503    ///
504    /// ```rust
505    /// use tenferro_df64_proof::extension::Df64Einsum;
506    ///
507    /// assert!(Df64Einsum::new_nary(&[&[0, 1], &[1, 2], &[2, 3]], &[0, 3]).is_ok());
508    /// assert!(Df64Einsum::new_nary(&[&[0, 1]], &[0]).is_err());
509    /// ```
510    pub fn new_nary(inputs: &[&[u32]], out: &[u32]) -> tenferro_runtime::Result<Self> {
511        let invalid = |message: &str| {
512            tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
513                "df64_einsum",
514                "pattern",
515                message,
516            ))
517        };
518        if inputs.len() < 2 {
519            return Err(invalid("a contraction takes at least two operands"));
520        }
521        if inputs.iter().any(|labels| labels.is_empty()) {
522            return Err(invalid("an operand must carry at least one label"));
523        }
524        for label in out {
525            if !inputs.iter().any(|labels| labels.contains(label)) {
526                return Err(invalid(
527                    "an output label must appear in at least one operand",
528                ));
529            }
530        }
531        let mut seen = out.to_vec();
532        seen.sort_unstable();
533        seen.dedup();
534        if seen.len() != out.len() {
535            return Err(invalid("an output label repeats"));
536        }
537        Ok(Self {
538            inputs: inputs.iter().map(|labels| labels.to_vec()).collect(),
539            out: out.to_vec(),
540        })
541    }
542
543    /// The pattern's label lists, when it has exactly two operands.
544    ///
545    /// The adjoint and tangent helpers are defined for the pairwise case, so they ask for this and
546    /// refuse anything wider rather than guessing.
547    ///
548    /// # Examples
549    ///
550    /// ```rust
551    /// use tenferro_df64_proof::extension::Df64Einsum;
552    ///
553    /// let op = Df64Einsum::new(&[0, 1], &[1, 2], &[0, 2]).expect("a contraction");
554    /// assert!(op.labels().is_some());
555    /// let wide = Df64Einsum::new_nary(&[&[0, 1], &[1, 2], &[2, 3]], &[0, 3]).expect("a contraction");
556    /// assert!(wide.labels().is_none());
557    /// ```
558    #[must_use]
559    pub fn labels(&self) -> Option<(&[u32], &[u32], &[u32])> {
560        match self.inputs.as_slice() {
561            [lhs, rhs] => Some((lhs, rhs, &self.out)),
562            _ => None,
563        }
564    }
565
566    /// Every operand's labels, in operand order.
567    ///
568    /// # Examples
569    ///
570    /// ```rust
571    /// use tenferro_df64_proof::extension::Df64Einsum;
572    ///
573    /// let op = Df64Einsum::new(&[0, 1], &[1, 2], &[0, 2]).expect("a contraction");
574    /// assert_eq!(op.input_labels(), &[vec![0, 1], vec![1, 2]]);
575    /// ```
576    #[must_use]
577    pub fn input_labels(&self) -> &[Vec<u32>] {
578        &self.inputs
579    }
580
581    /// The output's labels, in the output's axis order.
582    ///
583    /// # Examples
584    ///
585    /// ```rust
586    /// use tenferro_df64_proof::extension::Df64Einsum;
587    ///
588    /// let op = Df64Einsum::new(&[0, 1], &[1, 2], &[0, 2]).expect("a contraction");
589    /// assert_eq!(op.out_labels(), &[0, 2]);
590    /// ```
591    #[must_use]
592    pub fn out_labels(&self) -> &[u32] {
593        &self.out
594    }
595}
596
597impl ExtensionOp for Df64Einsum {
598    fn family_id(&self) -> &'static str {
599        DF64_OPS_FAMILY
600    }
601
602    fn payload_hash(&self, hasher: &mut dyn Hasher) {
603        hasher.write_usize(self.inputs.len());
604        for labels in self.inputs.iter().chain(core::iter::once(&self.out)) {
605            hasher.write_usize(labels.len());
606            for label in labels {
607                hasher.write_u32(*label);
608            }
609        }
610    }
611
612    fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
613        other
614            .as_any()
615            .downcast_ref::<Self>()
616            .is_some_and(|other| other == self)
617    }
618
619    fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
620        Arc::new(self.clone())
621    }
622
623    fn as_any(&self) -> &dyn Any {
624        self
625    }
626
627    fn input_count(&self) -> usize {
628        self.inputs.len()
629    }
630
631    fn output_count(&self) -> usize {
632        1
633    }
634
635    fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
636        tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
637    }
638
639    fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
640        tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
641    }
642
643    fn scalar_identity(&self) -> Option<&'static str> {
644        Some(DF64_SCALAR_IDENTITY)
645    }
646
647    fn infer_output_meta(
648        &self,
649        ctx: &mut ExtensionShapeContext<'_>,
650    ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
651        let dtype = ctx.input_dtype(0)?;
652        if !matches!(dtype, DType::External(_)) {
653            return Err(tenferro_tensor::Error::unsupported_dtype(
654                "df64_einsum",
655                dtype,
656                "df64_einsum takes an externally defined scalar",
657            ));
658        }
659        // The output's extent for a label is the extent the first operand that names it declares.
660        // Whether operands agree on a shared label is a value-level question, so the body checks it
661        // at execution rather than the metadata layer guessing.
662        let mut out_shape = Vec::with_capacity(self.out.len());
663        for label in &self.out {
664            let mut extent = None;
665            for (operand, labels) in self.inputs.iter().enumerate() {
666                if ctx.input_dtype(operand)? != dtype {
667                    return Err(tenferro_tensor::Error::invalid_argument(
668                        "df64_einsum",
669                        "inputs",
670                        "every operand must carry the same scalar",
671                    ));
672                }
673                if let Some(axis) = labels.iter().position(|candidate| candidate == label) {
674                    let shape = ctx.input_shape(operand)?;
675                    if shape.len() != labels.len() {
676                        return Err(tenferro_tensor::Error::rank_mismatch(
677                            "df64_einsum",
678                            labels.len(),
679                            shape.len(),
680                        ));
681                    }
682                    extent = Some(shape[axis].clone());
683                    break;
684                }
685            }
686            out_shape.push(extent.ok_or_else(|| {
687                tenferro_tensor::Error::invalid_argument(
688                    "df64_einsum",
689                    "pattern",
690                    "an output label must appear in at least one operand",
691                )
692            })?);
693        }
694        Ok(vec![(dtype, out_shape)])
695    }
696}
697
698/// The adjoint of a two-input contraction.
699///
700/// The adjoint of `out = einsum(lhs, rhs)` contracts the output cotangent with the other operand,
701/// which is the same operation with the labels rotated: `lhs_bar = einsum(out, rhs -> lhs)` and
702/// `rhs_bar = einsum(lhs, out -> rhs)`. It carries the pattern so the adjoint uses exactly the
703/// labels the primal used, and it has two outputs because the contraction has two inputs.
704///
705/// # Examples
706///
707/// ```rust
708/// use tenferro_ad::extension::ExtensionOp;
709/// use tenferro_df64_proof::extension::Df64EinsumVjp;
710///
711/// let adjoint = Df64EinsumVjp::of(&[&[0, 1], &[1, 2]], &[0, 2]).expect("a contraction");
712/// assert_eq!(<Df64EinsumVjp as ExtensionOp>::input_count(&adjoint), 3);
713/// assert_eq!(<Df64EinsumVjp as ExtensionOp>::output_count(&adjoint), 2);
714/// ```
715#[derive(Clone, Debug, PartialEq, Eq)]
716pub struct Df64EinsumVjp {
717    inputs: Vec<Vec<u32>>,
718    out: Vec<u32>,
719}
720
721impl Df64EinsumVjp {
722    /// Build the adjoint for the same labels as the primal contraction.
723    ///
724    /// # Errors
725    ///
726    /// Returns [`tenferro_tensor::Error::InvalidArgument`] when the pattern is not a pairwise contraction: an operand
727    /// carries no label, a label repeats inside one operand, the operands share no contracted
728    /// label, or an output label appears in no operand.
729    ///
730    /// # Examples
731    ///
732    /// ```rust
733    /// use tenferro_df64_proof::extension::Df64EinsumVjp;
734    ///
735    /// assert!(Df64EinsumVjp::of(&[&[0, 1], &[1, 2]], &[0, 2]).is_ok());
736    /// assert!(Df64EinsumVjp::of(&[&[], &[1, 2]], &[0, 2]).is_err());
737    /// ```
738    pub fn of(inputs: &[&[u32]], out: &[u32]) -> tenferro_runtime::Result<Self> {
739        Df64Einsum::new_nary(inputs, out)?;
740        Ok(Self {
741            inputs: inputs.iter().map(|labels| labels.to_vec()).collect(),
742            out: out.to_vec(),
743        })
744    }
745
746    /// The primal pattern's label lists.
747    ///
748    /// Every operand's labels, in operand order.
749    ///
750    /// # Examples
751    ///
752    /// ```rust
753    /// use tenferro_df64_proof::extension::Df64EinsumVjp;
754    ///
755    /// let adjoint = Df64EinsumVjp::of(&[&[0, 1], &[1, 2]], &[0, 2]).expect("a contraction");
756    /// assert_eq!(adjoint.input_labels(), &[vec![0, 1], vec![1, 2]]);
757    /// ```
758    #[must_use]
759    pub fn input_labels(&self) -> &[Vec<u32>] {
760        &self.inputs
761    }
762
763    /// The output's labels, in the output's axis order.
764    ///
765    /// # Examples
766    ///
767    /// ```rust
768    /// use tenferro_df64_proof::extension::Df64EinsumVjp;
769    ///
770    /// let adjoint = Df64EinsumVjp::of(&[&[0, 1], &[1, 2]], &[0, 2]).expect("a contraction");
771    /// assert_eq!(adjoint.out_labels(), &[0, 2]);
772    /// ```
773    #[must_use]
774    pub fn out_labels(&self) -> &[u32] {
775        &self.out
776    }
777}
778
779impl ExtensionOp for Df64EinsumVjp {
780    fn family_id(&self) -> &'static str {
781        DF64_OPS_FAMILY
782    }
783
784    fn payload_hash(&self, hasher: &mut dyn Hasher) {
785        hasher.write_usize(self.inputs.len());
786        for labels in self.inputs.iter().chain(core::iter::once(&self.out)) {
787            hasher.write_usize(labels.len());
788            for label in labels {
789                hasher.write_u32(*label);
790            }
791        }
792    }
793
794    fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
795        other
796            .as_any()
797            .downcast_ref::<Self>()
798            .is_some_and(|other| other == self)
799    }
800
801    fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
802        Arc::new(self.clone())
803    }
804
805    fn as_any(&self) -> &dyn Any {
806        self
807    }
808
809    fn input_count(&self) -> usize {
810        // Every operand and the output cotangent.
811        self.inputs.len() + 1
812    }
813
814    fn output_count(&self) -> usize {
815        // One cotangent per operand.
816        self.inputs.len()
817    }
818
819    fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
820        tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
821    }
822
823    fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
824        tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
825    }
826
827    fn scalar_identity(&self) -> Option<&'static str> {
828        Some(DF64_SCALAR_IDENTITY)
829    }
830
831    fn infer_output_meta(
832        &self,
833        ctx: &mut ExtensionShapeContext<'_>,
834    ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
835        let dtype = ctx.input_dtype(0)?;
836        if !matches!(dtype, DType::External(_)) {
837            return Err(tenferro_tensor::Error::unsupported_dtype(
838                "df64_einsum_vjp",
839                dtype,
840                "df64_einsum_vjp takes an externally defined scalar",
841            ));
842        }
843        let mut shapes = Vec::with_capacity(self.inputs.len());
844        for operand in 0..self.inputs.len() {
845            shapes.push((dtype, ctx.input_shape(operand)?.to_vec()));
846        }
847        Ok(shapes)
848    }
849}
850
851/// The forward tangent of a two-input contraction.
852///
853/// The tangent of `out = einsum(lhs, rhs)` is `einsum(lhs_dot, rhs) + einsum(lhs, rhs_dot)`, so the
854/// helper contracts each tangent with the other operand and adds the two results in the extended
855/// scalar. Its tangent availability is a payload field, because a linearization need not have a
856/// tangent for both operands, and the rule must not materialise a zero tangent for one that is
857/// absent.
858///
859/// # Examples
860///
861/// ```rust
862/// use tenferro_ad::extension::ExtensionOp;
863/// use tenferro_df64_proof::extension::Df64EinsumJvp;
864///
865/// let tangent = Df64EinsumJvp::of(&[&[0, 1], &[1, 2]], &[0, 2], &[true, false]).expect("a contraction");
866/// assert_eq!(<Df64EinsumJvp as ExtensionOp>::input_count(&tangent), 3);
867/// assert_eq!(<Df64EinsumJvp as ExtensionOp>::output_count(&tangent), 1);
868/// ```
869#[derive(Clone, Debug, PartialEq, Eq)]
870pub struct Df64EinsumJvp {
871    inputs: Vec<Vec<u32>>,
872    out: Vec<u32>,
873    tangents: Vec<bool>,
874}
875
876impl Df64EinsumJvp {
877    /// Build the tangent for a pattern and one tangent availability.
878    ///
879    /// # Errors
880    ///
881    /// Returns an error when the operand list is not a valid pattern (an operand carries no label, or
882    /// an output label appears in no operand), when the tangent mask does not have one entry per
883    /// operand, or when no operand carries a tangent, because then there is nothing to differentiate.
884    ///
885    /// # Examples
886    ///
887    /// ```rust
888    /// use tenferro_df64_proof::extension::Df64EinsumJvp;
889    ///
890    /// assert!(Df64EinsumJvp::of(&[&[0, 1], &[1, 2]], &[0, 2], &[true, true]).is_ok());
891    /// assert!(Df64EinsumJvp::of(&[&[0, 1], &[1, 2]], &[0, 2], &[false, false]).is_err());
892    /// ```
893    pub fn of(inputs: &[&[u32]], out: &[u32], tangents: &[bool]) -> tenferro_runtime::Result<Self> {
894        Df64Einsum::new_nary(inputs, out)?;
895        if tangents.len() != inputs.len() {
896            return Err(tenferro_runtime::Error::from(
897                tenferro_tensor::Error::invalid_argument(
898                    "df64_einsum_jvp",
899                    "tangents",
900                    "the tangent mask needs one entry per operand",
901                ),
902            ));
903        }
904        if !tangents.iter().any(|present| *present) {
905            return Err(tenferro_runtime::Error::from(
906                tenferro_tensor::Error::invalid_argument(
907                    "df64_einsum_jvp",
908                    "tangents",
909                    "at least one operand must carry a tangent",
910                ),
911            ));
912        }
913        Ok(Self {
914            inputs: inputs.iter().map(|labels| labels.to_vec()).collect(),
915            out: out.to_vec(),
916            tangents: tangents.to_vec(),
917        })
918    }
919
920    /// The tangent's pattern and the availability of each operand's tangent.
921    ///
922    /// # Examples
923    ///
924    /// ```rust
925    /// use tenferro_df64_proof::extension::Df64EinsumJvp;
926    ///
927    /// let tangent = Df64EinsumJvp::of(&[&[0, 1], &[1, 2]], &[0, 2], &[true, true]).expect("a tangent");
928    /// assert_eq!(tangent.tangents(), &[true, true]);
929    /// ```
930    #[must_use]
931    pub fn input_labels(&self) -> &[Vec<u32>] {
932        &self.inputs
933    }
934
935    /// The output's labels, in the output's axis order.
936    ///
937    /// # Examples
938    ///
939    /// ```rust
940    /// use tenferro_df64_proof::extension::Df64EinsumJvp;
941    ///
942    /// let tangent = Df64EinsumJvp::of(&[&[0, 1], &[1, 2]], &[0, 2], &[true, false])
943    ///     .expect("a tangent");
944    /// assert_eq!(tangent.out_labels(), &[0, 2]);
945    /// ```
946    #[must_use]
947    pub fn out_labels(&self) -> &[u32] {
948        &self.out
949    }
950
951    /// Which operands carry a tangent, in operand order.
952    ///
953    /// # Examples
954    ///
955    /// ```rust
956    /// use tenferro_df64_proof::extension::Df64EinsumJvp;
957    ///
958    /// let tangent = Df64EinsumJvp::of(&[&[0, 1], &[1, 2]], &[0, 2], &[true, false])
959    ///     .expect("a tangent");
960    /// assert_eq!(tangent.tangents(), &[true, false]);
961    /// ```
962    #[must_use]
963    pub fn tangents(&self) -> &[bool] {
964        &self.tangents
965    }
966}
967
968impl ExtensionOp for Df64EinsumJvp {
969    fn family_id(&self) -> &'static str {
970        DF64_OPS_FAMILY
971    }
972
973    fn payload_hash(&self, hasher: &mut dyn Hasher) {
974        hasher.write_usize(self.inputs.len());
975        for labels in self.inputs.iter().chain(core::iter::once(&self.out)) {
976            hasher.write_usize(labels.len());
977            for label in labels {
978                hasher.write_u32(*label);
979            }
980        }
981        hasher.write_usize(self.tangents.len());
982        for present in &self.tangents {
983            hasher.write_u8(u8::from(*present));
984        }
985    }
986
987    fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
988        other
989            .as_any()
990            .downcast_ref::<Self>()
991            .is_some_and(|other| other == self)
992    }
993
994    fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
995        Arc::new(self.clone())
996    }
997
998    fn as_any(&self) -> &dyn Any {
999        self
1000    }
1001
1002    fn input_count(&self) -> usize {
1003        // Every operand, plus one tangent per operand that carries one.
1004        self.inputs.len() + self.tangents.iter().filter(|present| **present).count()
1005    }
1006
1007    fn output_count(&self) -> usize {
1008        1
1009    }
1010
1011    fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
1012        tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
1013    }
1014
1015    fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
1016        tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
1017    }
1018
1019    fn scalar_identity(&self) -> Option<&'static str> {
1020        Some(DF64_SCALAR_IDENTITY)
1021    }
1022
1023    fn infer_output_meta(
1024        &self,
1025        ctx: &mut ExtensionShapeContext<'_>,
1026    ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
1027        let dtype = ctx.input_dtype(0)?;
1028        if !matches!(dtype, DType::External(_)) {
1029            return Err(tenferro_tensor::Error::unsupported_dtype(
1030                "df64_einsum_jvp",
1031                dtype,
1032                "df64_einsum_jvp takes an externally defined scalar",
1033            ));
1034        }
1035        let shapes: Vec<Vec<SymDim>> = (0..self.inputs.len())
1036            .map(|operand| {
1037                let shape = ctx.input_shape(operand)?;
1038                if shape.len() != self.inputs[operand].len() {
1039                    return Err(tenferro_tensor::Error::rank_mismatch(
1040                        "df64_einsum_jvp",
1041                        self.inputs[operand].len(),
1042                        shape.len(),
1043                    ));
1044                }
1045                Ok(shape.to_vec())
1046            })
1047            .collect::<tenferro_tensor::Result<Vec<_>>>()?;
1048        let mut out_shape = Vec::with_capacity(self.out.len());
1049        for label in &self.out {
1050            let mut extent = None;
1051            for (operand, labels) in self.inputs.iter().enumerate() {
1052                if let Some(axis) = labels.iter().position(|candidate| candidate == label) {
1053                    extent = Some(shapes[operand][axis].clone());
1054                    break;
1055                }
1056            }
1057            out_shape.push(extent.ok_or_else(|| {
1058                tenferro_tensor::Error::invalid_argument(
1059                    "df64_einsum_jvp",
1060                    "pattern",
1061                    "an output label must appear in at least one operand",
1062                )
1063            })?);
1064        }
1065        Ok(vec![(dtype, out_shape)])
1066    }
1067}
1068
1069/// Forward-mode tangent of the reduced QR factorization.
1070///
1071/// # Examples
1072///
1073/// ```rust
1074/// use tenferro_df64_proof::extension::Df64QrJvp;
1075/// use tenferro_ad::extension::ExtensionOp;
1076///
1077/// assert_eq!(<Df64QrJvp as ExtensionOp>::input_count(&Df64QrJvp), 3);
1078/// assert_eq!(<Df64QrJvp as ExtensionOp>::output_count(&Df64QrJvp), 2);
1079/// ```
1080#[derive(Clone, Copy, Debug, Default)]
1081pub struct Df64QrJvp;
1082
1083df64_operation!(
1084    Df64QrJvp,
1085    inputs = 3,
1086    outputs = 2,
1087    infer = |ctx| {
1088        let dtype = ctx.input_dtype(0)?;
1089        if !matches!(dtype, DType::External(_)) {
1090            return Err(tenferro_tensor::Error::unsupported_dtype(
1091                "df64_qr_jvp",
1092                dtype,
1093                "df64_qr_jvp takes an externally defined scalar",
1094            ));
1095        }
1096        Ok(vec![
1097            (dtype, ctx.input_shape(0)?.to_vec()),
1098            (dtype, ctx.input_shape(1)?.to_vec()),
1099        ])
1100    }
1101);
1102
1103/// The externally defined element type of the contribution's scalar.
1104///
1105/// This is what the runtime tag reports, next to the canonical identity.
1106pub const DF64_SCALAR: std::any::TypeId = std::any::TypeId::of::<Df64>();
1107
1108/// Reduced QR factorization of a real square or tall matrix.
1109///
1110/// The factorization is the contribution's own numerical body: modified
1111/// Gram-Schmidt with one re-orthogonalization pass, computed in the external scalar
1112/// so the factors keep its precision. `R` has a positive diagonal and `Q` has
1113/// orthonormal columns, returned as two externally defined tensors.
1114///
1115/// # Examples
1116///
1117/// ```rust
1118/// use tenferro_df64_proof::extension::Df64Qr;
1119/// use tenferro_ad::extension::ExtensionOp;
1120///
1121/// assert_eq!(<Df64Qr as ExtensionOp>::input_count(&Df64Qr), 1);
1122/// assert_eq!(<Df64Qr as ExtensionOp>::output_count(&Df64Qr), 2);
1123/// ```
1124#[derive(Clone, Copy, Debug, Default)]
1125pub struct Df64Qr;
1126
1127df64_operation!(
1128    Df64Qr,
1129    inputs = 1,
1130    outputs = 2,
1131    infer = |ctx| {
1132        let dtype = ctx.input_dtype(0)?;
1133        if !matches!(dtype, DType::External(_)) {
1134            return Err(tenferro_tensor::Error::unsupported_dtype(
1135                "df64_qr",
1136                dtype,
1137                "df64_qr takes an externally defined scalar",
1138            ));
1139        }
1140        // The factors have the input's extents, so the inference forwards them however
1141        // they are expressed. Whether the matrix is tall enough is a property of the
1142        // concrete input and is checked when the body runs.
1143        let shape = ctx.input_shape(0)?.to_vec();
1144        let [rows, columns] = match shape.as_slice() {
1145            [rows, columns] => [rows.clone(), columns.clone()],
1146            _ => {
1147                return Err(tenferro_tensor::Error::invalid_argument(
1148                    "df64_qr",
1149                    "input",
1150                    "df64_qr takes a rank-2 matrix",
1151                ));
1152            }
1153        };
1154        Ok(vec![
1155            (dtype, vec![rows.clone(), columns.clone()]),
1156            (dtype, vec![columns.clone(), columns]),
1157        ])
1158    }
1159);
1160
1161/// Reduced QR factorization of a column-major dense matrix in the external scalar.
1162///
1163/// # Errors
1164///
1165/// Returns a typed error when the input is not a rank-2 externally defined matrix, or
1166/// when a column is zero and the factorization has no unit vector for it.
1167fn qr_of(
1168    session: Option<&mut dyn tenferro_tensor::BackendSession>,
1169    inputs: &[TensorRead<'_>],
1170) -> tenferro_runtime::Result<Vec<Tensor>> {
1171    let input = sole_input("df64_qr", session, inputs)?;
1172    let tensor = input.tensor();
1173    let payload =
1174        external_payload::<Df64>("df64_qr", tensor).map_err(tenferro_runtime::Error::from)?;
1175    let shape = tensor.shape();
1176    let invalid = |message: &'static str| {
1177        tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
1178            "df64_qr", "input", message,
1179        ))
1180    };
1181    let [rows, columns] = match shape {
1182        [rows, columns] if *rows >= *columns => [*rows, *columns],
1183        _ => {
1184            return Err(invalid(
1185                "df64_qr takes a rank-2 matrix with at least as many rows as columns",
1186            ));
1187        }
1188    };
1189    let values = payload.as_slice();
1190    if values.len() != rows * columns {
1191        return Err(invalid("df64_qr takes a dense column-major matrix"));
1192    }
1193
1194    // Column-major access into the source matrix.
1195    let source = |row: usize, column: usize| values[row + column * rows];
1196    let mut q = vec![Df64::zero(); rows * columns];
1197    let mut r = vec![Df64::zero(); columns * columns];
1198    for column in 0..columns {
1199        for row in 0..rows {
1200            q[row + column * rows] = source(row, column);
1201        }
1202        // One re-orthogonalization pass after the first projection, so the columns stay
1203        // orthogonal to working precision even when the input is close to rank
1204        // deficient.
1205        for _pass in 0..2 {
1206            for previous in 0..column {
1207                let mut projection = Df64::zero();
1208                for row in 0..rows {
1209                    projection = projection + q[row + previous * rows] * q[row + column * rows];
1210                }
1211                for row in 0..rows {
1212                    q[row + column * rows] =
1213                        q[row + column * rows] - projection * q[row + previous * rows];
1214                }
1215                r[previous + column * columns] = r[previous + column * columns] + projection;
1216            }
1217        }
1218        let mut squares = Df64::zero();
1219        for row in 0..rows {
1220            squares = squares + q[row + column * rows] * q[row + column * rows];
1221        }
1222        let norm = squares.sqrt();
1223        if norm.hi == 0.0 {
1224            return Err(invalid("df64_qr takes a matrix with no zero column"));
1225        }
1226        // A positive diagonal is part of the factorization's contract, so a negative
1227        // norm flips the column and the entry together.
1228        let sign = if norm.hi < 0.0 {
1229            Df64::from_f64(-1.0)
1230        } else {
1231            Df64::from_f64(1.0)
1232        };
1233        let scale = sign * Df64::from_f64(1.0).ratio(norm);
1234        for row in 0..rows {
1235            q[row + column * rows] = q[row + column * rows] * scale;
1236        }
1237        r[column + column * columns] = sign * norm;
1238    }
1239
1240    let q = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![rows, columns], q)
1241        .map_err(tenferro_runtime::Error::from)?;
1242    let r = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![columns, columns], r)
1243        .map_err(tenferro_runtime::Error::from)?;
1244    Ok(vec![
1245        Tensor::external(ErasedHostTensor::new(q)),
1246        Tensor::external(ErasedHostTensor::new(r)),
1247    ])
1248}
1249
1250/// Narrow the external scalar to `f64` through the contribution's own conversion.
1251fn to_f64_of(
1252    session: Option<&mut dyn tenferro_tensor::BackendSession>,
1253    inputs: &[TensorRead<'_>],
1254) -> tenferro_runtime::Result<Vec<Tensor>> {
1255    let input = sole_input("df64_to_f64", session, inputs)?;
1256    let tensor = input.tensor();
1257    crate::conversion::to_f64(tensor)
1258        .map(|tensor| vec![tensor])
1259        .map_err(tenferro_runtime::Error::from)
1260}
1261
1262/// Widen a preset `f64` tensor into the external scalar.
1263fn from_f64_of(
1264    session: Option<&mut dyn tenferro_tensor::BackendSession>,
1265    inputs: &[TensorRead<'_>],
1266) -> tenferro_runtime::Result<Vec<Tensor>> {
1267    let input = match inputs.first() {
1268        // A borrowed read of the preset input is gathered by its own layout rather than
1269        // requiring a session the executor may not have.
1270        Some(read @ TensorRead::View(_)) if session.is_none() => {
1271            Input::Materialized(Box::new(owned_f64_read(read)?))
1272        }
1273        _ => sole_input("df64_from_f64", session, inputs)?,
1274    };
1275    let tensor = input.tensor();
1276    crate::conversion::to_df64(tensor)
1277        .map(|tensor| vec![tensor])
1278        .map_err(tenferro_runtime::Error::from)
1279}
1280
1281fn external_payload<'a, T: Scalar>(
1282    op: &'static str,
1283    tensor: &'a Tensor,
1284) -> tenferro_tensor::Result<&'a TypedTensor<T, DynRank, Host>> {
1285    match tensor.external_payload() {
1286        Some(payload) => payload.downcast_ref::<T>().ok_or_else(|| {
1287            tenferro_tensor::Error::unsupported_dtype(
1288                op,
1289                tensor.dtype(),
1290                "the external payload does not hold the expected scalar",
1291            )
1292        }),
1293        None => Err(tenferro_tensor::Error::unsupported_dtype(
1294            op,
1295            tensor.dtype(),
1296            "the operation takes an externally defined payload",
1297        )),
1298    }
1299}
1300
1301/// Reusable scratch the numerical bodies own between executions.
1302///
1303/// The entry lives in the runtime's accounted extension cache, so its retained bytes are
1304/// reported rather than hidden: this is the contribution's acquisition and return path for a
1305/// buffer of its own element type, and the cache's statistics are the evidence for it. The
1306/// buffers are named fields rather than a list so one can be written while another is read,
1307/// which the adjoint's chain of products needs.
1308#[derive(Debug, Default)]
1309struct Scratch {
1310    /// The transposed factor cotangent.
1311    q_bar_t: Vec<Df64>,
1312    /// The transposed triangular cotangent.
1313    r_bar_t: Vec<Df64>,
1314    /// The `Q_bar^T Q` product.
1315    q_bar_t_q: Vec<Df64>,
1316    /// The `R R_bar^T` product.
1317    r_r_bar: Vec<Df64>,
1318    /// The `R R_bar^T - Q_bar^T Q` difference.
1319    m: Vec<Df64>,
1320    /// `copyltu(M)`, built in place.
1321    s: Vec<Df64>,
1322    /// The `Q S` product.
1323    product: Vec<Df64>,
1324    /// The accumulator `Q_bar + Q S`.
1325    b: Vec<Df64>,
1326}
1327
1328impl Scratch {
1329    /// The cache namespace these buffers live under.
1330    const CACHE_NAME: &'static str = "scratch";
1331
1332    /// Every buffer this entry retains.
1333    fn buffers(&self) -> [&Vec<Df64>; 8] {
1334        [
1335            &self.q_bar_t,
1336            &self.r_bar_t,
1337            &self.q_bar_t_q,
1338            &self.r_r_bar,
1339            &self.m,
1340            &self.s,
1341            &self.product,
1342            &self.b,
1343        ]
1344    }
1345
1346    /// Resize one buffer and hand it back for writing, zero-filled.
1347    fn slot(buffer: &mut Vec<Df64>, length: usize) -> &mut [Df64] {
1348        buffer.clear();
1349        buffer.resize(length, Df64::zero());
1350        buffer.as_mut_slice()
1351    }
1352
1353    /// Bytes this entry retains, which the cache reports.
1354    fn retained_bytes(&self) -> usize {
1355        self.buffers()
1356            .iter()
1357            .map(|buffer| buffer.capacity() * std::mem::size_of::<Df64>())
1358            .sum()
1359    }
1360}
1361
1362/// Acquire the scratch for one operation shape, reusing what the cache already holds.
1363fn acquire_scratch(caches: &mut ExtensionCacheStore, shape: usize) -> Scratch {
1364    let key = ExtensionCacheKey::new(DF64_OPS_FAMILY, Scratch::CACHE_NAME, shape as u64);
1365    caches
1366        .get_mut::<Scratch>(&key)
1367        .map_or_else(Scratch::default, |scratch| Scratch {
1368            q_bar_t: std::mem::take(&mut scratch.q_bar_t),
1369            r_bar_t: std::mem::take(&mut scratch.r_bar_t),
1370            q_bar_t_q: std::mem::take(&mut scratch.q_bar_t_q),
1371            r_r_bar: std::mem::take(&mut scratch.r_r_bar),
1372            m: std::mem::take(&mut scratch.m),
1373            s: std::mem::take(&mut scratch.s),
1374            product: std::mem::take(&mut scratch.product),
1375            b: std::mem::take(&mut scratch.b),
1376        })
1377}
1378
1379/// Return the scratch to the cache with its retained bytes reported.
1380fn release_scratch(caches: &mut ExtensionCacheStore, shape: usize, scratch: Scratch) {
1381    let key = ExtensionCacheKey::new(DF64_OPS_FAMILY, Scratch::CACHE_NAME, shape as u64);
1382    let retained = scratch.retained_bytes();
1383    caches.put(key, scratch, retained);
1384}
1385
1386/// One operation input, borrowed when the runtime already owns it.
1387enum Input<'a> {
1388    Borrowed(&'a Tensor),
1389    // Boxed because a tensor value is much larger than a borrow, and every body only
1390    // reads it through `tensor`.
1391    Materialized(Box<Tensor>),
1392}
1393
1394impl Input<'_> {
1395    fn tensor(&self) -> &Tensor {
1396        match self {
1397            Self::Borrowed(tensor) => tensor,
1398            Self::Materialized(tensor) => tensor,
1399        }
1400    }
1401}
1402
1403/// Materialize a borrowed `f64` read without a session.
1404///
1405/// A reverse pass can hand the widening operation a strided view, such as the
1406/// broadcast of a seeding cotangent, and the widening operation's input is always a
1407/// preset `f64` tensor, so the elements are gathered by their own layout.
1408///
1409/// # Errors
1410///
1411/// Returns a typed error when the read is not a host `f64` tensor or its layout names
1412/// storage outside the borrowed buffer.
1413fn owned_f64_read(read: &TensorRead<'_>) -> tenferro_runtime::Result<Tensor> {
1414    let invalid = |message: &'static str| {
1415        tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
1416            "df64_from_f64",
1417            "input",
1418            message,
1419        ))
1420    };
1421    let TensorView::F64(view) = read.clone().tensor_view() else {
1422        return Err(invalid("df64_from_f64 takes a preset f64 tensor"));
1423    };
1424    let storage = view.host_storage().map_err(|source| {
1425        tenferro_runtime::Error::from(tenferro_tensor::Error::runtime_state_source(
1426            "df64_from_f64",
1427            source,
1428        ))
1429    })?;
1430    let shape = view.shape().to_vec();
1431    let mut values = Vec::with_capacity(storage.len().min(view.n_elements()));
1432    let mut index = vec![0usize; shape.len()];
1433    for _ in 0..view.n_elements() {
1434        let offset = view
1435            .linear_offset(&index)
1436            .ok_or_else(|| invalid("the borrowed f64 view is outside its buffer"))?;
1437        let value = storage
1438            .get(offset)
1439            .ok_or_else(|| invalid("the borrowed f64 view is outside its buffer"))?;
1440        values.push(*value);
1441        for (position, current) in index.iter_mut().enumerate() {
1442            *current += 1;
1443            if *current < shape[position] {
1444                break;
1445            }
1446            *current = 0;
1447        }
1448    }
1449    Tensor::from_vec_col_major(shape, values).map_err(|source| {
1450        tenferro_runtime::Error::from(tenferro_tensor::Error::runtime_state_source(
1451            "df64_from_f64",
1452            source,
1453        ))
1454    })
1455}
1456
1457/// Resolve every input of an operation.
1458///
1459/// A borrowed read is materialized in one pass first, so the session borrow does not
1460/// have to outlive the resolved inputs.
1461fn inputs_of<'a>(
1462    op: &'static str,
1463    session: Option<&mut dyn tenferro_tensor::BackendSession>,
1464    inputs: &'a [TensorRead<'a>],
1465) -> tenferro_runtime::Result<Vec<Input<'a>>> {
1466    let borrowed = inputs
1467        .iter()
1468        .any(|read| matches!(read, TensorRead::View(_)));
1469    if borrowed && session.is_none() {
1470        return Err(tenferro_runtime::Error::from(
1471            tenferro_tensor::Error::invalid_argument(
1472                op,
1473                "input",
1474                "the operation needs a session to read a borrowed input",
1475            ),
1476        ));
1477    }
1478    let mut materialized: Vec<Option<Tensor>> = (0..inputs.len()).map(|_| None).collect();
1479    if let Some(session) = session {
1480        for (index, read) in inputs.iter().enumerate() {
1481            if matches!(read, TensorRead::View(_)) {
1482                materialized[index] = Some(
1483                    session
1484                        .to_contiguous_read(read.clone())
1485                        .map_err(tenferro_runtime::Error::from)?,
1486                );
1487            }
1488        }
1489    }
1490    let mut resolved = Vec::with_capacity(inputs.len());
1491    for (read, owned) in inputs.iter().zip(materialized) {
1492        resolved.push(match (read, owned) {
1493            (TensorRead::Tensor(tensor), _) => Input::Borrowed(tensor),
1494            (TensorRead::View(_), Some(tensor)) => Input::Materialized(Box::new(tensor)),
1495            (TensorRead::View(_), None) => {
1496                return Err(tenferro_runtime::Error::from(
1497                    tenferro_tensor::Error::invalid_argument(
1498                        op,
1499                        "input",
1500                        "the operation needs a session to read a borrowed input",
1501                    ),
1502                ));
1503            }
1504        });
1505    }
1506    Ok(resolved)
1507}
1508
1509/// Resolve one input to a borrow or to materialized storage.
1510fn resolve_input<'a>(
1511    op: &'static str,
1512    session: Option<&mut dyn tenferro_tensor::BackendSession>,
1513    read: &TensorRead<'a>,
1514) -> tenferro_runtime::Result<Input<'a>> {
1515    match read {
1516        TensorRead::Tensor(tensor) => Ok(Input::Borrowed(tensor)),
1517        view @ TensorRead::View(_) => match session {
1518            Some(session) => Ok(Input::Materialized(Box::new(
1519                session
1520                    .to_contiguous_read(view.clone())
1521                    .map_err(tenferro_runtime::Error::from)?,
1522            ))),
1523            None => Err(tenferro_runtime::Error::from(
1524                tenferro_tensor::Error::invalid_argument(
1525                    op,
1526                    "input",
1527                    "the operation needs a session to read a borrowed input",
1528                ),
1529            )),
1530        },
1531    }
1532}
1533
1534/// Resolve the single input every operation in this crate takes.
1535///
1536/// The reverse pass can hand an operation a borrowed view of another value, so a
1537/// session materializes it into storage the operation owns. Without a session a view
1538/// is rejected explicitly rather than read through.
1539fn sole_input<'a>(
1540    op: &'static str,
1541    session: Option<&mut dyn tenferro_tensor::BackendSession>,
1542    inputs: &'a [TensorRead<'a>],
1543) -> tenferro_runtime::Result<Input<'a>> {
1544    match inputs.first() {
1545        Some(read) => resolve_input(op, session, read),
1546        None => Err(tenferro_runtime::Error::from(
1547            tenferro_tensor::Error::invalid_argument(op, "input", "the operation takes one input"),
1548        )),
1549    }
1550}
1551
1552/// Read one externally defined dense matrix from a tensor.
1553fn matrix_of<'a>(
1554    op: &'static str,
1555    tensor: &'a Tensor,
1556) -> tenferro_runtime::Result<crate::dense::Matrix<'a>> {
1557    let invalid = |message: &'static str| {
1558        tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
1559            op, "input", message,
1560        ))
1561    };
1562    let payload = external_payload::<Df64>(op, tensor).map_err(tenferro_runtime::Error::from)?;
1563    let [rows, columns] = match tensor.shape() {
1564        [rows, columns] => [*rows, *columns],
1565        _ => return Err(invalid("the operation takes a rank-2 matrix")),
1566    };
1567    if payload.as_slice().len() != rows * columns {
1568        return Err(invalid("the operation takes a dense column-major matrix"));
1569    }
1570    // The payload is caller-owned and borrowed for the length of the body, so no
1571    // tensor-sized copy is made on the way in.
1572    Ok(crate::dense::Matrix::borrowed(rows, payload.as_slice()))
1573}
1574
1575/// Wrap one dense matrix as an externally defined tensor.
1576fn tensor_of(
1577    op: &'static str,
1578    matrix: crate::dense::Matrix<'_>,
1579) -> tenferro_runtime::Result<Tensor> {
1580    let columns = matrix.columns();
1581    let host = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(
1582        vec![matrix.rows, columns],
1583        matrix.data.into_owned(),
1584    )
1585    .map_err(|source| {
1586        tenferro_runtime::Error::from(tenferro_tensor::Error::runtime_state_source(op, source))
1587    })?;
1588    Ok(Tensor::external(ErasedHostTensor::new(host)))
1589}
1590
1591/// Reverse-mode adjoint of the reduced QR factorization.
1592///
1593/// The adjoint of `A = Q R` for full column rank `A` is
1594/// `A_bar = (Q_bar + Q copyltu(R R_bar^T - Q_bar^T Q)) R^{-T}`, evaluated in the
1595/// external scalar, so the derivative keeps the factorization's precision.
1596fn qr_vjp_of(
1597    mask: (bool, bool),
1598    session: Option<&mut dyn tenferro_tensor::BackendSession>,
1599    caches: &mut ExtensionCacheStore,
1600    inputs: &[TensorRead<'_>],
1601) -> tenferro_runtime::Result<Vec<Tensor>> {
1602    let op = "df64_qr_vjp";
1603    let (has_q, has_r) = mask;
1604    let resolved = inputs_of(op, session, inputs)?;
1605    if resolved.len() != 2 + usize::from(has_q) + usize::from(has_r) {
1606        return Err(tenferro_runtime::Error::from(
1607            tenferro_tensor::Error::invalid_argument(
1608                op,
1609                "input",
1610                "the adjoint takes Q, R, and both cotangents",
1611            ),
1612        ));
1613    }
1614    let q = matrix_of(op, resolved[0].tensor())?;
1615    let r = matrix_of(op, resolved[1].tensor())?;
1616    let mut next = 2;
1617    // An absent cotangent is the zero cotangent, which is what an inactive derivative
1618    // means.
1619    let q_bar = if has_q {
1620        let matrix = matrix_of(op, resolved[next].tensor())?;
1621        next += 1;
1622        matrix
1623    } else {
1624        crate::dense::zeros(q.rows, q.columns())
1625    };
1626    let r_bar = if has_r {
1627        matrix_of(op, resolved[next].tensor())?
1628    } else {
1629        crate::dense::zeros(r.rows, r.columns())
1630    };
1631
1632    // Every intermediate comes from the accounted scratch, so an execution after the first
1633    // allocates only the factors it returns. The buffers are named fields, so one can be
1634    // written while another is read.
1635    let rows = q.rows;
1636    let columns = q.columns();
1637    let length = rows * columns;
1638    let square = r.rows * r.columns();
1639    let mut scratch = acquire_scratch(caches, square);
1640
1641    crate::dense::transpose_into(Scratch::slot(&mut scratch.q_bar_t, length), &q_bar);
1642    let q_bar_t = crate::dense::Matrix::borrowed(columns, scratch.q_bar_t.as_slice());
1643    crate::dense::transpose_into(Scratch::slot(&mut scratch.r_bar_t, square), &r_bar);
1644    {
1645        let r_bar_t = crate::dense::Matrix::borrowed(r.columns(), scratch.r_bar_t.as_slice());
1646        crate::dense::multiply_into(Scratch::slot(&mut scratch.r_r_bar, square), &r, &r_bar_t);
1647    }
1648    {
1649        let r = crate::dense::Matrix::borrowed(r.rows, scratch.r_r_bar.as_slice());
1650        crate::dense::multiply_into(Scratch::slot(&mut scratch.q_bar_t_q, square), &q_bar_t, &q);
1651        let q_bar_t_q = crate::dense::Matrix::borrowed(columns, scratch.q_bar_t_q.as_slice());
1652        // `M = R R_bar^T - Q_bar^T Q`.
1653        crate::dense::subtract_into(Scratch::slot(&mut scratch.m, square), &r, &q_bar_t_q);
1654    }
1655    {
1656        let m = crate::dense::Matrix::borrowed(r.rows, scratch.m.as_slice());
1657        // `copyltu(M)` is the lower triangle plus the strict lower triangle transposed.
1658        crate::dense::lower_triangle_into(Scratch::slot(&mut scratch.s, square), &m);
1659        // The second step adds to what the first wrote, so it must not clear the buffer.
1660        // Asking `slot` for it again would zero it and silently drop the lower triangle.
1661        crate::dense::add_strictly_lower_transposed_into(scratch.s.as_mut_slice(), &m);
1662    }
1663    {
1664        let s = crate::dense::Matrix::borrowed(r.rows, scratch.s.as_slice());
1665        crate::dense::multiply_into(Scratch::slot(&mut scratch.product, length), &q, &s);
1666        let product = crate::dense::Matrix::borrowed(rows, scratch.product.as_slice());
1667        let accumulated = Scratch::slot(&mut scratch.b, length);
1668        for (index, slot) in accumulated.iter_mut().enumerate() {
1669            let cotangent = q_bar.data.get(index).copied().unwrap_or_else(Df64::zero);
1670            *slot = cotangent + product.data[index];
1671        }
1672    }
1673    let b = crate::dense::Matrix::borrowed(rows, scratch.b.as_slice());
1674    let a_bar = crate::dense::solve_upper_from_the_right(&r, &b).ok_or_else(|| {
1675        tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
1676            op,
1677            "input",
1678            "the adjoint needs an invertible triangular factor",
1679        ))
1680    })?;
1681    release_scratch(caches, square, scratch);
1682    Ok(vec![tensor_of(op, a_bar)?])
1683}
1684
1685/// Forward-mode tangent of the reduced QR factorization.
1686///
1687/// With `A = Q R`, the tangent satisfies
1688/// `R_dot = triu(Q^T A_dot) R` and `Q_dot = (A_dot - Q R_dot) R^{-1}`, so the forward
1689/// rule emits both tangent outputs from the primal factors.
1690fn qr_jvp_of(
1691    session: Option<&mut dyn tenferro_tensor::BackendSession>,
1692    inputs: &[TensorRead<'_>],
1693) -> tenferro_runtime::Result<Vec<Tensor>> {
1694    let op = "df64_qr_jvp";
1695    let resolved = inputs_of(op, session, inputs)?;
1696    if resolved.len() != 3 {
1697        return Err(tenferro_runtime::Error::from(
1698            tenferro_tensor::Error::invalid_argument(
1699                op,
1700                "input",
1701                "the tangent takes Q, R, and the input tangent",
1702            ),
1703        ));
1704    }
1705    let q = matrix_of(op, resolved[0].tensor())?;
1706    let r = matrix_of(op, resolved[1].tensor())?;
1707    let a_dot = matrix_of(op, resolved[2].tensor())?;
1708
1709    // With A = Q R, the differential identity W = Q^T A_dot = S R + R_dot holds, where
1710    // S = Q^T Q_dot is skew and R_dot is upper triangular. The strictly lower part of W
1711    // therefore determines S by forward substitution, and R_dot = W - S R follows. Taking
1712    // the upper triangle of W directly would drop the S R term, which is zero only for a
1713    // single column.
1714    let w = crate::dense::multiply(&crate::dense::transpose(&q), &a_dot);
1715    let n = r.columns();
1716    let mut s = crate::dense::zeros(n, n);
1717    for column in 0..n {
1718        for row in (column + 1)..n {
1719            let mut value = w.at(row, column);
1720            for earlier in 0..column {
1721                value = value - s.at(row, earlier) * r.at(earlier, column);
1722            }
1723            let skew = value / r.at(column, column);
1724            s.set(row, column, skew);
1725            s.set(column, row, -skew);
1726        }
1727    }
1728    let r_dot = crate::dense::subtract(&w, &crate::dense::multiply(&s, &r));
1729    let residual = crate::dense::subtract(&a_dot, &crate::dense::multiply(&q, &r_dot));
1730    let q_dot = crate::dense::solve_upper_from_the_right(&r, &residual).ok_or_else(|| {
1731        tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
1732            op,
1733            "input",
1734            "the tangent needs an invertible triangular factor",
1735        ))
1736    })?;
1737    Ok(vec![tensor_of(op, q_dot)?, tensor_of(op, r_dot)?])
1738}
1739
1740/// The output shape a label list describes, given the extents the operands declare.
1741fn shape_of_labels(labels: &[Box<[u32]>], shapes: &[Vec<usize>], out_labels: &[u32]) -> Vec<usize> {
1742    let mut extents: Vec<(u32, usize)> = Vec::new();
1743    for (operand, operand_labels) in labels.iter().enumerate() {
1744        for (axis, label) in operand_labels.iter().enumerate() {
1745            if !extents.iter().any(|(existing, _)| existing == label) {
1746                extents.push((*label, shapes[operand][axis]));
1747            }
1748        }
1749    }
1750    out_labels
1751        .iter()
1752        .map(|label| {
1753            extents
1754                .iter()
1755                .find(|(existing, _)| existing == label)
1756                .map(|(_, extent)| *extent)
1757                .unwrap_or(1)
1758        })
1759        .collect()
1760}
1761
1762/// Contract two operand payloads over their shared labels, in the external scalar's arithmetic.
1763///
1764/// The output index space is walked once and, for each of its points, the contracted index space is
1765/// summed. A label an operand names but the output does not is summed, which is what ordinary einsum
1766/// notation means by it, and a label that repeats inside one operand was refused when the operation
1767/// was built.
1768fn contract_in_scalar(
1769    op: &'static str,
1770    labels: &[Box<[u32]>],
1771    out_labels: &[u32],
1772    lhs_values: &[Df64],
1773    lhs_shape: &[usize],
1774    rhs_values: &[Df64],
1775    rhs_shape: &[usize],
1776) -> tenferro_runtime::Result<Vec<Df64>> {
1777    if element_count(lhs_shape) != lhs_values.len() || element_count(rhs_shape) != rhs_values.len()
1778    {
1779        return Err(tenferro_runtime::Error::from(
1780            tenferro_tensor::Error::invalid_argument(
1781                op,
1782                "inputs",
1783                "the payload length does not match the declared shape",
1784            ),
1785        ));
1786    }
1787
1788    let mut extents: Vec<(u32, usize)> = Vec::new();
1789    let label_extent = |label: u32, extent: usize, extents: &mut Vec<(u32, usize)>| match extents
1790        .iter()
1791        .find(|(existing, _)| *existing == label)
1792    {
1793        Some((_, existing)) => *existing == extent,
1794        None => {
1795            extents.push((label, extent));
1796            true
1797        }
1798    };
1799    let mut ok = true;
1800    for (axis, label) in labels[0].iter().enumerate() {
1801        ok &= label_extent(*label, lhs_shape[axis], &mut extents);
1802    }
1803    for (axis, label) in labels[1].iter().enumerate() {
1804        ok &= label_extent(*label, rhs_shape[axis], &mut extents);
1805    }
1806    if !ok {
1807        return Err(tenferro_runtime::Error::from(
1808            tenferro_tensor::Error::invalid_argument(
1809                op,
1810                "inputs",
1811                "the inputs disagree on the extent of a shared label",
1812            ),
1813        ));
1814    }
1815    let extent_of = |label: u32, extents: &[(u32, usize)]| {
1816        extents
1817            .iter()
1818            .find(|(existing, _)| *existing == label)
1819            .map(|(_, extent)| *extent)
1820            .unwrap_or(1)
1821    };
1822    let out_shape: Vec<usize> = out_labels
1823        .iter()
1824        .map(|label| extent_of(*label, &extents))
1825        .collect();
1826    let mut summed_labels: Vec<u32> = Vec::new();
1827    for label in labels[0].iter().chain(labels[1].iter()) {
1828        if !out_labels.contains(label) && !summed_labels.contains(label) {
1829            summed_labels.push(*label);
1830        }
1831    }
1832    let summed_shape: Vec<usize> = summed_labels
1833        .iter()
1834        .map(|label| extent_of(*label, &extents))
1835        .collect();
1836
1837    let out_count = element_count(&out_shape);
1838    let summed_count = element_count(&summed_shape);
1839    let mut result = vec![Df64::zero(); out_count];
1840    let mut out_index = vec![0usize; out_shape.len()];
1841    let mut summed_index = vec![0usize; summed_shape.len()];
1842    for slot in result.iter_mut() {
1843        for value in summed_index.iter_mut() {
1844            *value = 0;
1845        }
1846        let mut accumulator = Df64::zero();
1847        for _ in 0..summed_count {
1848            let lhs_offset = offset_for(
1849                &labels[0],
1850                lhs_shape,
1851                out_labels,
1852                &out_index,
1853                &summed_labels,
1854                &summed_index,
1855            );
1856            let rhs_offset = offset_for(
1857                &labels[1],
1858                rhs_shape,
1859                out_labels,
1860                &out_index,
1861                &summed_labels,
1862                &summed_index,
1863            );
1864            accumulator = accumulator + lhs_values[lhs_offset] * rhs_values[rhs_offset];
1865            advance(&mut summed_index, &summed_shape);
1866        }
1867        *slot = accumulator;
1868        advance(&mut out_index, &out_shape);
1869    }
1870    Ok(result)
1871}
1872
1873/// Wrap a payload as an external tensor of the shape the labels describe.
1874fn external_of(values: Vec<Df64>, shape: Vec<usize>) -> tenferro_runtime::Result<Tensor> {
1875    let tensor = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(shape, values)
1876        .map_err(tenferro_runtime::Error::from)?;
1877    Ok(Tensor::external(ErasedHostTensor::new(tensor)))
1878}
1879
1880/// Contract every operand, folding from the left.
1881///
1882/// An intermediate keeps exactly the labels the remaining operands or the output still need, so a
1883/// label that only the already-contracted operands name is summed by that step, which is what the
1884/// notation means by a label the output omits.
1885/// Fold a list of operands from the left, returning the contracted values and their shape.
1886///
1887/// An intermediate keeps exactly the labels the remaining operands or the output still need, so a
1888/// label only the already-contracted operands name is summed by that step, which is what the notation
1889/// means by a label the output omits.
1890fn fold_in_scalar(
1891    op: &'static str,
1892    labels: &[Box<[u32]>],
1893    out_labels: &[u32],
1894    operands: &[(Vec<Df64>, Vec<usize>)],
1895) -> tenferro_runtime::Result<(Vec<Df64>, Vec<usize>)> {
1896    let mut accumulator = operands[0].0.clone();
1897    let mut accumulator_shape = operands[0].1.clone();
1898    let mut accumulator_labels: Vec<u32> = labels[0].to_vec();
1899    for index in 1..labels.len() {
1900        let next_labels: Vec<u32> = labels[index].to_vec();
1901        let keep: Vec<u32> = if index + 1 == labels.len() {
1902            out_labels.to_vec()
1903        } else {
1904            let mut keep: Vec<u32> = Vec::new();
1905            for label in accumulator_labels.iter().chain(next_labels.iter()) {
1906                let needed_later = out_labels.contains(label)
1907                    || labels[index + 1..].iter().any(|rest| rest.contains(label));
1908                if needed_later && !keep.contains(label) {
1909                    keep.push(*label);
1910                }
1911            }
1912            keep
1913        };
1914        let pair = [
1915            accumulator_labels.clone().into_boxed_slice(),
1916            next_labels.clone().into_boxed_slice(),
1917        ];
1918        let contracted = contract_in_scalar(
1919            op,
1920            &pair,
1921            &keep,
1922            &accumulator,
1923            &accumulator_shape,
1924            &operands[index].0,
1925            &operands[index].1,
1926        )?;
1927        let contracted_shape = shape_of_labels(
1928            &pair,
1929            &[accumulator_shape.clone(), operands[index].1.clone()],
1930            &keep,
1931        );
1932        accumulator = contracted;
1933        accumulator_shape = contracted_shape;
1934        accumulator_labels = keep;
1935    }
1936    Ok((accumulator, accumulator_shape))
1937}
1938
1939/// Contract every operand, returning the result as an external tensor.
1940fn einsum_of(
1941    labels: &[Box<[u32]>],
1942    out_labels: &[u32],
1943    session: Option<&mut dyn tenferro_tensor::BackendSession>,
1944    inputs: &[TensorRead<'_>],
1945) -> tenferro_runtime::Result<Vec<Tensor>> {
1946    let op = "df64_einsum";
1947    let resolved = inputs_of(op, session, inputs)?;
1948    if resolved.len() != labels.len() || labels.len() < 2 {
1949        return Err(tenferro_runtime::Error::from(
1950            tenferro_tensor::Error::invalid_argument(
1951                op,
1952                "input",
1953                "a contraction takes one input per operand",
1954            ),
1955        ));
1956    }
1957    let mut operands: Vec<(Vec<Df64>, Vec<usize>)> = Vec::with_capacity(resolved.len());
1958    for operand in &resolved {
1959        operands.push((
1960            payload_of::<Df64>(op, operand.tensor())?,
1961            operand.tensor().shape().to_vec(),
1962        ));
1963    }
1964    let (values, shape) = fold_in_scalar(op, labels, out_labels, &operands)?;
1965    Ok(vec![external_of(values, shape)?])
1966}
1967
1968/// The adjoint of a contraction: the cotangent of each operand, in the extended scalar.
1969///
1970/// The labels rotate rather than change, so the adjoint contracts the output cotangent with the
1971/// other operand and lands on the operand it differentiates.
1972fn einsum_vjp_of(
1973    labels: &[Box<[u32]>],
1974    out_labels: &[u32],
1975    session: Option<&mut dyn tenferro_tensor::BackendSession>,
1976    inputs: &[TensorRead<'_>],
1977) -> tenferro_runtime::Result<Vec<Tensor>> {
1978    let op = "df64_einsum_vjp";
1979    let resolved = inputs_of(op, session, inputs)?;
1980    if resolved.len() != labels.len() + 1 || labels.len() < 2 {
1981        return Err(tenferro_runtime::Error::from(
1982            tenferro_tensor::Error::invalid_argument(
1983                op,
1984                "input",
1985                "the adjoint takes one input per operand and the output cotangent",
1986            ),
1987        ));
1988    }
1989    let mut operands: Vec<(Vec<Df64>, Vec<usize>)> = Vec::with_capacity(labels.len());
1990    for operand in &resolved[..labels.len()] {
1991        operands.push((
1992            payload_of::<Df64>(op, operand.tensor())?,
1993            operand.tensor().shape().to_vec(),
1994        ));
1995    }
1996    let cotangent = payload_of::<Df64>(op, resolved[labels.len()].tensor())?;
1997    let cotangent_shape = resolved[labels.len()].tensor().shape().to_vec();
1998
1999    // The cotangent of each operand is the contraction of the output cotangent with every other
2000    // operand, so the cotangent takes that operand's place and carries the output's labels.
2001    let mut outputs = Vec::with_capacity(labels.len());
2002    for position in 0..labels.len() {
2003        let mut substituted: Vec<(Vec<Df64>, Vec<usize>)> = Vec::with_capacity(labels.len());
2004        for (index, operand) in operands.iter().enumerate() {
2005            if index == position {
2006                substituted.push((cotangent.clone(), cotangent_shape.clone()));
2007            } else {
2008                substituted.push(operand.clone());
2009            }
2010        }
2011        let mut step_labels: Vec<Box<[u32]>> = labels.to_vec();
2012        step_labels[position] = out_labels.to_vec().into_boxed_slice();
2013        let (values, shape) = fold_in_scalar(op, &step_labels, &labels[position], &substituted)?;
2014        outputs.push(external_of(values, shape)?);
2015    }
2016    Ok(outputs)
2017}
2018
2019/// The forward tangent of a contraction: each operand's tangent contracted with the other.
2020fn einsum_jvp_of(
2021    labels: &[Box<[u32]>],
2022    out_labels: &[u32],
2023    tangents: &[bool],
2024    session: Option<&mut dyn tenferro_tensor::BackendSession>,
2025    inputs: &[TensorRead<'_>],
2026) -> tenferro_runtime::Result<Vec<Tensor>> {
2027    let op = "df64_einsum_jvp";
2028    let resolved = inputs_of(op, session, inputs)?;
2029    let expected = labels.len() + tangents.iter().filter(|present| **present).count();
2030    if resolved.len() != expected || labels.len() != tangents.len() || labels.len() < 2 {
2031        return Err(tenferro_runtime::Error::from(
2032            tenferro_tensor::Error::invalid_argument(
2033                op,
2034                "input",
2035                "the tangent takes one input per operand and one per tangent",
2036            ),
2037        ));
2038    }
2039    let mut operands: Vec<(Vec<Df64>, Vec<usize>)> = Vec::with_capacity(labels.len());
2040    let mut shapes: Vec<Vec<usize>> = Vec::with_capacity(labels.len());
2041    for operand in &resolved[..labels.len()] {
2042        operands.push((
2043            payload_of::<Df64>(op, operand.tensor())?,
2044            operand.tensor().shape().to_vec(),
2045        ));
2046        shapes.push(operand.tensor().shape().to_vec());
2047    }
2048    let mut dots: Vec<Option<(Vec<Df64>, Vec<usize>)>> = vec![None; labels.len()];
2049    let mut slot = labels.len();
2050    for (index, present) in tangents.iter().enumerate() {
2051        if !present {
2052            continue;
2053        }
2054        dots[index] = Some((
2055            payload_of::<Df64>(op, resolved[slot].tensor())?,
2056            resolved[slot].tensor().shape().to_vec(),
2057        ));
2058        slot += 1;
2059    }
2060    let out_shape = shape_of_labels(labels, &shapes, out_labels);
2061
2062    // Each tangent takes its operand's place, under that operand's own labels, and the parts sum.
2063    let mut total: Option<Vec<Df64>> = None;
2064    for (position, dot) in dots.iter().enumerate() {
2065        let Some((values, shape)) = dot else {
2066            continue;
2067        };
2068        let mut substituted = operands.clone();
2069        substituted[position] = (values.clone(), shape.clone());
2070        let (part, _) = fold_in_scalar(op, labels, out_labels, &substituted)?;
2071        total = Some(match total {
2072            Some(existing) => existing
2073                .iter()
2074                .zip(&part)
2075                .map(|(left, right)| *left + *right)
2076                .collect(),
2077            None => part,
2078        });
2079    }
2080    let values = total.ok_or_else(|| {
2081        tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
2082            op,
2083            "tangents",
2084            "at least one operand must carry a tangent",
2085        ))
2086    })?;
2087    Ok(vec![external_of(values, out_shape)?])
2088}
2089
2090/// The column-major offset an input's labels select at one output and contracted index.
2091fn offset_for(
2092    input_labels: &[u32],
2093    input_shape: &[usize],
2094    out_labels: &[u32],
2095    out_index: &[usize],
2096    summed_labels: &[u32],
2097    summed_index: &[usize],
2098) -> usize {
2099    let mut offset = 0usize;
2100    let mut stride = 1usize;
2101    for (axis, label) in input_labels.iter().enumerate() {
2102        let position = out_labels
2103            .iter()
2104            .position(|candidate| candidate == label)
2105            .map(|index| out_index[index])
2106            .or_else(|| {
2107                summed_labels
2108                    .iter()
2109                    .position(|candidate| candidate == label)
2110                    .map(|index| summed_index[index])
2111            })
2112            .unwrap_or(0);
2113        offset += position * stride;
2114        stride *= input_shape[axis];
2115    }
2116    offset
2117}
2118
2119/// Advance a column-major odometer by one, wrapping the fastest-varying axis first.
2120fn advance(index: &mut [usize], shape: &[usize]) {
2121    for axis in 0..shape.len() {
2122        index[axis] += 1;
2123        if index[axis] < shape[axis] {
2124            return;
2125        }
2126        index[axis] = 0;
2127    }
2128}
2129
2130/// The number of elements a shape describes.
2131fn element_count(shape: &[usize]) -> usize {
2132    shape.iter().product()
2133}
2134
2135/// The externally defined payload of a tensor.
2136fn payload_of<T: tenferro_tensor_core::Scalar>(
2137    op: &'static str,
2138    tensor: &Tensor,
2139) -> tenferro_runtime::Result<Vec<T>> {
2140    external_payload::<T>(op, tensor)
2141        .map(|payload| payload.as_slice().to_vec())
2142        .map_err(tenferro_runtime::Error::from)
2143}
2144
2145/// Fill a tensor of `shape` with the scalar input's value.
2146fn expand_of(
2147    shape: &[usize],
2148    session: Option<&mut dyn tenferro_tensor::BackendSession>,
2149    inputs: &[TensorRead<'_>],
2150) -> tenferro_runtime::Result<Vec<Tensor>> {
2151    let input = sole_input("df64_expand", session, inputs)?;
2152    let tensor = input.tensor();
2153    let payload =
2154        external_payload::<Df64>("df64_expand", tensor).map_err(tenferro_runtime::Error::from)?;
2155    let value = payload.as_slice().first().copied().ok_or_else(|| {
2156        tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
2157            "df64_expand",
2158            "input",
2159            "df64_expand takes a scalar payload",
2160        ))
2161    })?;
2162    let count: usize = shape.iter().product();
2163    let output = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(
2164        shape.to_vec(),
2165        vec![value; count],
2166    )
2167    .map_err(tenferro_runtime::Error::from)?;
2168    Ok(vec![Tensor::external(ErasedHostTensor::new(output))])
2169}
2170
2171fn total_of(
2172    session: Option<&mut dyn tenferro_tensor::BackendSession>,
2173    inputs: &[TensorRead<'_>],
2174) -> tenferro_runtime::Result<Vec<Tensor>> {
2175    let input = sole_input("df64_total", session, inputs)?;
2176    let tensor = input.tensor();
2177    let payload =
2178        external_payload::<Df64>("df64_total", tensor).map_err(tenferro_runtime::Error::from)?;
2179    let total = scalar_fold::<Df64, Df64Add>("df64_total", payload, Df64::zero())
2180        .map_err(tenferro_runtime::Error::from)?;
2181    let output = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![], vec![total])
2182        .map_err(tenferro_runtime::Error::from)?;
2183    Ok(vec![Tensor::external(ErasedHostTensor::new(output))])
2184}
2185
2186/// Which numerical body a prepared operation runs.
2187#[derive(Debug)]
2188enum Df64Body {
2189    /// Total sum of the input.
2190    Total,
2191    /// The scalar input broadcast to this shape.
2192    Expand(Box<[usize]>),
2193    /// Reduced QR factorization of the matrix input.
2194    Qr,
2195    /// Narrow the external scalar to `f64`.
2196    ToF64,
2197    /// Widen a preset `f64` tensor into the external scalar.
2198    FromF64,
2199    /// The adjoint of the factorization, with the cotangents the caller supplied.
2200    QrVjp((bool, bool)),
2201    /// The tangent of the factorization.
2202    QrJvp,
2203    /// The forward tangent of a pairwise contraction.
2204    EinsumJvp {
2205        /// The labels of each operand, in that operand's own axis order.
2206        inputs: Box<[Box<[u32]>]>,
2207        /// The labels of the contraction's output, in the output's axis order.
2208        out: Box<[u32]>,
2209        /// Which operands carry a tangent, in operand order.
2210        tangents: Box<[bool]>,
2211    },
2212    /// The adjoint of a pairwise contraction.
2213    EinsumVjp {
2214        /// The labels of each operand, in that operand's own axis order.
2215        inputs: Box<[Box<[u32]>]>,
2216        /// The labels of the contraction's output, in the output's axis order.
2217        out: Box<[u32]>,
2218    },
2219    /// A pairwise contraction of two external tensors over their shared labels.
2220    Einsum {
2221        /// The labels of each input, in that input's own axis order.
2222        inputs: Box<[Box<[u32]>]>,
2223        /// The labels of the output, in the output's axis order.
2224        out: Box<[u32]>,
2225    },
2226}
2227
2228impl Df64Body {
2229    fn execute(
2230        &self,
2231        session: Option<&mut dyn tenferro_tensor::BackendSession>,
2232        caches: &mut ExtensionCacheStore,
2233        inputs: &[TensorRead<'_>],
2234    ) -> tenferro_runtime::Result<Vec<Tensor>> {
2235        match self {
2236            Self::Total => total_of(session, inputs),
2237            Self::Expand(shape) => expand_of(shape, session, inputs),
2238            Self::Qr => qr_of(session, inputs),
2239            Self::ToF64 => to_f64_of(session, inputs),
2240            Self::FromF64 => from_f64_of(session, inputs),
2241            Self::QrVjp(mask) => qr_vjp_of(*mask, session, caches, inputs),
2242            Self::QrJvp => qr_jvp_of(session, inputs),
2243            Self::Einsum {
2244                inputs: labels,
2245                out,
2246            } => einsum_of(labels, out, session, inputs),
2247            Self::EinsumVjp {
2248                inputs: labels,
2249                out,
2250            } => einsum_vjp_of(labels, out, session, inputs),
2251            Self::EinsumJvp {
2252                inputs: labels,
2253                out,
2254                tangents,
2255            } => einsum_jvp_of(labels, out, tangents, session, inputs),
2256        }
2257    }
2258}
2259
2260#[derive(Debug)]
2261struct Df64Prepared {
2262    binding: PreparedOperationBinding,
2263    specialization: SpecializationProjection,
2264    body: Df64Body,
2265}
2266
2267impl PreparedOperation for Df64Prepared {
2268    fn binding(&self) -> &PreparedOperationBinding {
2269        &self.binding
2270    }
2271
2272    fn specialization(&self) -> &SpecializationProjection {
2273        &self.specialization
2274    }
2275
2276    fn retained_bytes(&self) -> usize {
2277        0
2278    }
2279}
2280
2281impl PreparedOperationExecutor for Df64Prepared {
2282    fn execute(
2283        &self,
2284        _context: &mut ErasedExecutionContext<'_>,
2285        caches: &mut ExtensionCacheStore,
2286        inputs: &[TensorRead<'_>],
2287    ) -> tenferro_runtime::Result<Vec<Tensor>> {
2288        // The erased context does not expose a session, so a body that needs one relies
2289        // on the session entry point and gathers what it can without one.
2290        self.body.execute(None, caches, inputs)
2291    }
2292
2293    fn supports_session(&self) -> bool {
2294        true
2295    }
2296
2297    fn execute_in_session(
2298        &self,
2299        session: &mut dyn tenferro_tensor::BackendSession,
2300        caches: &mut ExtensionCacheStore,
2301        inputs: &[TensorRead<'_>],
2302    ) -> tenferro_runtime::Result<Vec<Tensor>> {
2303        self.body.execute(Some(session), caches, inputs)
2304    }
2305}
2306
2307#[derive(Debug)]
2308struct Df64Engine {
2309    family_id: &'static str,
2310    engine_id: EngineId,
2311}
2312
2313impl ExtensionEngine for Df64Engine {
2314    fn family_id(&self) -> &'static str {
2315        self.family_id
2316    }
2317
2318    fn engine_id(&self) -> &EngineId {
2319        &self.engine_id
2320    }
2321
2322    fn context_identity(&self) -> ExecutionContextIdentity {
2323        ExecutionContextIdentity::of::<CpuBackend>()
2324    }
2325
2326    fn prepare(
2327        &self,
2328        request: ExtensionPrepareRequest<'_>,
2329    ) -> Result<PrepareCapability, PrepareError> {
2330        let operation = request.operation().as_any();
2331        let body = if let Some(expand) = operation.downcast_ref::<Df64Expand>() {
2332            Df64Body::Expand(expand.shape.clone())
2333        } else if operation.downcast_ref::<Df64Qr>().is_some() {
2334            Df64Body::Qr
2335        } else if operation.downcast_ref::<Df64ToF64>().is_some() {
2336            Df64Body::ToF64
2337        } else if operation.downcast_ref::<Df64FromF64>().is_some() {
2338            Df64Body::FromF64
2339        } else if let Some(adjoint) = operation.downcast_ref::<Df64QrVjp>() {
2340            Df64Body::QrVjp((adjoint.has_q, adjoint.has_r))
2341        } else if operation.downcast_ref::<Df64QrJvp>().is_some() {
2342            Df64Body::QrJvp
2343        } else if let Some(tangent) = operation.downcast_ref::<Df64EinsumJvp>() {
2344            Df64Body::EinsumJvp {
2345                inputs: tangent
2346                    .input_labels()
2347                    .iter()
2348                    .map(|labels| labels.clone().into_boxed_slice())
2349                    .collect(),
2350                out: tangent.out_labels().to_vec().into_boxed_slice(),
2351                tangents: tangent.tangents().to_vec().into_boxed_slice(),
2352            }
2353        } else if let Some(adjoint) = operation.downcast_ref::<Df64EinsumVjp>() {
2354            Df64Body::EinsumVjp {
2355                inputs: adjoint
2356                    .input_labels()
2357                    .iter()
2358                    .map(|labels| labels.clone().into_boxed_slice())
2359                    .collect(),
2360                out: adjoint.out_labels().to_vec().into_boxed_slice(),
2361            }
2362        } else if let Some(contraction) = operation.downcast_ref::<Df64Einsum>() {
2363            Df64Body::Einsum {
2364                inputs: contraction
2365                    .input_labels()
2366                    .iter()
2367                    .map(|labels| labels.clone().into_boxed_slice())
2368                    .collect(),
2369                out: contraction.out_labels().to_vec().into_boxed_slice(),
2370            }
2371        } else {
2372            Df64Body::Total
2373        };
2374        let prepared = Arc::new(Df64Prepared {
2375            binding: request.binding().clone(),
2376            specialization: request.specialization().clone(),
2377            body,
2378        });
2379        let operation: PreparedOperationHandle = Arc::clone(&prepared) as PreparedOperationHandle;
2380        let executor: PreparedOperationExecutorHandle = prepared as PreparedOperationExecutorHandle;
2381        Ok(PrepareCapability::Prepared(
2382            PreparedOperationPlan::executable(operation, executor),
2383        ))
2384    }
2385}
2386
2387#[derive(Debug)]
2388struct Df64Config {
2389    family_id: &'static str,
2390}
2391
2392impl ExtensionPlanningConfig for Df64Config {
2393    fn family_id(&self) -> &'static str {
2394        self.family_id
2395    }
2396
2397    fn as_any(&self) -> &dyn Any {
2398        self
2399    }
2400
2401    fn payload_hash(&self, state: &mut dyn Hasher) {
2402        state.write(self.family_id.as_bytes());
2403    }
2404
2405    fn payload_eq(&self, other: &dyn ExtensionPlanningConfig) -> bool {
2406        other
2407            .as_any()
2408            .downcast_ref::<Self>()
2409            .is_some_and(|other| self.family_id == other.family_id)
2410    }
2411
2412    fn retained_bytes(&self) -> usize {
2413        0
2414    }
2415}
2416
2417#[derive(Debug)]
2418struct Df64TotalModule {
2419    module_id: ExtensionModuleId,
2420    engine_id: EngineId,
2421}
2422
2423impl ExtensionModule for Df64TotalModule {
2424    fn module_id(&self) -> &ExtensionModuleId {
2425        &self.module_id
2426    }
2427
2428    fn configure(
2429        &self,
2430        registrar: &mut ExtensionModuleRegistrar<'_>,
2431    ) -> Result<(), ExtensionModuleError> {
2432        registrar.register_engine(Arc::new(Df64Engine {
2433            family_id: DF64_OPS_FAMILY,
2434            engine_id: self.engine_id.clone(),
2435        }))?;
2436        registrar.register_planning_config(
2437            self.engine_id.clone(),
2438            Arc::new(Df64Config {
2439                family_id: DF64_OPS_FAMILY,
2440            }),
2441        )
2442    }
2443}
2444
2445/// The module a downstream application installs to enable [`Df64Total`].
2446///
2447/// # Errors
2448///
2449/// Returns an error when the CPU runtime engine identifier is unavailable in this
2450/// process.
2451///
2452/// # Examples
2453///
2454/// ```rust
2455/// use tenferro_df64_proof::extension::module;
2456///
2457/// assert!(module().is_ok());
2458/// ```
2459pub fn module() -> Result<Arc<dyn ExtensionModule>, tenferro_runtime::RuntimeConfigError> {
2460    module_for_engine(tenferro_cpu::runtime_engine_id()?)
2461}
2462
2463/// The module a downstream application installs when it composes more than one CPU
2464/// backend in one runtime.
2465///
2466/// The module binds the contribution's operations to `engine_id`, so an application that
2467/// keeps a standard backend and a contribution backend under distinct engine identities
2468/// installs this module against the contribution's engine.
2469///
2470/// # Errors
2471///
2472/// Returns [`tenferro_runtime::RuntimeConfigError`] when the module's configured
2473/// identifier is invalid.
2474///
2475/// # Examples
2476///
2477/// ```rust
2478/// use tenferro_df64_proof::extension::module_for_engine;
2479/// use tenferro_runtime::EngineId;
2480///
2481/// let module = module_for_engine(EngineId::new("example.df64.engine.v1")?)?;
2482/// assert_eq!(module.module_id().as_str(), "tenferro-df64-proof.module");
2483/// # Ok::<(), Box<dyn std::error::Error>>(())
2484/// ```
2485pub fn module_for_engine(
2486    engine_id: EngineId,
2487) -> Result<Arc<dyn ExtensionModule>, tenferro_runtime::RuntimeConfigError> {
2488    Ok(Arc::new(Df64TotalModule {
2489        module_id: ExtensionModuleId::new("tenferro-df64-proof.module")?,
2490        engine_id,
2491    }))
2492}
2493
2494/// Run [`Df64Total`] on an eagerly held external tensor.
2495///
2496/// # Errors
2497///
2498/// Returns [`tenferro_runtime::Error::RuntimeStateSource`] when the contribution's module
2499/// cannot be registered, or when the operation cannot be prepared or executed for this
2500/// input.
2501///
2502/// # Examples
2503///
2504/// ```rust
2505/// use tenferro_ad::extension::ExtensionOp;
2506/// use tenferro_df64_proof::extension::{module, Df64Total};
2507///
2508/// assert_eq!(<Df64Total as ExtensionOp>::input_count(&Df64Total), 1);
2509/// assert!(module().is_ok());
2510/// ```
2511pub fn apply_total(input: &EagerTensor) -> tenferro_runtime::Result<Vec<EagerTensor>> {
2512    let module = module().map_err(|source| {
2513        tenferro_runtime::Error::runtime_state_source(
2514            "df64_total",
2515            tenferro_runtime::ErrorPhase::Execution,
2516            source,
2517        )
2518    })?;
2519    apply_eager_with_extension_session(Arc::new(Df64Total), &[input], module)
2520}