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            // The fused solve is the AD carrier of `solve`, so it reports the
169            // same support entry as the plain solve kernel.
170            LinalgOp::LuFactorSolve => Self::FullPivLuSolve,
171            LinalgOp::Svd { .. } => Self::Svd,
172            LinalgOp::SvdFull => Self::SvdFull,
173            LinalgOp::SvdVals { .. } => Self::SvdVals,
174            LinalgOp::Qr { .. } => Self::Qr,
175            LinalgOp::RankRevealingQr { .. } => Self::RankRevealingQr,
176            LinalgOp::HouseholderQrFactor => Self::HouseholderQrFactor,
177            LinalgOp::HouseholderQrFromFactors => Self::HouseholderQrFromFactors,
178            LinalgOp::HouseholderQrAppend => Self::HouseholderQrAppend,
179            LinalgOp::HouseholderQrR { .. } => Self::HouseholderQrR,
180            LinalgOp::HouseholderQrQColumns { .. } => Self::HouseholderQrQColumns,
181            LinalgOp::HouseholderQrThinQ { .. } => Self::HouseholderQrThinQ,
182            LinalgOp::HouseholderQrAppendTangent => Self::HouseholderQrAppendTangent,
183            LinalgOp::HouseholderQrSplitTangent { .. } => Self::HouseholderQrSplitTangent,
184            LinalgOp::Eigh { .. } => Self::Eigh,
185            LinalgOp::EighVals { .. } => Self::EighVals,
186            LinalgOp::Eig { .. } => Self::Eig,
187            LinalgOp::EigVals { .. } => Self::EigVals,
188            LinalgOp::TriangularSolve { .. } => Self::TriangularSolve,
189        }
190    }
191}
192
193/// AD support status for one output of a linalg operation.
194///
195/// # Examples
196///
197/// ```rust
198/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind, LinalgAdRuleSupport};
199///
200/// let full_piv_lu = linalg_ad_support(LinalgAdOpKind::FullPivLu);
201/// let l_output = full_piv_lu.outputs.iter().find(|output| output.name == "l").unwrap();
202/// assert_eq!(l_output.status, LinalgAdRuleSupport::SupportedViaLinearize);
203/// ```
204#[derive(Clone, Copy, Debug, PartialEq, Eq)]
205pub struct LinalgAdOutputSupport {
206    /// Output position in the linalg operation result tuple.
207    pub index: usize,
208    /// Stable output name used by tests and support dashboards.
209    pub name: &'static str,
210    /// AD support status for this specific output.
211    pub status: LinalgAdRuleSupport,
212}
213
214/// AD support manifest entry for one linalg operation.
215///
216/// # Examples
217///
218/// ```rust
219/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind, LinalgAdRuleSupport};
220///
221/// let solve = linalg_ad_support(LinalgAdOpKind::TriangularSolve);
222/// assert_eq!(solve.vjp.route, tenferro_linalg::LinalgAdRoute::LinearizeThenCustomLinearTranspose);
223/// ```
224#[derive(Clone, Copy, Debug, PartialEq, Eq)]
225pub struct LinalgAdSupport {
226    /// Operation kind described by this manifest entry.
227    pub kind: LinalgAdOpKind,
228    /// User-visible JVP support and route.
229    pub jvp: LinalgAdModeSupport,
230    /// User-visible VJP support and route.
231    pub vjp: LinalgAdModeSupport,
232    /// Definitional linearize rule implementation status.
233    pub linearize_rule: LinalgAdRuleSupport,
234    /// Direct primal VJP rule implementation status.
235    pub custom_vjp_rule: LinalgAdRuleSupport,
236    /// Custom transposed-linear rule implementation status.
237    pub custom_linear_transpose_rule: LinalgAdRuleSupport,
238    /// Forward-mode graph emission support.
239    pub linearize: LinalgAdRuleSupport,
240    /// Transposed-linear graph emission support.
241    pub transpose: LinalgAdRuleSupport,
242    /// Per-output support status for multi-output operations.
243    pub outputs: &'static [LinalgAdOutputSupport],
244    /// Numerical or semantic caveats for this operation family.
245    pub caveats: &'static [&'static str],
246}
247
248const fn mode(status: LinalgAdRuleSupport, route: LinalgAdRoute) -> LinalgAdModeSupport {
249    LinalgAdModeSupport { status, route }
250}
251
252const fn jvp_route(status: LinalgAdRuleSupport) -> LinalgAdRoute {
253    match status {
254        LinalgAdRuleSupport::Unsupported
255        | LinalgAdRuleSupport::NonDifferentiable
256        | LinalgAdRuleSupport::PendingOracle => LinalgAdRoute::Unsupported,
257        LinalgAdRuleSupport::Supported
258        | LinalgAdRuleSupport::SupportedViaLinearize
259        | LinalgAdRuleSupport::PartiallySupported => LinalgAdRoute::Linearize,
260    }
261}
262
263const fn support_entry(
264    kind: LinalgAdOpKind,
265    linearize: LinalgAdRuleSupport,
266    transpose: LinalgAdRuleSupport,
267    vjp: LinalgAdModeSupport,
268    custom_linear_transpose_rule: LinalgAdRuleSupport,
269    outputs: &'static [LinalgAdOutputSupport],
270    caveats: &'static [&'static str],
271) -> LinalgAdSupport {
272    LinalgAdSupport {
273        kind,
274        jvp: mode(linearize, jvp_route(linearize)),
275        vjp,
276        linearize_rule: linearize,
277        custom_vjp_rule: LinalgAdRuleSupport::Unsupported,
278        custom_linear_transpose_rule,
279        linearize,
280        transpose,
281        outputs,
282        caveats,
283    }
284}
285
286const fn output(
287    index: usize,
288    name: &'static str,
289    status: LinalgAdRuleSupport,
290) -> LinalgAdOutputSupport {
291    LinalgAdOutputSupport {
292        index,
293        name,
294        status,
295    }
296}
297
298static CHOLESKY_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
299    0,
300    "factor",
301    LinalgAdRuleSupport::SupportedViaLinearize,
302)];
303static LU_OUTPUTS: [LinalgAdOutputSupport; 4] = [
304    output(0, "p", LinalgAdRuleSupport::NonDifferentiable),
305    output(1, "l", LinalgAdRuleSupport::SupportedViaLinearize),
306    output(2, "u", LinalgAdRuleSupport::SupportedViaLinearize),
307    output(3, "parity", LinalgAdRuleSupport::NonDifferentiable),
308];
309static LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 3] = [
310    output(0, "packed_lu", LinalgAdRuleSupport::Unsupported),
311    output(1, "pivots", LinalgAdRuleSupport::NonDifferentiable),
312    output(2, "parity", LinalgAdRuleSupport::NonDifferentiable),
313];
314static SIGNDET_FROM_LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
315    0,
316    "sign",
317    LinalgAdRuleSupport::SupportedViaLinearize,
318)];
319static LOGABSDET_FROM_LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
320    0,
321    "logabsdet",
322    LinalgAdRuleSupport::SupportedViaLinearize,
323)];
324static SOLUTION_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
325    0,
326    "solution",
327    LinalgAdRuleSupport::SupportedViaLinearize,
328)];
329static FULL_PIV_LU_OUTPUTS: [LinalgAdOutputSupport; 5] = [
330    output(0, "p", LinalgAdRuleSupport::NonDifferentiable),
331    output(1, "l", LinalgAdRuleSupport::SupportedViaLinearize),
332    output(2, "u", LinalgAdRuleSupport::SupportedViaLinearize),
333    output(3, "q", LinalgAdRuleSupport::NonDifferentiable),
334    output(4, "parity", LinalgAdRuleSupport::NonDifferentiable),
335];
336static FULL_PIV_LU_SOLVE_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
337    0,
338    "solution",
339    LinalgAdRuleSupport::SupportedViaLinearize,
340)];
341static SVD_OUTPUTS: [LinalgAdOutputSupport; 3] = [
342    output(0, "u", LinalgAdRuleSupport::SupportedViaLinearize),
343    output(
344        1,
345        "singular_values",
346        LinalgAdRuleSupport::SupportedViaLinearize,
347    ),
348    output(2, "vt", LinalgAdRuleSupport::SupportedViaLinearize),
349];
350static SVD_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
351    0,
352    "singular_values",
353    LinalgAdRuleSupport::SupportedViaLinearize,
354)];
355static SVD_FULL_OUTPUTS: [LinalgAdOutputSupport; 3] = [
356    output(0, "u", LinalgAdRuleSupport::Unsupported),
357    output(1, "singular_values", LinalgAdRuleSupport::Unsupported),
358    output(2, "vt", LinalgAdRuleSupport::Unsupported),
359];
360static QR_OUTPUTS: [LinalgAdOutputSupport; 2] = [
361    output(0, "q", LinalgAdRuleSupport::SupportedViaLinearize),
362    output(1, "r", LinalgAdRuleSupport::SupportedViaLinearize),
363];
364static RANK_REVEALING_QR_OUTPUTS: [LinalgAdOutputSupport; 4] = [
365    output(0, "q", LinalgAdRuleSupport::Unsupported),
366    output(1, "r", LinalgAdRuleSupport::Unsupported),
367    output(2, "column_permutation", LinalgAdRuleSupport::Unsupported),
368    output(3, "rank", LinalgAdRuleSupport::Unsupported),
369];
370static HOUSEHOLDER_QR_STATE_OUTPUTS: [LinalgAdOutputSupport; 2] = [
371    output(0, "packed", LinalgAdRuleSupport::SupportedViaLinearize),
372    output(1, "coeff", LinalgAdRuleSupport::NonDifferentiable),
373];
374static HOUSEHOLDER_QR_VALUE_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
375    0,
376    "value",
377    LinalgAdRuleSupport::SupportedViaLinearize,
378)];
379static HOUSEHOLDER_QR_RESIDUAL_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
380    0,
381    "internal_value",
382    LinalgAdRuleSupport::Unsupported,
383)];
384static HOUSEHOLDER_QR_CAVEATS: [&str; 1] =
385    ["Rank-deficient states are outside the differentiable domain."];
386static HOUSEHOLDER_QR_Q_COLUMNS_CAVEATS: [&str; 2] = [
387    "Rank-deficient states are outside the differentiable domain.",
388    "Column ranges reaching past the thin-Q width k = min(m, n) are \
389     value-only: the complement basis is defined only up to a rotation inside \
390     the nullspace, so the rule returns a typed Unsupported instead of a \
391     derivative.",
392];
393static EIGH_OUTPUTS: [LinalgAdOutputSupport; 2] = [
394    output(0, "eigenvalues", LinalgAdRuleSupport::SupportedViaLinearize),
395    output(
396        1,
397        "eigenvectors",
398        LinalgAdRuleSupport::SupportedViaLinearize,
399    ),
400];
401static EIGH_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
402    0,
403    "eigenvalues",
404    LinalgAdRuleSupport::SupportedViaLinearize,
405)];
406static EIG_OUTPUTS: [LinalgAdOutputSupport; 2] = [
407    output(0, "eigenvalues", LinalgAdRuleSupport::SupportedViaLinearize),
408    output(1, "eigenvectors", LinalgAdRuleSupport::Unsupported),
409];
410static EIG_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
411    0,
412    "eigenvalues",
413    LinalgAdRuleSupport::SupportedViaLinearize,
414)];
415
416static DECOMPOSITION_CAVEATS: [&str; 1] = [
417    "Derivative regularization handles near-degenerate spectra but does not make exact degeneracies smoothly differentiable.",
418];
419
420static LINALG_AD_SUPPORT: [LinalgAdSupport; LinalgAdOpKind::COUNT] = [
421    support_entry(
422        LinalgAdOpKind::Cholesky,
423        LinalgAdRuleSupport::SupportedViaLinearize,
424        LinalgAdRuleSupport::Unsupported,
425        mode(
426            LinalgAdRuleSupport::SupportedViaLinearize,
427            LinalgAdRoute::LinearizeThenTranspose,
428        ),
429        LinalgAdRuleSupport::Unsupported,
430        &CHOLESKY_OUTPUTS,
431        &DECOMPOSITION_CAVEATS,
432    ),
433    support_entry(
434        LinalgAdOpKind::Lu,
435        LinalgAdRuleSupport::PartiallySupported,
436        LinalgAdRuleSupport::Unsupported,
437        mode(
438            LinalgAdRuleSupport::PartiallySupported,
439            LinalgAdRoute::LinearizeThenTranspose,
440        ),
441        LinalgAdRuleSupport::Unsupported,
442        &LU_OUTPUTS,
443        &DECOMPOSITION_CAVEATS,
444    ),
445    support_entry(
446        LinalgAdOpKind::LuFactor,
447        LinalgAdRuleSupport::Unsupported,
448        LinalgAdRuleSupport::Unsupported,
449        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
450        LinalgAdRuleSupport::Unsupported,
451        &LU_FACTOR_OUTPUTS,
452        &[],
453    ),
454    support_entry(
455        LinalgAdOpKind::SignDetFromLuFactor,
456        LinalgAdRuleSupport::SupportedViaLinearize,
457        LinalgAdRuleSupport::Unsupported,
458        mode(
459            LinalgAdRuleSupport::SupportedViaLinearize,
460            LinalgAdRoute::LinearizeThenTranspose,
461        ),
462        LinalgAdRuleSupport::Unsupported,
463        &SIGNDET_FROM_LU_FACTOR_OUTPUTS,
464        &[],
465    ),
466    support_entry(
467        LinalgAdOpKind::LogAbsDetFromLuFactor,
468        LinalgAdRuleSupport::SupportedViaLinearize,
469        LinalgAdRuleSupport::Unsupported,
470        mode(
471            LinalgAdRuleSupport::SupportedViaLinearize,
472            LinalgAdRoute::LinearizeThenTranspose,
473        ),
474        LinalgAdRuleSupport::Unsupported,
475        &LOGABSDET_FROM_LU_FACTOR_OUTPUTS,
476        &[],
477    ),
478    support_entry(
479        LinalgAdOpKind::LuSolvePrepared,
480        LinalgAdRuleSupport::SupportedViaLinearize,
481        LinalgAdRuleSupport::PartiallySupported,
482        mode(
483            LinalgAdRuleSupport::SupportedViaLinearize,
484            LinalgAdRoute::LinearizeThenCustomLinearTranspose,
485        ),
486        LinalgAdRuleSupport::PartiallySupported,
487        &SOLUTION_OUTPUTS,
488        &[],
489    ),
490    support_entry(
491        LinalgAdOpKind::FullPivLu,
492        LinalgAdRuleSupport::SupportedViaLinearize,
493        LinalgAdRuleSupport::Unsupported,
494        mode(
495            LinalgAdRuleSupport::SupportedViaLinearize,
496            LinalgAdRoute::LinearizeThenTranspose,
497        ),
498        LinalgAdRuleSupport::Unsupported,
499        &FULL_PIV_LU_OUTPUTS,
500        &DECOMPOSITION_CAVEATS,
501    ),
502    support_entry(
503        LinalgAdOpKind::FullPivLuSolve,
504        LinalgAdRuleSupport::SupportedViaLinearize,
505        LinalgAdRuleSupport::Supported,
506        mode(
507            LinalgAdRuleSupport::SupportedViaLinearize,
508            LinalgAdRoute::LinearizeThenCustomLinearTranspose,
509        ),
510        LinalgAdRuleSupport::Supported,
511        &FULL_PIV_LU_SOLVE_OUTPUTS,
512        &[],
513    ),
514    support_entry(
515        LinalgAdOpKind::Svd,
516        LinalgAdRuleSupport::SupportedViaLinearize,
517        LinalgAdRuleSupport::Unsupported,
518        mode(
519            LinalgAdRuleSupport::SupportedViaLinearize,
520            LinalgAdRoute::LinearizeThenTranspose,
521        ),
522        LinalgAdRuleSupport::Unsupported,
523        &SVD_OUTPUTS,
524        &DECOMPOSITION_CAVEATS,
525    ),
526    support_entry(
527        LinalgAdOpKind::SvdVals,
528        LinalgAdRuleSupport::SupportedViaLinearize,
529        LinalgAdRuleSupport::Unsupported,
530        mode(
531            LinalgAdRuleSupport::SupportedViaLinearize,
532            LinalgAdRoute::LinearizeThenTranspose,
533        ),
534        LinalgAdRuleSupport::Unsupported,
535        &SVD_VALS_OUTPUTS,
536        &DECOMPOSITION_CAVEATS,
537    ),
538    support_entry(
539        LinalgAdOpKind::Qr,
540        LinalgAdRuleSupport::SupportedViaLinearize,
541        LinalgAdRuleSupport::Unsupported,
542        mode(
543            LinalgAdRuleSupport::SupportedViaLinearize,
544            LinalgAdRoute::LinearizeThenTranspose,
545        ),
546        LinalgAdRuleSupport::Unsupported,
547        &QR_OUTPUTS,
548        &DECOMPOSITION_CAVEATS,
549    ),
550    support_entry(
551        LinalgAdOpKind::Eigh,
552        LinalgAdRuleSupport::SupportedViaLinearize,
553        LinalgAdRuleSupport::Unsupported,
554        mode(
555            LinalgAdRuleSupport::SupportedViaLinearize,
556            LinalgAdRoute::LinearizeThenTranspose,
557        ),
558        LinalgAdRuleSupport::Unsupported,
559        &EIGH_OUTPUTS,
560        &DECOMPOSITION_CAVEATS,
561    ),
562    support_entry(
563        LinalgAdOpKind::EighVals,
564        LinalgAdRuleSupport::SupportedViaLinearize,
565        LinalgAdRuleSupport::Unsupported,
566        mode(
567            LinalgAdRuleSupport::SupportedViaLinearize,
568            LinalgAdRoute::LinearizeThenTranspose,
569        ),
570        LinalgAdRuleSupport::Unsupported,
571        &EIGH_VALS_OUTPUTS,
572        &DECOMPOSITION_CAVEATS,
573    ),
574    support_entry(
575        LinalgAdOpKind::Eig,
576        LinalgAdRuleSupport::PartiallySupported,
577        LinalgAdRuleSupport::Unsupported,
578        mode(
579            LinalgAdRuleSupport::PartiallySupported,
580            LinalgAdRoute::LinearizeThenTranspose,
581        ),
582        LinalgAdRuleSupport::Unsupported,
583        &EIG_OUTPUTS,
584        &DECOMPOSITION_CAVEATS,
585    ),
586    support_entry(
587        LinalgAdOpKind::EigVals,
588        LinalgAdRuleSupport::SupportedViaLinearize,
589        LinalgAdRuleSupport::Unsupported,
590        mode(
591            LinalgAdRuleSupport::SupportedViaLinearize,
592            LinalgAdRoute::LinearizeThenTranspose,
593        ),
594        LinalgAdRuleSupport::Unsupported,
595        &EIG_VALS_OUTPUTS,
596        &DECOMPOSITION_CAVEATS,
597    ),
598    support_entry(
599        LinalgAdOpKind::TriangularSolve,
600        LinalgAdRuleSupport::SupportedViaLinearize,
601        LinalgAdRuleSupport::Supported,
602        mode(
603            LinalgAdRuleSupport::SupportedViaLinearize,
604            LinalgAdRoute::LinearizeThenCustomLinearTranspose,
605        ),
606        LinalgAdRuleSupport::Supported,
607        &SOLUTION_OUTPUTS,
608        &[],
609    ),
610    // Full-matrices SVD is a value-only route: nullspace/kernel extraction does
611    // not require derivatives, and the thin-SVD linearize rule does not extend
612    // to the square factors. AD is intentionally unsupported for every output.
613    support_entry(
614        LinalgAdOpKind::SvdFull,
615        LinalgAdRuleSupport::Unsupported,
616        LinalgAdRuleSupport::Unsupported,
617        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
618        LinalgAdRuleSupport::Unsupported,
619        &SVD_FULL_OUTPUTS,
620        &[],
621    ),
622    support_entry(
623        LinalgAdOpKind::HouseholderQrFactor,
624        LinalgAdRuleSupport::SupportedViaLinearize,
625        LinalgAdRuleSupport::Unsupported,
626        mode(
627            LinalgAdRuleSupport::SupportedViaLinearize,
628            LinalgAdRoute::LinearizeThenTranspose,
629        ),
630        LinalgAdRuleSupport::Unsupported,
631        &HOUSEHOLDER_QR_STATE_OUTPUTS,
632        &HOUSEHOLDER_QR_CAVEATS,
633    ),
634    support_entry(
635        LinalgAdOpKind::HouseholderQrFromFactors,
636        LinalgAdRuleSupport::SupportedViaLinearize,
637        LinalgAdRuleSupport::Unsupported,
638        mode(
639            LinalgAdRuleSupport::SupportedViaLinearize,
640            LinalgAdRoute::LinearizeThenTranspose,
641        ),
642        LinalgAdRuleSupport::Unsupported,
643        &HOUSEHOLDER_QR_STATE_OUTPUTS,
644        &HOUSEHOLDER_QR_CAVEATS,
645    ),
646    support_entry(
647        LinalgAdOpKind::HouseholderQrAppend,
648        LinalgAdRuleSupport::SupportedViaLinearize,
649        LinalgAdRuleSupport::Unsupported,
650        mode(
651            LinalgAdRuleSupport::SupportedViaLinearize,
652            LinalgAdRoute::LinearizeThenTranspose,
653        ),
654        LinalgAdRuleSupport::Unsupported,
655        &HOUSEHOLDER_QR_STATE_OUTPUTS,
656        &HOUSEHOLDER_QR_CAVEATS,
657    ),
658    support_entry(
659        LinalgAdOpKind::HouseholderQrR,
660        LinalgAdRuleSupport::SupportedViaLinearize,
661        LinalgAdRuleSupport::Unsupported,
662        mode(
663            LinalgAdRuleSupport::SupportedViaLinearize,
664            LinalgAdRoute::LinearizeThenTranspose,
665        ),
666        LinalgAdRuleSupport::Unsupported,
667        &HOUSEHOLDER_QR_VALUE_OUTPUTS,
668        &HOUSEHOLDER_QR_CAVEATS,
669    ),
670    support_entry(
671        LinalgAdOpKind::HouseholderQrQColumns,
672        LinalgAdRuleSupport::SupportedViaLinearize,
673        LinalgAdRuleSupport::Unsupported,
674        mode(
675            LinalgAdRuleSupport::SupportedViaLinearize,
676            LinalgAdRoute::LinearizeThenTranspose,
677        ),
678        LinalgAdRuleSupport::Unsupported,
679        &HOUSEHOLDER_QR_VALUE_OUTPUTS,
680        &HOUSEHOLDER_QR_Q_COLUMNS_CAVEATS,
681    ),
682    support_entry(
683        LinalgAdOpKind::HouseholderQrThinQ,
684        LinalgAdRuleSupport::Unsupported,
685        LinalgAdRuleSupport::Unsupported,
686        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
687        LinalgAdRuleSupport::Unsupported,
688        &HOUSEHOLDER_QR_RESIDUAL_OUTPUTS,
689        &["Internal fixed residual; no public differentiable surface."],
690    ),
691    support_entry(
692        LinalgAdOpKind::HouseholderQrAppendTangent,
693        LinalgAdRuleSupport::Unsupported,
694        LinalgAdRuleSupport::Supported,
695        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
696        LinalgAdRuleSupport::Supported,
697        &HOUSEHOLDER_QR_RESIDUAL_OUTPUTS,
698        &["Internal linear append operation; no public differentiable surface."],
699    ),
700    support_entry(
701        LinalgAdOpKind::HouseholderQrSplitTangent,
702        LinalgAdRuleSupport::Unsupported,
703        LinalgAdRuleSupport::Unsupported,
704        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
705        LinalgAdRuleSupport::Unsupported,
706        &HOUSEHOLDER_QR_RESIDUAL_OUTPUTS,
707        &["Internal transpose residual; no public differentiable surface."],
708    ),
709    support_entry(
710        LinalgAdOpKind::RankRevealingQr,
711        LinalgAdRuleSupport::Unsupported,
712        LinalgAdRuleSupport::Unsupported,
713        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
714        LinalgAdRuleSupport::Unsupported,
715        &RANK_REVEALING_QR_OUTPUTS,
716        &["Pivot selection and numerical rank are discontinuous; all outputs are initially unsupported for AD."],
717    ),
718];
719
720/// Return the complete linalg AD support manifest.
721///
722/// # Examples
723///
724/// ```rust
725/// let manifest = tenferro_linalg::all_linalg_ad_support();
726/// assert_eq!(manifest.len(), tenferro_linalg::LinalgAdOpKind::COUNT);
727/// ```
728pub fn all_linalg_ad_support() -> &'static [LinalgAdSupport; LinalgAdOpKind::COUNT] {
729    &LINALG_AD_SUPPORT
730}
731
732/// Return the support manifest entry for one linalg operation kind.
733///
734/// # Examples
735///
736/// ```rust
737/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind};
738///
739/// let entry = linalg_ad_support(LinalgAdOpKind::Eigh);
740/// assert_eq!(entry.kind, LinalgAdOpKind::Eigh);
741/// ```
742pub fn linalg_ad_support(kind: LinalgAdOpKind) -> &'static LinalgAdSupport {
743    &LINALG_AD_SUPPORT[kind.as_index()]
744}
745
746#[cfg(test)]
747pub(crate) fn linalg_ad_support_for_op(op: LinalgOp) -> &'static LinalgAdSupport {
748    linalg_ad_support(LinalgAdOpKind::from_linalg_op(op))
749}
750
751#[cfg(test)]
752mod tests;