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::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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
205pub struct LinalgAdOutputSupport {
206 pub index: usize,
208 pub name: &'static str,
210 pub status: LinalgAdRuleSupport,
212}
213
214#[derive(Clone, Copy, Debug, PartialEq, Eq)]
225pub struct LinalgAdSupport {
226 pub kind: LinalgAdOpKind,
228 pub jvp: LinalgAdModeSupport,
230 pub vjp: LinalgAdModeSupport,
232 pub linearize_rule: LinalgAdRuleSupport,
234 pub custom_vjp_rule: LinalgAdRuleSupport,
236 pub custom_linear_transpose_rule: LinalgAdRuleSupport,
238 pub linearize: LinalgAdRuleSupport,
240 pub transpose: LinalgAdRuleSupport,
242 pub outputs: &'static [LinalgAdOutputSupport],
244 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 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
720pub fn all_linalg_ad_support() -> &'static [LinalgAdSupport; LinalgAdOpKind::COUNT] {
729 &LINALG_AD_SUPPORT
730}
731
732pub 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;