Skip to main content

tenferro_df64_proof/
ad.rs

1//! First-order AD for the contribution's externally defined operations.
2//!
3//! The runtime keys one rule per family and role, so the rules below dispatch on the
4//! operation payload. Each rule emits the contribution's own operations for the
5//! derivative work: a preset broadcast is not available for a scalar tenferro does
6//! not declare, and the factorization's adjoint is the contribution's numerical body
7//! too.
8
9use std::sync::Arc;
10
11use tenferro_ad::semantic_extension::{
12    AdValue, ResidualSpec, SemanticAdError, SemanticAdRuleRole, SemanticLinearizeRequest,
13    SemanticLinearizeResult, SemanticLinearizeRule, SemanticPrimalVjpRequest,
14    SemanticPrimalVjpRule,
15};
16use tenferro_ops::ext_op::ExtensionOp;
17use tenferro_runtime::program::SemanticProgramBuilder;
18
19use crate::extension::{
20    Df64Einsum, Df64EinsumJvp, Df64EinsumVjp, Df64Expand, Df64FromF64, Df64Qr, Df64QrJvp,
21    Df64QrVjp, Df64ToF64, Df64Total, DF64_OPS_FAMILY,
22};
23
24/// One operation of the contribution's family.
25#[derive(Clone, Copy, Debug, PartialEq, Eq)]
26enum Df64Op {
27    /// The total sum.
28    Total,
29    /// The scalar broadcast, which is the total sum's adjoint.
30    Expand,
31    /// The reduced QR factorization.
32    Qr,
33    /// A matrix contraction.
34    Einsum,
35    /// The narrowing conversion to `f64`.
36    ToF64,
37    /// The widening conversion from `f64`.
38    FromF64,
39    /// The factorization's adjoint.
40    QrVjp,
41    /// The factorization's tangent.
42    QrJvp,
43}
44
45impl Df64Op {
46    /// Classify the payload a rule was asked about.
47    fn of(op: &dyn ExtensionOp) -> Option<Self> {
48        let any = op.as_any();
49        if any.downcast_ref::<Df64Total>().is_some() {
50            Some(Self::Total)
51        } else if any.downcast_ref::<Df64Expand>().is_some() {
52            Some(Self::Expand)
53        } else if any.downcast_ref::<Df64Einsum>().is_some() {
54            Some(Self::Einsum)
55        } else if any.downcast_ref::<Df64Qr>().is_some() {
56            Some(Self::Qr)
57        } else if any.downcast_ref::<Df64ToF64>().is_some() {
58            Some(Self::ToF64)
59        } else if any.downcast_ref::<Df64FromF64>().is_some() {
60            Some(Self::FromF64)
61        } else if any.downcast_ref::<Df64QrVjp>().is_some() {
62            Some(Self::QrVjp)
63        } else if any.downcast_ref::<Df64QrJvp>().is_some() {
64            Some(Self::QrJvp)
65        } else {
66            None
67        }
68    }
69}
70
71/// An operation outside a rule's domain is an error rather than a guess.
72fn unsupported(op: Df64Op, role: SemanticAdRuleRole) -> SemanticAdError {
73    SemanticAdError::Rule {
74        family_id: DF64_OPS_FAMILY,
75        role,
76        source: Box::new(std::io::Error::other(format!(
77            "the Df64 operations family has no {role:?} rule for {op:?}"
78        ))),
79    }
80}
81
82/// Reverse-mode rule for the contribution's operations.
83///
84/// The total sum's adjoint is a broadcast, and the conversions are linear, so their
85/// adjoint swaps the two directions. None of them reads the primal value, so the rule
86/// declares no residual.
87///
88/// # Examples
89///
90/// ```rust
91/// use tenferro_ad::semantic_extension::{SemanticExtensionRuleSet, SemanticPrimalVjpRule};
92/// use tenferro_df64_proof::ad::Df64VjpRule;
93///
94/// let rule = Df64VjpRule;
95/// assert_eq!(
96///     <Df64VjpRule as SemanticPrimalVjpRule>::family_id(&rule),
97///     "tenferro-df64-proof.df64_ops.v1"
98/// );
99///
100/// // The rule set a downstream application installs.
101/// let rules = SemanticExtensionRuleSet::new()
102///     .with_primal_vjp(std::sync::Arc::new(rule))
103///     .expect("one rule per family");
104/// assert!(rules.lookup_primal_vjp("tenferro-df64-proof.df64_ops.v1").is_some());
105/// ```
106#[derive(Debug)]
107pub struct Df64VjpRule;
108
109impl SemanticPrimalVjpRule for Df64VjpRule {
110    fn family_id(&self) -> &'static str {
111        DF64_OPS_FAMILY
112    }
113
114    fn residual_mask(&self) -> ResidualSpec {
115        // The union over the family: the factorization's adjoint reads its primal
116        // factors, and the contraction's adjoint reads its operands (the case the
117        // residual specification names for an einsum rule).
118        ResidualSpec::all_outputs().with_all_inputs()
119    }
120
121    fn primal_vjp(
122        &self,
123        request: SemanticPrimalVjpRequest<'_>,
124        builder: &mut SemanticProgramBuilder,
125    ) -> Result<Box<[AdValue]>, SemanticAdError> {
126        let role = SemanticAdRuleRole::PrimalVjp;
127        let op = Df64Op::of(request.op()).ok_or_else(|| SemanticAdError::Rule {
128            family_id: DF64_OPS_FAMILY,
129            role,
130            source: Box::new(std::io::Error::other(
131                "the payload is not one of the Df64 operations",
132            )),
133        })?;
134        let inactive = || vec![AdValue::Absent; request.primal_input_count()].into_boxed_slice();
135        let emitted = match op {
136            Df64Op::Total => {
137                let Some(AdValue::Value(cotangent)) = request.cotangent_outputs().first().copied()
138                else {
139                    return Ok(inactive());
140                };
141                // The adjoint places the cotangent back into the input's shape, which
142                // the broadcast payload carries, so the shape is read from the primal
143                // input's metadata.
144                let shape = exact_shape(&request)?;
145                builder.add_extension(Arc::new(Df64Expand::new(shape)), &[cotangent])
146            }
147            // A linear conversion's adjoint is the opposite conversion.
148            Df64Op::ToF64 | Df64Op::FromF64 => {
149                let Some(AdValue::Value(cotangent)) = request.cotangent_outputs().first().copied()
150                else {
151                    return Ok(inactive());
152                };
153                let operation = if op == Df64Op::ToF64 {
154                    Arc::new(Df64FromF64) as Arc<dyn ExtensionOp>
155                } else {
156                    Arc::new(Df64ToF64) as Arc<dyn ExtensionOp>
157                };
158                builder.add_extension(operation, &[cotangent])
159            }
160            Df64Op::Qr => {
161                // A loss need not depend on both factors, so the adjoint is told which
162                // cotangents are present and reads the primal factors it needs.
163                let has_q = matches!(request.cotangent_outputs().first(), Some(AdValue::Value(_)));
164                let has_r = matches!(request.cotangent_outputs().get(1), Some(AdValue::Value(_)));
165                if !has_q && !has_r {
166                    return Ok(inactive());
167                }
168                let mut operands = vec![
169                    request.primal_output_value(0)?,
170                    request.primal_output_value(1)?,
171                ];
172                for (present, cotangent) in [
173                    (has_q, request.cotangent_outputs().first().copied()),
174                    (has_r, request.cotangent_outputs().get(1).copied()),
175                ] {
176                    if !present {
177                        continue;
178                    }
179                    match cotangent {
180                        Some(AdValue::Value(value)) => operands.push(value),
181                        _ => return Ok(inactive()),
182                    }
183                }
184                builder.add_extension(Arc::new(Df64QrVjp::of(has_q, has_r)), &operands)
185            }
186            Df64Op::Einsum => {
187                // The adjoint of a contraction contracts the output cotangent with the other
188                // operand, so the helper reads both operands and the cotangent and produces one
189                // cotangent per operand.
190                let Some(contraction) = request.op().as_any().downcast_ref::<Df64Einsum>() else {
191                    return Err(unsupported(op, role));
192                };
193                let Some(AdValue::Value(cotangent)) = request.cotangent_outputs().first().copied()
194                else {
195                    return Ok(inactive());
196                };
197                // The adjoint contracts the cotangent with the other operands in every operand's
198                // place, so it carries the whole operand list rather than a pair.
199                let operands: Vec<&[u32]> = contraction
200                    .input_labels()
201                    .iter()
202                    .map(|labels| labels.as_slice())
203                    .collect();
204                // The primal operation accepted this pattern, so the adjoint's validation is a
205                // rule-level invariant rather than a user error.
206                let Ok(adjoint) = Df64EinsumVjp::of(&operands, contraction.out_labels()) else {
207                    return Err(unsupported(op, role));
208                };
209                // Every operand reaches the helper, followed by the cotangent.
210                let mut call_operands = Vec::with_capacity(contraction.input_labels().len() + 1);
211                for index in 0..contraction.input_labels().len() {
212                    call_operands.push(request.primal_input_value(index)?);
213                }
214                call_operands.push(cotangent);
215                builder.add_extension(Arc::new(adjoint), &call_operands)
216            }
217            Df64Op::Expand | Df64Op::QrVjp | Df64Op::QrJvp => {
218                return Err(unsupported(op, role));
219            }
220        }
221        .map_err(SemanticAdError::Build)?;
222
223        Ok(emitted
224            .iter()
225            .enumerate()
226            .map(|(index, value)| {
227                if index < request.primal_input_count() && request.active_inputs()[index] {
228                    AdValue::Value(*value)
229                } else {
230                    AdValue::Absent
231                }
232            })
233            .collect::<Vec<_>>()
234            .into_boxed_slice())
235    }
236}
237
238/// Forward-mode rule for the contribution's operations.
239///
240/// The total sum is linear, so its tangent output is the sum of the tangent inputs,
241/// and each conversion's tangent passes through the same conversion.
242///
243/// # Examples
244///
245/// ```rust
246/// use tenferro_ad::semantic_extension::{SemanticExtensionRuleSet, SemanticLinearizeRule};
247/// use tenferro_df64_proof::ad::Df64LinearizeRule;
248///
249/// let rules = SemanticExtensionRuleSet::new()
250///     .with_linearize(std::sync::Arc::new(Df64LinearizeRule))
251///     .expect("one linearize rule per family");
252/// assert!(rules.lookup_linearize("tenferro-df64-proof.df64_ops.v1").is_some());
253/// ```
254#[derive(Debug)]
255pub struct Df64LinearizeRule;
256
257impl SemanticLinearizeRule for Df64LinearizeRule {
258    fn family_id(&self) -> &'static str {
259        DF64_OPS_FAMILY
260    }
261
262    fn linearize(
263        &self,
264        request: SemanticLinearizeRequest<'_>,
265        builder: &mut SemanticProgramBuilder,
266    ) -> Result<SemanticLinearizeResult, SemanticAdError> {
267        let role = SemanticAdRuleRole::Linearize;
268        let op = Df64Op::of(request.op()).ok_or_else(|| SemanticAdError::Rule {
269            family_id: DF64_OPS_FAMILY,
270            role,
271            source: Box::new(std::io::Error::other(
272                "the payload is not one of the Df64 operations",
273            )),
274        })?;
275        let (operation, tangents): (Arc<dyn ExtensionOp>, Vec<_>) = match op {
276            Df64Op::Total => (
277                Arc::new(Df64Total),
278                request
279                    .tangent_inputs()
280                    .iter()
281                    .filter_map(|value| value.value())
282                    .collect(),
283            ),
284            Df64Op::ToF64 | Df64Op::FromF64 => {
285                let tangent = request
286                    .tangent_inputs()
287                    .iter()
288                    .find_map(|value| value.value());
289                let operation: Arc<dyn ExtensionOp> = if op == Df64Op::ToF64 {
290                    Arc::new(Df64ToF64)
291                } else {
292                    Arc::new(Df64FromF64)
293                };
294                (operation, tangent.into_iter().collect())
295            }
296            Df64Op::Qr => {
297                // The tangent of the factorization is its own body: it needs both primal
298                // factors and the input tangent, so the rule emits one operation.
299                let Some(tangent) = request
300                    .tangent_inputs()
301                    .first()
302                    .and_then(|value| value.value())
303                else {
304                    let inactive = (0..request.primal_outputs().len())
305                        .map(|_| AdValue::Absent)
306                        .collect::<Vec<_>>();
307                    return Ok(SemanticLinearizeResult::new(inactive, Vec::new()));
308                };
309                (
310                    Arc::new(Df64QrJvp) as Arc<dyn ExtensionOp>,
311                    vec![
312                        request.primal_outputs()[0],
313                        request.primal_outputs()[1],
314                        tangent,
315                    ],
316                )
317            }
318            Df64Op::Einsum => {
319                // The tangent of a contraction contracts each tangent with the other operand, so
320                // the rule emits one helper carrying both operands and whichever tangents exist.
321                let Some(contraction) = request.op().as_any().downcast_ref::<Df64Einsum>() else {
322                    return Err(unsupported(op, role));
323                };
324                let has_lhs = request
325                    .tangent_inputs()
326                    .first()
327                    .and_then(|value| value.value())
328                    .is_some();
329                let has_rhs = request
330                    .tangent_inputs()
331                    .get(1)
332                    .and_then(|value| value.value())
333                    .is_some();
334                if !has_lhs && !has_rhs {
335                    let inactive = (0..request.primal_inputs().len())
336                        .map(|_| AdValue::Absent)
337                        .collect::<Vec<_>>();
338                    return Ok(SemanticLinearizeResult::new(inactive, Vec::new()));
339                }
340                // The helpers are defined for the pairwise case, so a wider pattern is refused
341                // rather than differentiated as if it were pairwise.
342                // One tangent flag per operand, which is what the helper carries.
343                let mask: Vec<bool> = request
344                    .tangent_inputs()
345                    .iter()
346                    .map(|value| value.value().is_some())
347                    .collect();
348                let operand_labels: Vec<&[u32]> = contraction
349                    .input_labels()
350                    .iter()
351                    .map(|labels| labels.as_slice())
352                    .collect();
353                let Ok(tangent) =
354                    Df64EinsumJvp::of(&operand_labels, contraction.out_labels(), &mask)
355                else {
356                    return Err(unsupported(op, role));
357                };
358                let mut operands: Vec<_> = request.primal_inputs().to_vec();
359                for value in request.tangent_inputs() {
360                    if let Some(value) = value.value() {
361                        operands.push(value);
362                    }
363                }
364                (Arc::new(tangent) as Arc<dyn ExtensionOp>, operands)
365            }
366            Df64Op::Expand | Df64Op::QrVjp | Df64Op::QrJvp => {
367                return Err(unsupported(op, role));
368            }
369        };
370        if tangents.is_empty() {
371            let inactive = (0..request.primal_outputs().len())
372                .map(|_| AdValue::Absent)
373                .collect::<Vec<_>>();
374            return Ok(SemanticLinearizeResult::new(inactive, Vec::new()));
375        }
376        let emitted = builder
377            .add_extension(operation, &tangents)
378            .map_err(SemanticAdError::Build)?;
379        let outputs = emitted
380            .iter()
381            .enumerate()
382            .map(|(index, value)| {
383                if request
384                    .active_outputs()
385                    .get(index)
386                    .copied()
387                    .unwrap_or(false)
388                {
389                    AdValue::Value(*value)
390                } else {
391                    AdValue::Absent
392                }
393            })
394            .collect::<Vec<_>>();
395        Ok(SemanticLinearizeResult::new(outputs, Vec::new()))
396    }
397}
398
399/// Read the primal input's shape as concrete extents.
400///
401/// # Errors
402///
403/// Returns [`SemanticAdError::Invariant`] when the shape is not fully concrete,
404/// because a broadcast of an unknown extent has no declared shape for the op payload.
405fn exact_shape(request: &SemanticPrimalVjpRequest<'_>) -> Result<Vec<usize>, SemanticAdError> {
406    let metadata = request.primal_input_meta(0)?;
407    metadata
408        .shape()
409        .iter()
410        .map(|extent| match extent.as_exact() {
411            Some(tenferro_ops::dim_expr::DimExpr::Const(value)) => Ok(*value),
412            _ => Err(SemanticAdError::Invariant {
413                family_id: DF64_OPS_FAMILY,
414                role: SemanticAdRuleRole::PrimalVjp,
415                message: format!(
416                    "the adjoint of {DF64_OPS_FAMILY} needs a concrete input shape, got {extent:?}"
417                ),
418            }),
419        })
420        .collect()
421}