1#[cfg(test)]
2use crate::extension::LinalgOp;
3
4#[derive(Clone, Copy, Debug, PartialEq, Eq)]
15pub enum LinalgAdRuleSupport {
16 Supported,
17 SupportedViaLinearize,
18 PartiallySupported,
19 NonDifferentiable,
20 Unsupported,
21 PendingOracle,
22}
23
24#[derive(Clone, Copy, Debug, PartialEq, Eq)]
35pub enum LinalgAdRoute {
36 Unsupported,
38 Linearize,
40 LinearizeThenTranspose,
43 LinearizeThenCustomLinearTranspose,
46 CustomVjp,
48 CustomPreferredWithLinearizeFallback,
51}
52
53#[derive(Clone, Copy, Debug, PartialEq, Eq)]
64pub struct LinalgAdModeSupport {
65 pub status: LinalgAdRuleSupport,
67 pub route: LinalgAdRoute,
69}
70
71#[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 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 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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
202pub struct LinalgAdOutputSupport {
203 pub index: usize,
205 pub name: &'static str,
207 pub status: LinalgAdRuleSupport,
209}
210
211#[derive(Clone, Copy, Debug, PartialEq, Eq)]
222pub struct LinalgAdSupport {
223 pub kind: LinalgAdOpKind,
225 pub jvp: LinalgAdModeSupport,
227 pub vjp: LinalgAdModeSupport,
229 pub linearize_rule: LinalgAdRuleSupport,
231 pub custom_vjp_rule: LinalgAdRuleSupport,
233 pub custom_linear_transpose_rule: LinalgAdRuleSupport,
235 pub linearize: LinalgAdRuleSupport,
237 pub transpose: LinalgAdRuleSupport,
239 pub outputs: &'static [LinalgAdOutputSupport],
241 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 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
710pub fn all_linalg_ad_support() -> &'static [LinalgAdSupport; LinalgAdOpKind::COUNT] {
719 &LINALG_AD_SUPPORT
720}
721
722pub 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;