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}
100
101impl LinalgAdOpKind {
102    pub const COUNT: usize = 17;
103
104    /// Return the manifest index for this operation kind.
105    ///
106    /// # Examples
107    ///
108    /// ```rust
109    /// use tenferro_linalg::LinalgAdOpKind;
110    ///
111    /// assert_eq!(LinalgAdOpKind::Cholesky.as_index(), 0);
112    /// ```
113    pub const fn as_index(self) -> usize {
114        match self {
115            Self::Cholesky => 0,
116            Self::Lu => 1,
117            Self::LuFactor => 2,
118            Self::SignDetFromLuFactor => 3,
119            Self::LogAbsDetFromLuFactor => 4,
120            Self::LuSolvePrepared => 5,
121            Self::FullPivLu => 6,
122            Self::FullPivLuSolve => 7,
123            Self::Svd => 8,
124            Self::SvdVals => 9,
125            Self::Qr => 10,
126            Self::Eigh => 11,
127            Self::EighVals => 12,
128            Self::Eig => 13,
129            Self::EigVals => 14,
130            Self::TriangularSolve => 15,
131            Self::SvdFull => 16,
132        }
133    }
134
135    #[cfg(test)]
136    pub(crate) const fn from_linalg_op(op: LinalgOp) -> Self {
137        match op {
138            LinalgOp::Cholesky => Self::Cholesky,
139            LinalgOp::Lu => Self::Lu,
140            LinalgOp::LuFactor => Self::LuFactor,
141            LinalgOp::SignDetFromLuFactor => Self::SignDetFromLuFactor,
142            LinalgOp::LogAbsDetFromLuFactor => Self::LogAbsDetFromLuFactor,
143            LinalgOp::LuSolvePrepared { .. } => Self::LuSolvePrepared,
144            LinalgOp::FullPivLu => Self::FullPivLu,
145            LinalgOp::FullPivLuSolve { .. } => Self::FullPivLuSolve,
146            // The partial-pivot single-op solve shares the solve-family AD
147            // route (linearize + custom linear transpose) with FullPivLuSolve
148            // and is not separately exposed in the public manifest.
149            LinalgOp::Solve => Self::FullPivLuSolve,
150            LinalgOp::Svd { .. } => Self::Svd,
151            LinalgOp::SvdFull => Self::SvdFull,
152            LinalgOp::SvdVals { .. } => Self::SvdVals,
153            LinalgOp::Qr { .. } => Self::Qr,
154            LinalgOp::Eigh { .. } => Self::Eigh,
155            LinalgOp::EighVals { .. } => Self::EighVals,
156            LinalgOp::Eig { .. } => Self::Eig,
157            LinalgOp::EigVals { .. } => Self::EigVals,
158            LinalgOp::TriangularSolve { .. } => Self::TriangularSolve,
159        }
160    }
161}
162
163/// AD support status for one output of a linalg operation.
164///
165/// # Examples
166///
167/// ```rust
168/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind, LinalgAdRuleSupport};
169///
170/// let full_piv_lu = linalg_ad_support(LinalgAdOpKind::FullPivLu);
171/// let l_output = full_piv_lu.outputs.iter().find(|output| output.name == "l").unwrap();
172/// assert_eq!(l_output.status, LinalgAdRuleSupport::SupportedViaLinearize);
173/// ```
174#[derive(Clone, Copy, Debug, PartialEq, Eq)]
175pub struct LinalgAdOutputSupport {
176    /// Output position in the linalg operation result tuple.
177    pub index: usize,
178    /// Stable output name used by tests and support dashboards.
179    pub name: &'static str,
180    /// AD support status for this specific output.
181    pub status: LinalgAdRuleSupport,
182}
183
184/// AD support manifest entry for one linalg operation.
185///
186/// # Examples
187///
188/// ```rust
189/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind, LinalgAdRuleSupport};
190///
191/// let solve = linalg_ad_support(LinalgAdOpKind::TriangularSolve);
192/// assert_eq!(solve.vjp.route, tenferro_linalg::LinalgAdRoute::LinearizeThenCustomLinearTranspose);
193/// ```
194#[derive(Clone, Copy, Debug, PartialEq, Eq)]
195pub struct LinalgAdSupport {
196    /// Operation kind described by this manifest entry.
197    pub kind: LinalgAdOpKind,
198    /// User-visible JVP support and route.
199    pub jvp: LinalgAdModeSupport,
200    /// User-visible VJP support and route.
201    pub vjp: LinalgAdModeSupport,
202    /// Definitional linearize rule implementation status.
203    pub linearize_rule: LinalgAdRuleSupport,
204    /// Direct primal VJP rule implementation status.
205    pub custom_vjp_rule: LinalgAdRuleSupport,
206    /// Custom transposed-linear rule implementation status.
207    pub custom_linear_transpose_rule: LinalgAdRuleSupport,
208    /// Forward-mode graph emission support.
209    pub linearize: LinalgAdRuleSupport,
210    /// Transposed-linear graph emission support.
211    pub transpose: LinalgAdRuleSupport,
212    /// Per-output support status for multi-output operations.
213    pub outputs: &'static [LinalgAdOutputSupport],
214    /// Numerical or semantic caveats for this operation family.
215    pub caveats: &'static [&'static str],
216}
217
218const fn mode(status: LinalgAdRuleSupport, route: LinalgAdRoute) -> LinalgAdModeSupport {
219    LinalgAdModeSupport { status, route }
220}
221
222const fn jvp_route(status: LinalgAdRuleSupport) -> LinalgAdRoute {
223    match status {
224        LinalgAdRuleSupport::Unsupported
225        | LinalgAdRuleSupport::NonDifferentiable
226        | LinalgAdRuleSupport::PendingOracle => LinalgAdRoute::Unsupported,
227        LinalgAdRuleSupport::Supported
228        | LinalgAdRuleSupport::SupportedViaLinearize
229        | LinalgAdRuleSupport::PartiallySupported => LinalgAdRoute::Linearize,
230    }
231}
232
233const fn support_entry(
234    kind: LinalgAdOpKind,
235    linearize: LinalgAdRuleSupport,
236    transpose: LinalgAdRuleSupport,
237    vjp: LinalgAdModeSupport,
238    custom_linear_transpose_rule: LinalgAdRuleSupport,
239    outputs: &'static [LinalgAdOutputSupport],
240    caveats: &'static [&'static str],
241) -> LinalgAdSupport {
242    LinalgAdSupport {
243        kind,
244        jvp: mode(linearize, jvp_route(linearize)),
245        vjp,
246        linearize_rule: linearize,
247        custom_vjp_rule: LinalgAdRuleSupport::Unsupported,
248        custom_linear_transpose_rule,
249        linearize,
250        transpose,
251        outputs,
252        caveats,
253    }
254}
255
256const fn output(
257    index: usize,
258    name: &'static str,
259    status: LinalgAdRuleSupport,
260) -> LinalgAdOutputSupport {
261    LinalgAdOutputSupport {
262        index,
263        name,
264        status,
265    }
266}
267
268static CHOLESKY_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
269    0,
270    "factor",
271    LinalgAdRuleSupport::SupportedViaLinearize,
272)];
273static LU_OUTPUTS: [LinalgAdOutputSupport; 4] = [
274    output(0, "p", LinalgAdRuleSupport::NonDifferentiable),
275    output(1, "l", LinalgAdRuleSupport::SupportedViaLinearize),
276    output(2, "u", LinalgAdRuleSupport::SupportedViaLinearize),
277    output(3, "parity", LinalgAdRuleSupport::NonDifferentiable),
278];
279static LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 3] = [
280    output(0, "packed_lu", LinalgAdRuleSupport::Unsupported),
281    output(1, "pivots", LinalgAdRuleSupport::NonDifferentiable),
282    output(2, "parity", LinalgAdRuleSupport::NonDifferentiable),
283];
284static SIGNDET_FROM_LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
285    0,
286    "sign",
287    LinalgAdRuleSupport::SupportedViaLinearize,
288)];
289static LOGABSDET_FROM_LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
290    0,
291    "logabsdet",
292    LinalgAdRuleSupport::SupportedViaLinearize,
293)];
294static SOLUTION_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
295    0,
296    "solution",
297    LinalgAdRuleSupport::SupportedViaLinearize,
298)];
299static FULL_PIV_LU_OUTPUTS: [LinalgAdOutputSupport; 5] = [
300    output(0, "p", LinalgAdRuleSupport::NonDifferentiable),
301    output(1, "l", LinalgAdRuleSupport::SupportedViaLinearize),
302    output(2, "u", LinalgAdRuleSupport::SupportedViaLinearize),
303    output(3, "q", LinalgAdRuleSupport::NonDifferentiable),
304    output(4, "parity", LinalgAdRuleSupport::NonDifferentiable),
305];
306static FULL_PIV_LU_SOLVE_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
307    0,
308    "solution",
309    LinalgAdRuleSupport::SupportedViaLinearize,
310)];
311static SVD_OUTPUTS: [LinalgAdOutputSupport; 3] = [
312    output(0, "u", LinalgAdRuleSupport::SupportedViaLinearize),
313    output(
314        1,
315        "singular_values",
316        LinalgAdRuleSupport::SupportedViaLinearize,
317    ),
318    output(2, "vt", LinalgAdRuleSupport::SupportedViaLinearize),
319];
320static SVD_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
321    0,
322    "singular_values",
323    LinalgAdRuleSupport::SupportedViaLinearize,
324)];
325static SVD_FULL_OUTPUTS: [LinalgAdOutputSupport; 3] = [
326    output(0, "u", LinalgAdRuleSupport::Unsupported),
327    output(1, "singular_values", LinalgAdRuleSupport::Unsupported),
328    output(2, "vt", LinalgAdRuleSupport::Unsupported),
329];
330static QR_OUTPUTS: [LinalgAdOutputSupport; 2] = [
331    output(0, "q", LinalgAdRuleSupport::SupportedViaLinearize),
332    output(1, "r", LinalgAdRuleSupport::SupportedViaLinearize),
333];
334static EIGH_OUTPUTS: [LinalgAdOutputSupport; 2] = [
335    output(0, "eigenvalues", LinalgAdRuleSupport::SupportedViaLinearize),
336    output(
337        1,
338        "eigenvectors",
339        LinalgAdRuleSupport::SupportedViaLinearize,
340    ),
341];
342static EIGH_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
343    0,
344    "eigenvalues",
345    LinalgAdRuleSupport::SupportedViaLinearize,
346)];
347static EIG_OUTPUTS: [LinalgAdOutputSupport; 2] = [
348    output(0, "eigenvalues", LinalgAdRuleSupport::SupportedViaLinearize),
349    output(1, "eigenvectors", LinalgAdRuleSupport::Unsupported),
350];
351static EIG_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
352    0,
353    "eigenvalues",
354    LinalgAdRuleSupport::SupportedViaLinearize,
355)];
356
357static DECOMPOSITION_CAVEATS: [&str; 1] = [
358    "Derivative regularization handles near-degenerate spectra but does not make exact degeneracies smoothly differentiable.",
359];
360
361static LINALG_AD_SUPPORT: [LinalgAdSupport; LinalgAdOpKind::COUNT] = [
362    support_entry(
363        LinalgAdOpKind::Cholesky,
364        LinalgAdRuleSupport::SupportedViaLinearize,
365        LinalgAdRuleSupport::Unsupported,
366        mode(
367            LinalgAdRuleSupport::SupportedViaLinearize,
368            LinalgAdRoute::LinearizeThenTranspose,
369        ),
370        LinalgAdRuleSupport::Unsupported,
371        &CHOLESKY_OUTPUTS,
372        &DECOMPOSITION_CAVEATS,
373    ),
374    support_entry(
375        LinalgAdOpKind::Lu,
376        LinalgAdRuleSupport::PartiallySupported,
377        LinalgAdRuleSupport::Unsupported,
378        mode(
379            LinalgAdRuleSupport::PartiallySupported,
380            LinalgAdRoute::LinearizeThenTranspose,
381        ),
382        LinalgAdRuleSupport::Unsupported,
383        &LU_OUTPUTS,
384        &DECOMPOSITION_CAVEATS,
385    ),
386    support_entry(
387        LinalgAdOpKind::LuFactor,
388        LinalgAdRuleSupport::Unsupported,
389        LinalgAdRuleSupport::Unsupported,
390        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
391        LinalgAdRuleSupport::Unsupported,
392        &LU_FACTOR_OUTPUTS,
393        &[],
394    ),
395    support_entry(
396        LinalgAdOpKind::SignDetFromLuFactor,
397        LinalgAdRuleSupport::SupportedViaLinearize,
398        LinalgAdRuleSupport::Unsupported,
399        mode(
400            LinalgAdRuleSupport::SupportedViaLinearize,
401            LinalgAdRoute::LinearizeThenTranspose,
402        ),
403        LinalgAdRuleSupport::Unsupported,
404        &SIGNDET_FROM_LU_FACTOR_OUTPUTS,
405        &[],
406    ),
407    support_entry(
408        LinalgAdOpKind::LogAbsDetFromLuFactor,
409        LinalgAdRuleSupport::SupportedViaLinearize,
410        LinalgAdRuleSupport::Unsupported,
411        mode(
412            LinalgAdRuleSupport::SupportedViaLinearize,
413            LinalgAdRoute::LinearizeThenTranspose,
414        ),
415        LinalgAdRuleSupport::Unsupported,
416        &LOGABSDET_FROM_LU_FACTOR_OUTPUTS,
417        &[],
418    ),
419    support_entry(
420        LinalgAdOpKind::LuSolvePrepared,
421        LinalgAdRuleSupport::SupportedViaLinearize,
422        LinalgAdRuleSupport::PartiallySupported,
423        mode(
424            LinalgAdRuleSupport::SupportedViaLinearize,
425            LinalgAdRoute::LinearizeThenCustomLinearTranspose,
426        ),
427        LinalgAdRuleSupport::PartiallySupported,
428        &SOLUTION_OUTPUTS,
429        &[],
430    ),
431    support_entry(
432        LinalgAdOpKind::FullPivLu,
433        LinalgAdRuleSupport::SupportedViaLinearize,
434        LinalgAdRuleSupport::Unsupported,
435        mode(
436            LinalgAdRuleSupport::SupportedViaLinearize,
437            LinalgAdRoute::LinearizeThenTranspose,
438        ),
439        LinalgAdRuleSupport::Unsupported,
440        &FULL_PIV_LU_OUTPUTS,
441        &DECOMPOSITION_CAVEATS,
442    ),
443    support_entry(
444        LinalgAdOpKind::FullPivLuSolve,
445        LinalgAdRuleSupport::SupportedViaLinearize,
446        LinalgAdRuleSupport::Supported,
447        mode(
448            LinalgAdRuleSupport::SupportedViaLinearize,
449            LinalgAdRoute::LinearizeThenCustomLinearTranspose,
450        ),
451        LinalgAdRuleSupport::Supported,
452        &FULL_PIV_LU_SOLVE_OUTPUTS,
453        &[],
454    ),
455    support_entry(
456        LinalgAdOpKind::Svd,
457        LinalgAdRuleSupport::SupportedViaLinearize,
458        LinalgAdRuleSupport::Unsupported,
459        mode(
460            LinalgAdRuleSupport::SupportedViaLinearize,
461            LinalgAdRoute::LinearizeThenTranspose,
462        ),
463        LinalgAdRuleSupport::Unsupported,
464        &SVD_OUTPUTS,
465        &DECOMPOSITION_CAVEATS,
466    ),
467    support_entry(
468        LinalgAdOpKind::SvdVals,
469        LinalgAdRuleSupport::SupportedViaLinearize,
470        LinalgAdRuleSupport::Unsupported,
471        mode(
472            LinalgAdRuleSupport::SupportedViaLinearize,
473            LinalgAdRoute::LinearizeThenTranspose,
474        ),
475        LinalgAdRuleSupport::Unsupported,
476        &SVD_VALS_OUTPUTS,
477        &DECOMPOSITION_CAVEATS,
478    ),
479    support_entry(
480        LinalgAdOpKind::Qr,
481        LinalgAdRuleSupport::SupportedViaLinearize,
482        LinalgAdRuleSupport::Unsupported,
483        mode(
484            LinalgAdRuleSupport::SupportedViaLinearize,
485            LinalgAdRoute::LinearizeThenTranspose,
486        ),
487        LinalgAdRuleSupport::Unsupported,
488        &QR_OUTPUTS,
489        &DECOMPOSITION_CAVEATS,
490    ),
491    support_entry(
492        LinalgAdOpKind::Eigh,
493        LinalgAdRuleSupport::SupportedViaLinearize,
494        LinalgAdRuleSupport::Unsupported,
495        mode(
496            LinalgAdRuleSupport::SupportedViaLinearize,
497            LinalgAdRoute::LinearizeThenTranspose,
498        ),
499        LinalgAdRuleSupport::Unsupported,
500        &EIGH_OUTPUTS,
501        &DECOMPOSITION_CAVEATS,
502    ),
503    support_entry(
504        LinalgAdOpKind::EighVals,
505        LinalgAdRuleSupport::SupportedViaLinearize,
506        LinalgAdRuleSupport::Unsupported,
507        mode(
508            LinalgAdRuleSupport::SupportedViaLinearize,
509            LinalgAdRoute::LinearizeThenTranspose,
510        ),
511        LinalgAdRuleSupport::Unsupported,
512        &EIGH_VALS_OUTPUTS,
513        &DECOMPOSITION_CAVEATS,
514    ),
515    support_entry(
516        LinalgAdOpKind::Eig,
517        LinalgAdRuleSupport::PartiallySupported,
518        LinalgAdRuleSupport::Unsupported,
519        mode(
520            LinalgAdRuleSupport::PartiallySupported,
521            LinalgAdRoute::LinearizeThenTranspose,
522        ),
523        LinalgAdRuleSupport::Unsupported,
524        &EIG_OUTPUTS,
525        &DECOMPOSITION_CAVEATS,
526    ),
527    support_entry(
528        LinalgAdOpKind::EigVals,
529        LinalgAdRuleSupport::SupportedViaLinearize,
530        LinalgAdRuleSupport::Unsupported,
531        mode(
532            LinalgAdRuleSupport::SupportedViaLinearize,
533            LinalgAdRoute::LinearizeThenTranspose,
534        ),
535        LinalgAdRuleSupport::Unsupported,
536        &EIG_VALS_OUTPUTS,
537        &DECOMPOSITION_CAVEATS,
538    ),
539    support_entry(
540        LinalgAdOpKind::TriangularSolve,
541        LinalgAdRuleSupport::SupportedViaLinearize,
542        LinalgAdRuleSupport::Supported,
543        mode(
544            LinalgAdRuleSupport::SupportedViaLinearize,
545            LinalgAdRoute::LinearizeThenCustomLinearTranspose,
546        ),
547        LinalgAdRuleSupport::Supported,
548        &SOLUTION_OUTPUTS,
549        &[],
550    ),
551    // Full-matrices SVD is a value-only route: nullspace/kernel extraction does
552    // not require derivatives, and the thin-SVD linearize rule does not extend
553    // to the square factors. AD is intentionally unsupported for every output.
554    support_entry(
555        LinalgAdOpKind::SvdFull,
556        LinalgAdRuleSupport::Unsupported,
557        LinalgAdRuleSupport::Unsupported,
558        mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
559        LinalgAdRuleSupport::Unsupported,
560        &SVD_FULL_OUTPUTS,
561        &[],
562    ),
563];
564
565/// Return the complete linalg AD support manifest.
566///
567/// # Examples
568///
569/// ```rust
570/// let manifest = tenferro_linalg::all_linalg_ad_support();
571/// assert_eq!(manifest.len(), tenferro_linalg::LinalgAdOpKind::COUNT);
572/// ```
573pub fn all_linalg_ad_support() -> &'static [LinalgAdSupport; LinalgAdOpKind::COUNT] {
574    &LINALG_AD_SUPPORT
575}
576
577/// Return the support manifest entry for one linalg operation kind.
578///
579/// # Examples
580///
581/// ```rust
582/// use tenferro_linalg::{linalg_ad_support, LinalgAdOpKind};
583///
584/// let entry = linalg_ad_support(LinalgAdOpKind::Eigh);
585/// assert_eq!(entry.kind, LinalgAdOpKind::Eigh);
586/// ```
587pub fn linalg_ad_support(kind: LinalgAdOpKind) -> &'static LinalgAdSupport {
588    &LINALG_AD_SUPPORT[kind.as_index()]
589}
590
591#[cfg(test)]
592pub(crate) fn linalg_ad_support_for_op(op: LinalgOp) -> &'static LinalgAdSupport {
593    linalg_ad_support(LinalgAdOpKind::from_linalg_op(op))
594}
595
596#[cfg(test)]
597mod tests;