Skip to main content

tenferro_linalg/ad/
support.rs

1#[cfg(test)]
2use crate::extension::LinalgOp;
3
4/// AD rule support status for a linalg operation or output.
5///
6/// # Examples
7///
8/// ```rust
9/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind, LinalgAdRuleSupport};
10///
11/// let svd = linalg_ad_support(LinalgAdOpKind::Svd);
12/// assert_eq!(svd.linearize, LinalgAdRuleSupport::SupportedViaLinearize);
13/// ```
14#[derive(Clone, Copy, Debug, PartialEq, Eq)]
15pub enum LinalgAdRuleSupport {
16    Supported,
17    SupportedViaLinearize,
18    PartiallySupported,
19    NonDifferentiable,
20    Unsupported,
21    PendingOracle,
22}
23
24/// Implementation route used for a user-visible AD mode.
25///
26/// # Examples
27///
28/// ```rust
29/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind, LinalgAdRoute};
30///
31/// let svd = linalg_ad_support(LinalgAdOpKind::Svd);
32/// assert_eq!(svd.vjp.route, LinalgAdRoute::LinearizeThenTranspose);
33/// ```
34#[derive(Clone, Copy, Debug, PartialEq, Eq)]
35pub enum LinalgAdRoute {
36    /// No supported route exists.
37    Unsupported,
38    /// The mode is emitted directly by the operation's linearize rule.
39    Linearize,
40    /// Reverse mode is supported by linearizing the primal graph and then
41    /// transposing that linear graph.
42    LinearizeThenTranspose,
43    /// Reverse mode is supported by linearizing the primal graph and using a
44    /// custom transposed-linear rule for the emitted linear operation.
45    LinearizeThenCustomLinearTranspose,
46    /// Reverse mode is emitted by a direct operation-specific primal VJP rule.
47    CustomVjp,
48    /// A custom VJP is preferred, with the canonical linearize-then-transpose
49    /// route retained as an intentional not-applicable fallback.
50    CustomPreferredWithLinearizeFallback,
51}
52
53/// User-visible AD support for one mode.
54///
55/// # Examples
56///
57/// ```rust
58/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind, LinalgAdRoute};
59///
60/// let solve = linalg_ad_support(LinalgAdOpKind::TriangularSolve);
61/// assert_eq!(solve.vjp.route, LinalgAdRoute::LinearizeThenCustomLinearTranspose);
62/// ```
63#[derive(Clone, Copy, Debug, PartialEq, Eq)]
64pub struct LinalgAdModeSupport {
65    /// Whether the user-visible mode is supported.
66    pub status: LinalgAdRuleSupport,
67    /// How the mode is implemented.
68    pub route: LinalgAdRoute,
69}
70
71/// Operation keys covered by the linalg AD support manifest.
72///
73/// # Examples
74///
75/// ```rust
76/// use tenferro_linalg::LinalgAdOpKind;
77///
78/// assert!(LinalgAdOpKind::Svd.as_index() < LinalgAdOpKind::COUNT);
79/// ```
80#[derive(Clone, Copy, Debug, PartialEq, Eq)]
81pub enum LinalgAdOpKind {
82    Cholesky,
83    Lu,
84    LuFactor,
85    SignDetFromLuFactor,
86    LogAbsDetFromLuFactor,
87    LuSolvePrepared,
88    FullPivLu,
89    FullPivLuSolve,
90    Svd,
91    SvdVals,
92    Qr,
93    Eigh,
94    EighVals,
95    Eig,
96    EigVals,
97    TriangularSolve,
98    SvdFull,
99    HouseholderQrFactor,
100    HouseholderQrFromFactors,
101    HouseholderQrAppend,
102    HouseholderQrR,
103    HouseholderQrQColumns,
104    HouseholderQrThinQ,
105    HouseholderQrAppendTangent,
106    HouseholderQrSplitTangent,
107    RankRevealingQr,
108}
109
110impl LinalgAdOpKind {
111    pub const COUNT: usize = 26;
112
113    /// Return the manifest index for this operation kind.
114    ///
115    /// # Examples
116    ///
117    /// ```rust
118    /// use tenferro_linalg::LinalgAdOpKind;
119    ///
120    /// assert_eq!(LinalgAdOpKind::Cholesky.as_index(), 0);
121    /// ```
122    pub const fn as_index(self) -> usize {
123        match self {
124            Self::Cholesky => 0,
125            Self::Lu => 1,
126            Self::LuFactor => 2,
127            Self::SignDetFromLuFactor => 3,
128            Self::LogAbsDetFromLuFactor => 4,
129            Self::LuSolvePrepared => 5,
130            Self::FullPivLu => 6,
131            Self::FullPivLuSolve => 7,
132            Self::Svd => 8,
133            Self::SvdVals => 9,
134            Self::Qr => 10,
135            Self::Eigh => 11,
136            Self::EighVals => 12,
137            Self::Eig => 13,
138            Self::EigVals => 14,
139            Self::TriangularSolve => 15,
140            Self::SvdFull => 16,
141            Self::HouseholderQrFactor => 17,
142            Self::HouseholderQrFromFactors => 18,
143            Self::HouseholderQrAppend => 19,
144            Self::HouseholderQrR => 20,
145            Self::HouseholderQrQColumns => 21,
146            Self::HouseholderQrThinQ => 22,
147            Self::HouseholderQrAppendTangent => 23,
148            Self::HouseholderQrSplitTangent => 24,
149            Self::RankRevealingQr => 25,
150        }
151    }
152
153    #[cfg(test)]
154    pub(crate) const fn from_linalg_op(op: LinalgOp) -> Self {
155        match op {
156            LinalgOp::Cholesky => Self::Cholesky,
157            LinalgOp::Lu => Self::Lu,
158            LinalgOp::LuFactor => Self::LuFactor,
159            LinalgOp::SignDetFromLuFactor => Self::SignDetFromLuFactor,
160            LinalgOp::LogAbsDetFromLuFactor => Self::LogAbsDetFromLuFactor,
161            LinalgOp::LuSolvePrepared { .. } => Self::LuSolvePrepared,
162            LinalgOp::FullPivLu => Self::FullPivLu,
163            LinalgOp::FullPivLuSolve { .. } => Self::FullPivLuSolve,
164            // The partial-pivot single-op solve shares the solve-family AD
165            // route (linearize + custom linear transpose) with FullPivLuSolve
166            // and is not separately exposed in the public manifest.
167            LinalgOp::Solve => Self::FullPivLuSolve,
168            LinalgOp::Svd { .. } => Self::Svd,
169            LinalgOp::SvdFull => Self::SvdFull,
170            LinalgOp::SvdVals { .. } => Self::SvdVals,
171            LinalgOp::Qr { .. } => Self::Qr,
172            LinalgOp::RankRevealingQr { .. } => Self::RankRevealingQr,
173            LinalgOp::HouseholderQrFactor => Self::HouseholderQrFactor,
174            LinalgOp::HouseholderQrFromFactors => Self::HouseholderQrFromFactors,
175            LinalgOp::HouseholderQrAppend => Self::HouseholderQrAppend,
176            LinalgOp::HouseholderQrR { .. } => Self::HouseholderQrR,
177            LinalgOp::HouseholderQrQColumns { .. } => Self::HouseholderQrQColumns,
178            LinalgOp::HouseholderQrThinQ { .. } => Self::HouseholderQrThinQ,
179            LinalgOp::HouseholderQrAppendTangent => Self::HouseholderQrAppendTangent,
180            LinalgOp::HouseholderQrSplitTangent { .. } => Self::HouseholderQrSplitTangent,
181            LinalgOp::Eigh { .. } => Self::Eigh,
182            LinalgOp::EighVals { .. } => Self::EighVals,
183            LinalgOp::Eig { .. } => Self::Eig,
184            LinalgOp::EigVals { .. } => Self::EigVals,
185            LinalgOp::TriangularSolve { .. } => Self::TriangularSolve,
186        }
187    }
188}
189
190/// AD support status for one output of a linalg operation.
191///
192/// # Examples
193///
194/// ```rust
195/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind, LinalgAdRuleSupport};
196///
197/// let full_piv_lu = linalg_ad_support(LinalgAdOpKind::FullPivLu);
198/// let l_output = full_piv_lu.outputs.iter().find(|output| output.name == "l").unwrap();
199/// assert_eq!(l_output.status, LinalgAdRuleSupport::SupportedViaLinearize);
200/// ```
201#[derive(Clone, Copy, Debug, PartialEq, Eq)]
202pub struct LinalgAdOutputSupport {
203    /// Output position in the linalg operation result tuple.
204    pub index: usize,
205    /// Stable output name used by tests and support dashboards.
206    pub name: &'static str,
207    /// AD support status for this specific output.
208    pub status: LinalgAdRuleSupport,
209}
210
211/// AD support manifest entry for one linalg operation.
212///
213/// # Examples
214///
215/// ```rust
216/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind, LinalgAdRuleSupport};
217///
218/// let solve = linalg_ad_support(LinalgAdOpKind::TriangularSolve);
219/// assert_eq!(solve.vjp.route, tenferro_linalg::LinalgAdRoute::LinearizeThenCustomLinearTranspose);
220/// ```
221#[derive(Clone, Copy, Debug, PartialEq, Eq)]
222pub struct LinalgAdSupport {
223    /// Operation kind described by this manifest entry.
224    pub kind: LinalgAdOpKind,
225    /// User-visible JVP support and route.
226    pub jvp: LinalgAdModeSupport,
227    /// User-visible VJP support and route.
228    pub vjp: LinalgAdModeSupport,
229    /// Definitional linearize rule implementation status.
230    pub linearize_rule: LinalgAdRuleSupport,
231    /// Direct primal VJP rule implementation status.
232    pub custom_vjp_rule: LinalgAdRuleSupport,
233    /// Custom transposed-linear rule implementation status.
234    pub custom_linear_transpose_rule: LinalgAdRuleSupport,
235    /// Forward-mode graph emission support.
236    pub linearize: LinalgAdRuleSupport,
237    /// Transposed-linear graph emission support.
238    pub transpose: LinalgAdRuleSupport,
239    /// Per-output support status for multi-output operations.
240    pub outputs: &'static [LinalgAdOutputSupport],
241    /// Numerical or semantic caveats for this operation family.
242    pub caveats: &'static [&'static str],
243}
244
245const fn mode(status: LinalgAdRuleSupport, route: LinalgAdRoute) -> LinalgAdModeSupport {
246    LinalgAdModeSupport { status, route }
247}
248
249const fn jvp_route(status: LinalgAdRuleSupport) -> LinalgAdRoute {
250    match status {
251        LinalgAdRuleSupport::Unsupported
252        | LinalgAdRuleSupport::NonDifferentiable
253        | LinalgAdRuleSupport::PendingOracle => LinalgAdRoute::Unsupported,
254        LinalgAdRuleSupport::Supported
255        | LinalgAdRuleSupport::SupportedViaLinearize
256        | LinalgAdRuleSupport::PartiallySupported => LinalgAdRoute::Linearize,
257    }
258}
259
260const fn support_entry(
261    kind: LinalgAdOpKind,
262    linearize: LinalgAdRuleSupport,
263    transpose: LinalgAdRuleSupport,
264    vjp: LinalgAdModeSupport,
265    custom_linear_transpose_rule: LinalgAdRuleSupport,
266    outputs: &'static [LinalgAdOutputSupport],
267    caveats: &'static [&'static str],
268) -> LinalgAdSupport {
269    LinalgAdSupport {
270        kind,
271        jvp: mode(linearize, jvp_route(linearize)),
272        vjp,
273        linearize_rule: linearize,
274        custom_vjp_rule: LinalgAdRuleSupport::Unsupported,
275        custom_linear_transpose_rule,
276        linearize,
277        transpose,
278        outputs,
279        caveats,
280    }
281}
282
283const fn output(
284    index: usize,
285    name: &'static str,
286    status: LinalgAdRuleSupport,
287) -> LinalgAdOutputSupport {
288    LinalgAdOutputSupport {
289        index,
290        name,
291        status,
292    }
293}
294
295static CHOLESKY_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
296    0,
297    "factor",
298    LinalgAdRuleSupport::SupportedViaLinearize,
299)];
300static LU_OUTPUTS: [LinalgAdOutputSupport; 4] = [
301    output(0, "p", LinalgAdRuleSupport::NonDifferentiable),
302    output(1, "l", LinalgAdRuleSupport::SupportedViaLinearize),
303    output(2, "u", LinalgAdRuleSupport::SupportedViaLinearize),
304    output(3, "parity", LinalgAdRuleSupport::NonDifferentiable),
305];
306static LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 3] = [
307    output(0, "packed_lu", LinalgAdRuleSupport::Unsupported),
308    output(1, "pivots", LinalgAdRuleSupport::NonDifferentiable),
309    output(2, "parity", LinalgAdRuleSupport::NonDifferentiable),
310];
311static SIGNDET_FROM_LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
312    0,
313    "sign",
314    LinalgAdRuleSupport::SupportedViaLinearize,
315)];
316static LOGABSDET_FROM_LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
317    0,
318    "logabsdet",
319    LinalgAdRuleSupport::SupportedViaLinearize,
320)];
321static SOLUTION_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
322    0,
323    "solution",
324    LinalgAdRuleSupport::SupportedViaLinearize,
325)];
326static FULL_PIV_LU_OUTPUTS: [LinalgAdOutputSupport; 5] = [
327    output(0, "p", LinalgAdRuleSupport::NonDifferentiable),
328    output(1, "l", LinalgAdRuleSupport::SupportedViaLinearize),
329    output(2, "u", LinalgAdRuleSupport::SupportedViaLinearize),
330    output(3, "q", LinalgAdRuleSupport::NonDifferentiable),
331    output(4, "parity", LinalgAdRuleSupport::NonDifferentiable),
332];
333static FULL_PIV_LU_SOLVE_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
334    0,
335    "solution",
336    LinalgAdRuleSupport::SupportedViaLinearize,
337)];
338static SVD_OUTPUTS: [LinalgAdOutputSupport; 3] = [
339    output(0, "u", LinalgAdRuleSupport::SupportedViaLinearize),
340    output(
341        1,
342        "singular_values",
343        LinalgAdRuleSupport::SupportedViaLinearize,
344    ),
345    output(2, "vt", LinalgAdRuleSupport::SupportedViaLinearize),
346];
347static SVD_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
348    0,
349    "singular_values",
350    LinalgAdRuleSupport::SupportedViaLinearize,
351)];
352static SVD_FULL_OUTPUTS: [LinalgAdOutputSupport; 3] = [
353    output(0, "u", LinalgAdRuleSupport::Unsupported),
354    output(1, "singular_values", LinalgAdRuleSupport::Unsupported),
355    output(2, "vt", LinalgAdRuleSupport::Unsupported),
356];
357static QR_OUTPUTS: [LinalgAdOutputSupport; 2] = [
358    output(0, "q", LinalgAdRuleSupport::SupportedViaLinearize),
359    output(1, "r", LinalgAdRuleSupport::SupportedViaLinearize),
360];
361static RANK_REVEALING_QR_OUTPUTS: [LinalgAdOutputSupport; 4] = [
362    output(0, "q", LinalgAdRuleSupport::Unsupported),
363    output(1, "r", LinalgAdRuleSupport::Unsupported),
364    output(2, "column_permutation", LinalgAdRuleSupport::Unsupported),
365    output(3, "rank", LinalgAdRuleSupport::Unsupported),
366];
367static HOUSEHOLDER_QR_STATE_OUTPUTS: [LinalgAdOutputSupport; 2] = [
368    output(0, "packed", LinalgAdRuleSupport::SupportedViaLinearize),
369    output(1, "coeff", LinalgAdRuleSupport::NonDifferentiable),
370];
371static HOUSEHOLDER_QR_VALUE_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
372    0,
373    "value",
374    LinalgAdRuleSupport::SupportedViaLinearize,
375)];
376static HOUSEHOLDER_QR_RESIDUAL_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
377    0,
378    "internal_value",
379    LinalgAdRuleSupport::Unsupported,
380)];
381static HOUSEHOLDER_QR_CAVEATS: [&str; 1] =
382    ["Rank-deficient states are outside the differentiable domain."];
383static EIGH_OUTPUTS: [LinalgAdOutputSupport; 2] = [
384    output(0, "eigenvalues", LinalgAdRuleSupport::SupportedViaLinearize),
385    output(
386        1,
387        "eigenvectors",
388        LinalgAdRuleSupport::SupportedViaLinearize,
389    ),
390];
391static EIGH_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
392    0,
393    "eigenvalues",
394    LinalgAdRuleSupport::SupportedViaLinearize,
395)];
396static EIG_OUTPUTS: [LinalgAdOutputSupport; 2] = [
397    output(0, "eigenvalues", LinalgAdRuleSupport::SupportedViaLinearize),
398    output(1, "eigenvectors", LinalgAdRuleSupport::Unsupported),
399];
400static EIG_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
401    0,
402    "eigenvalues",
403    LinalgAdRuleSupport::SupportedViaLinearize,
404)];
405
406static DECOMPOSITION_CAVEATS: [&str; 1] = [
407    "Derivative regularization handles near-degenerate spectra but does not make exact degeneracies smoothly differentiable.",
408];
409
410static LINALG_AD_SUPPORT: [LinalgAdSupport; LinalgAdOpKind::COUNT] = [
411    support_entry(
412        LinalgAdOpKind::Cholesky,
413        LinalgAdRuleSupport::SupportedViaLinearize,
414        LinalgAdRuleSupport::Unsupported,
415        mode(
416            LinalgAdRuleSupport::SupportedViaLinearize,
417            LinalgAdRoute::LinearizeThenTranspose,
418        ),
419        LinalgAdRuleSupport::Unsupported,
420        &CHOLESKY_OUTPUTS,
421        &DECOMPOSITION_CAVEATS,
422    ),
423    support_entry(
424        LinalgAdOpKind::Lu,
425        LinalgAdRuleSupport::PartiallySupported,
426        LinalgAdRuleSupport::Unsupported,
427        mode(
428            LinalgAdRuleSupport::PartiallySupported,
429            LinalgAdRoute::LinearizeThenTranspose,
430        ),
431        LinalgAdRuleSupport::Unsupported,
432        &LU_OUTPUTS,
433        &DECOMPOSITION_CAVEATS,
434    ),
435    support_entry(
436        LinalgAdOpKind::LuFactor,
437        LinalgAdRuleSupport::Unsupported,
438        LinalgAdRuleSupport::Unsupported,
439        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
440        LinalgAdRuleSupport::Unsupported,
441        &LU_FACTOR_OUTPUTS,
442        &[],
443    ),
444    support_entry(
445        LinalgAdOpKind::SignDetFromLuFactor,
446        LinalgAdRuleSupport::SupportedViaLinearize,
447        LinalgAdRuleSupport::Unsupported,
448        mode(
449            LinalgAdRuleSupport::SupportedViaLinearize,
450            LinalgAdRoute::LinearizeThenTranspose,
451        ),
452        LinalgAdRuleSupport::Unsupported,
453        &SIGNDET_FROM_LU_FACTOR_OUTPUTS,
454        &[],
455    ),
456    support_entry(
457        LinalgAdOpKind::LogAbsDetFromLuFactor,
458        LinalgAdRuleSupport::SupportedViaLinearize,
459        LinalgAdRuleSupport::Unsupported,
460        mode(
461            LinalgAdRuleSupport::SupportedViaLinearize,
462            LinalgAdRoute::LinearizeThenTranspose,
463        ),
464        LinalgAdRuleSupport::Unsupported,
465        &LOGABSDET_FROM_LU_FACTOR_OUTPUTS,
466        &[],
467    ),
468    support_entry(
469        LinalgAdOpKind::LuSolvePrepared,
470        LinalgAdRuleSupport::SupportedViaLinearize,
471        LinalgAdRuleSupport::PartiallySupported,
472        mode(
473            LinalgAdRuleSupport::SupportedViaLinearize,
474            LinalgAdRoute::LinearizeThenCustomLinearTranspose,
475        ),
476        LinalgAdRuleSupport::PartiallySupported,
477        &SOLUTION_OUTPUTS,
478        &[],
479    ),
480    support_entry(
481        LinalgAdOpKind::FullPivLu,
482        LinalgAdRuleSupport::SupportedViaLinearize,
483        LinalgAdRuleSupport::Unsupported,
484        mode(
485            LinalgAdRuleSupport::SupportedViaLinearize,
486            LinalgAdRoute::LinearizeThenTranspose,
487        ),
488        LinalgAdRuleSupport::Unsupported,
489        &FULL_PIV_LU_OUTPUTS,
490        &DECOMPOSITION_CAVEATS,
491    ),
492    support_entry(
493        LinalgAdOpKind::FullPivLuSolve,
494        LinalgAdRuleSupport::SupportedViaLinearize,
495        LinalgAdRuleSupport::Supported,
496        mode(
497            LinalgAdRuleSupport::SupportedViaLinearize,
498            LinalgAdRoute::LinearizeThenCustomLinearTranspose,
499        ),
500        LinalgAdRuleSupport::Supported,
501        &FULL_PIV_LU_SOLVE_OUTPUTS,
502        &[],
503    ),
504    support_entry(
505        LinalgAdOpKind::Svd,
506        LinalgAdRuleSupport::SupportedViaLinearize,
507        LinalgAdRuleSupport::Unsupported,
508        mode(
509            LinalgAdRuleSupport::SupportedViaLinearize,
510            LinalgAdRoute::LinearizeThenTranspose,
511        ),
512        LinalgAdRuleSupport::Unsupported,
513        &SVD_OUTPUTS,
514        &DECOMPOSITION_CAVEATS,
515    ),
516    support_entry(
517        LinalgAdOpKind::SvdVals,
518        LinalgAdRuleSupport::SupportedViaLinearize,
519        LinalgAdRuleSupport::Unsupported,
520        mode(
521            LinalgAdRuleSupport::SupportedViaLinearize,
522            LinalgAdRoute::LinearizeThenTranspose,
523        ),
524        LinalgAdRuleSupport::Unsupported,
525        &SVD_VALS_OUTPUTS,
526        &DECOMPOSITION_CAVEATS,
527    ),
528    support_entry(
529        LinalgAdOpKind::Qr,
530        LinalgAdRuleSupport::SupportedViaLinearize,
531        LinalgAdRuleSupport::Unsupported,
532        mode(
533            LinalgAdRuleSupport::SupportedViaLinearize,
534            LinalgAdRoute::LinearizeThenTranspose,
535        ),
536        LinalgAdRuleSupport::Unsupported,
537        &QR_OUTPUTS,
538        &DECOMPOSITION_CAVEATS,
539    ),
540    support_entry(
541        LinalgAdOpKind::Eigh,
542        LinalgAdRuleSupport::SupportedViaLinearize,
543        LinalgAdRuleSupport::Unsupported,
544        mode(
545            LinalgAdRuleSupport::SupportedViaLinearize,
546            LinalgAdRoute::LinearizeThenTranspose,
547        ),
548        LinalgAdRuleSupport::Unsupported,
549        &EIGH_OUTPUTS,
550        &DECOMPOSITION_CAVEATS,
551    ),
552    support_entry(
553        LinalgAdOpKind::EighVals,
554        LinalgAdRuleSupport::SupportedViaLinearize,
555        LinalgAdRuleSupport::Unsupported,
556        mode(
557            LinalgAdRuleSupport::SupportedViaLinearize,
558            LinalgAdRoute::LinearizeThenTranspose,
559        ),
560        LinalgAdRuleSupport::Unsupported,
561        &EIGH_VALS_OUTPUTS,
562        &DECOMPOSITION_CAVEATS,
563    ),
564    support_entry(
565        LinalgAdOpKind::Eig,
566        LinalgAdRuleSupport::PartiallySupported,
567        LinalgAdRuleSupport::Unsupported,
568        mode(
569            LinalgAdRuleSupport::PartiallySupported,
570            LinalgAdRoute::LinearizeThenTranspose,
571        ),
572        LinalgAdRuleSupport::Unsupported,
573        &EIG_OUTPUTS,
574        &DECOMPOSITION_CAVEATS,
575    ),
576    support_entry(
577        LinalgAdOpKind::EigVals,
578        LinalgAdRuleSupport::SupportedViaLinearize,
579        LinalgAdRuleSupport::Unsupported,
580        mode(
581            LinalgAdRuleSupport::SupportedViaLinearize,
582            LinalgAdRoute::LinearizeThenTranspose,
583        ),
584        LinalgAdRuleSupport::Unsupported,
585        &EIG_VALS_OUTPUTS,
586        &DECOMPOSITION_CAVEATS,
587    ),
588    support_entry(
589        LinalgAdOpKind::TriangularSolve,
590        LinalgAdRuleSupport::SupportedViaLinearize,
591        LinalgAdRuleSupport::Supported,
592        mode(
593            LinalgAdRuleSupport::SupportedViaLinearize,
594            LinalgAdRoute::LinearizeThenCustomLinearTranspose,
595        ),
596        LinalgAdRuleSupport::Supported,
597        &SOLUTION_OUTPUTS,
598        &[],
599    ),
600    // Full-matrices SVD is a value-only route: nullspace/kernel extraction does
601    // not require derivatives, and the thin-SVD linearize rule does not extend
602    // to the square factors. AD is intentionally unsupported for every output.
603    support_entry(
604        LinalgAdOpKind::SvdFull,
605        LinalgAdRuleSupport::Unsupported,
606        LinalgAdRuleSupport::Unsupported,
607        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
608        LinalgAdRuleSupport::Unsupported,
609        &SVD_FULL_OUTPUTS,
610        &[],
611    ),
612    support_entry(
613        LinalgAdOpKind::HouseholderQrFactor,
614        LinalgAdRuleSupport::SupportedViaLinearize,
615        LinalgAdRuleSupport::Unsupported,
616        mode(
617            LinalgAdRuleSupport::SupportedViaLinearize,
618            LinalgAdRoute::LinearizeThenTranspose,
619        ),
620        LinalgAdRuleSupport::Unsupported,
621        &HOUSEHOLDER_QR_STATE_OUTPUTS,
622        &HOUSEHOLDER_QR_CAVEATS,
623    ),
624    support_entry(
625        LinalgAdOpKind::HouseholderQrFromFactors,
626        LinalgAdRuleSupport::SupportedViaLinearize,
627        LinalgAdRuleSupport::Unsupported,
628        mode(
629            LinalgAdRuleSupport::SupportedViaLinearize,
630            LinalgAdRoute::LinearizeThenTranspose,
631        ),
632        LinalgAdRuleSupport::Unsupported,
633        &HOUSEHOLDER_QR_STATE_OUTPUTS,
634        &HOUSEHOLDER_QR_CAVEATS,
635    ),
636    support_entry(
637        LinalgAdOpKind::HouseholderQrAppend,
638        LinalgAdRuleSupport::SupportedViaLinearize,
639        LinalgAdRuleSupport::Unsupported,
640        mode(
641            LinalgAdRuleSupport::SupportedViaLinearize,
642            LinalgAdRoute::LinearizeThenTranspose,
643        ),
644        LinalgAdRuleSupport::Unsupported,
645        &HOUSEHOLDER_QR_STATE_OUTPUTS,
646        &HOUSEHOLDER_QR_CAVEATS,
647    ),
648    support_entry(
649        LinalgAdOpKind::HouseholderQrR,
650        LinalgAdRuleSupport::SupportedViaLinearize,
651        LinalgAdRuleSupport::Unsupported,
652        mode(
653            LinalgAdRuleSupport::SupportedViaLinearize,
654            LinalgAdRoute::LinearizeThenTranspose,
655        ),
656        LinalgAdRuleSupport::Unsupported,
657        &HOUSEHOLDER_QR_VALUE_OUTPUTS,
658        &HOUSEHOLDER_QR_CAVEATS,
659    ),
660    support_entry(
661        LinalgAdOpKind::HouseholderQrQColumns,
662        LinalgAdRuleSupport::SupportedViaLinearize,
663        LinalgAdRuleSupport::Unsupported,
664        mode(
665            LinalgAdRuleSupport::SupportedViaLinearize,
666            LinalgAdRoute::LinearizeThenTranspose,
667        ),
668        LinalgAdRuleSupport::Unsupported,
669        &HOUSEHOLDER_QR_VALUE_OUTPUTS,
670        &HOUSEHOLDER_QR_CAVEATS,
671    ),
672    support_entry(
673        LinalgAdOpKind::HouseholderQrThinQ,
674        LinalgAdRuleSupport::Unsupported,
675        LinalgAdRuleSupport::Unsupported,
676        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
677        LinalgAdRuleSupport::Unsupported,
678        &HOUSEHOLDER_QR_RESIDUAL_OUTPUTS,
679        &["Internal fixed residual; no public differentiable surface."],
680    ),
681    support_entry(
682        LinalgAdOpKind::HouseholderQrAppendTangent,
683        LinalgAdRuleSupport::Unsupported,
684        LinalgAdRuleSupport::Supported,
685        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
686        LinalgAdRuleSupport::Supported,
687        &HOUSEHOLDER_QR_RESIDUAL_OUTPUTS,
688        &["Internal linear append operation; no public differentiable surface."],
689    ),
690    support_entry(
691        LinalgAdOpKind::HouseholderQrSplitTangent,
692        LinalgAdRuleSupport::Unsupported,
693        LinalgAdRuleSupport::Unsupported,
694        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
695        LinalgAdRuleSupport::Unsupported,
696        &HOUSEHOLDER_QR_RESIDUAL_OUTPUTS,
697        &["Internal transpose residual; no public differentiable surface."],
698    ),
699    support_entry(
700        LinalgAdOpKind::RankRevealingQr,
701        LinalgAdRuleSupport::Unsupported,
702        LinalgAdRuleSupport::Unsupported,
703        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
704        LinalgAdRuleSupport::Unsupported,
705        &RANK_REVEALING_QR_OUTPUTS,
706        &["Pivot selection and numerical rank are discontinuous; all outputs are initially unsupported for AD."],
707    ),
708];
709
710/// Return the complete linalg AD support manifest.
711///
712/// # Examples
713///
714/// ```rust
715/// let manifest = tenferro_linalg::all_linalg_ad_support();
716/// assert_eq!(manifest.len(), tenferro_linalg::LinalgAdOpKind::COUNT);
717/// ```
718pub fn all_linalg_ad_support() -> &'static [LinalgAdSupport; LinalgAdOpKind::COUNT] {
719    &LINALG_AD_SUPPORT
720}
721
722/// Return the support manifest entry for one linalg operation kind.
723///
724/// # Examples
725///
726/// ```rust
727/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind};
728///
729/// let entry = linalg_ad_support(LinalgAdOpKind::Eigh);
730/// assert_eq!(entry.kind, LinalgAdOpKind::Eigh);
731/// ```
732pub fn linalg_ad_support(kind: LinalgAdOpKind) -> &'static LinalgAdSupport {
733    &LINALG_AD_SUPPORT[kind.as_index()]
734}
735
736#[cfg(test)]
737pub(crate) fn linalg_ad_support_for_op(op: LinalgOp) -> &'static LinalgAdSupport {
738    linalg_ad_support(LinalgAdOpKind::from_linalg_op(op))
739}
740
741#[cfg(test)]
742mod tests;