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}
100
101impl LinalgAdOpKind {
102 pub const COUNT: usize = 17;
103
104 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 LinalgOp::Svd { .. } => Self::Svd,
147 LinalgOp::SvdFull => Self::SvdFull,
148 LinalgOp::SvdVals { .. } => Self::SvdVals,
149 LinalgOp::Qr { .. } => Self::Qr,
150 LinalgOp::Eigh { .. } => Self::Eigh,
151 LinalgOp::EighVals { .. } => Self::EighVals,
152 LinalgOp::Eig { .. } => Self::Eig,
153 LinalgOp::EigVals { .. } => Self::EigVals,
154 LinalgOp::TriangularSolve { .. } => Self::TriangularSolve,
155 }
156 }
157}
158
159#[derive(Clone, Copy, Debug, PartialEq, Eq)]
171pub struct LinalgAdOutputSupport {
172 pub index: usize,
174 pub name: &'static str,
176 pub status: LinalgAdRuleSupport,
178}
179
180#[derive(Clone, Copy, Debug, PartialEq, Eq)]
191pub struct LinalgAdSupport {
192 pub kind: LinalgAdOpKind,
194 pub jvp: LinalgAdModeSupport,
196 pub vjp: LinalgAdModeSupport,
198 pub linearize_rule: LinalgAdRuleSupport,
200 pub custom_vjp_rule: LinalgAdRuleSupport,
202 pub custom_linear_transpose_rule: LinalgAdRuleSupport,
204 pub linearize: LinalgAdRuleSupport,
206 pub transpose: LinalgAdRuleSupport,
208 pub outputs: &'static [LinalgAdOutputSupport],
210 pub caveats: &'static [&'static str],
212}
213
214const fn mode(status: LinalgAdRuleSupport, route: LinalgAdRoute) -> LinalgAdModeSupport {
215 LinalgAdModeSupport { status, route }
216}
217
218const fn jvp_route(status: LinalgAdRuleSupport) -> LinalgAdRoute {
219 match status {
220 LinalgAdRuleSupport::Unsupported
221 | LinalgAdRuleSupport::NonDifferentiable
222 | LinalgAdRuleSupport::PendingOracle => LinalgAdRoute::Unsupported,
223 LinalgAdRuleSupport::Supported
224 | LinalgAdRuleSupport::SupportedViaLinearize
225 | LinalgAdRuleSupport::PartiallySupported => LinalgAdRoute::Linearize,
226 }
227}
228
229const fn support_entry(
230 kind: LinalgAdOpKind,
231 linearize: LinalgAdRuleSupport,
232 transpose: LinalgAdRuleSupport,
233 vjp: LinalgAdModeSupport,
234 custom_linear_transpose_rule: LinalgAdRuleSupport,
235 outputs: &'static [LinalgAdOutputSupport],
236 caveats: &'static [&'static str],
237) -> LinalgAdSupport {
238 LinalgAdSupport {
239 kind,
240 jvp: mode(linearize, jvp_route(linearize)),
241 vjp,
242 linearize_rule: linearize,
243 custom_vjp_rule: LinalgAdRuleSupport::Unsupported,
244 custom_linear_transpose_rule,
245 linearize,
246 transpose,
247 outputs,
248 caveats,
249 }
250}
251
252const fn output(
253 index: usize,
254 name: &'static str,
255 status: LinalgAdRuleSupport,
256) -> LinalgAdOutputSupport {
257 LinalgAdOutputSupport {
258 index,
259 name,
260 status,
261 }
262}
263
264static CHOLESKY_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
265 0,
266 "factor",
267 LinalgAdRuleSupport::SupportedViaLinearize,
268)];
269static LU_OUTPUTS: [LinalgAdOutputSupport; 4] = [
270 output(0, "p", LinalgAdRuleSupport::NonDifferentiable),
271 output(1, "l", LinalgAdRuleSupport::SupportedViaLinearize),
272 output(2, "u", LinalgAdRuleSupport::SupportedViaLinearize),
273 output(3, "parity", LinalgAdRuleSupport::NonDifferentiable),
274];
275static LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 3] = [
276 output(0, "packed_lu", LinalgAdRuleSupport::Unsupported),
277 output(1, "pivots", LinalgAdRuleSupport::NonDifferentiable),
278 output(2, "parity", LinalgAdRuleSupport::NonDifferentiable),
279];
280static SIGNDET_FROM_LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
281 0,
282 "sign",
283 LinalgAdRuleSupport::SupportedViaLinearize,
284)];
285static LOGABSDET_FROM_LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
286 0,
287 "logabsdet",
288 LinalgAdRuleSupport::SupportedViaLinearize,
289)];
290static SOLUTION_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
291 0,
292 "solution",
293 LinalgAdRuleSupport::SupportedViaLinearize,
294)];
295static FULL_PIV_LU_OUTPUTS: [LinalgAdOutputSupport; 5] = [
296 output(0, "p", LinalgAdRuleSupport::NonDifferentiable),
297 output(1, "l", LinalgAdRuleSupport::SupportedViaLinearize),
298 output(2, "u", LinalgAdRuleSupport::SupportedViaLinearize),
299 output(3, "q", LinalgAdRuleSupport::NonDifferentiable),
300 output(4, "parity", LinalgAdRuleSupport::NonDifferentiable),
301];
302static FULL_PIV_LU_SOLVE_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
303 0,
304 "solution",
305 LinalgAdRuleSupport::SupportedViaLinearize,
306)];
307static SVD_OUTPUTS: [LinalgAdOutputSupport; 3] = [
308 output(0, "u", LinalgAdRuleSupport::SupportedViaLinearize),
309 output(
310 1,
311 "singular_values",
312 LinalgAdRuleSupport::SupportedViaLinearize,
313 ),
314 output(2, "vt", LinalgAdRuleSupport::SupportedViaLinearize),
315];
316static SVD_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
317 0,
318 "singular_values",
319 LinalgAdRuleSupport::SupportedViaLinearize,
320)];
321static SVD_FULL_OUTPUTS: [LinalgAdOutputSupport; 3] = [
322 output(0, "u", LinalgAdRuleSupport::Unsupported),
323 output(1, "singular_values", LinalgAdRuleSupport::Unsupported),
324 output(2, "vt", LinalgAdRuleSupport::Unsupported),
325];
326static QR_OUTPUTS: [LinalgAdOutputSupport; 2] = [
327 output(0, "q", LinalgAdRuleSupport::SupportedViaLinearize),
328 output(1, "r", LinalgAdRuleSupport::SupportedViaLinearize),
329];
330static EIGH_OUTPUTS: [LinalgAdOutputSupport; 2] = [
331 output(0, "eigenvalues", LinalgAdRuleSupport::SupportedViaLinearize),
332 output(
333 1,
334 "eigenvectors",
335 LinalgAdRuleSupport::SupportedViaLinearize,
336 ),
337];
338static EIGH_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
339 0,
340 "eigenvalues",
341 LinalgAdRuleSupport::SupportedViaLinearize,
342)];
343static EIG_OUTPUTS: [LinalgAdOutputSupport; 2] = [
344 output(0, "eigenvalues", LinalgAdRuleSupport::SupportedViaLinearize),
345 output(1, "eigenvectors", LinalgAdRuleSupport::Unsupported),
346];
347static EIG_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
348 0,
349 "eigenvalues",
350 LinalgAdRuleSupport::SupportedViaLinearize,
351)];
352
353static DECOMPOSITION_CAVEATS: [&str; 1] = [
354 "Derivative regularization handles near-degenerate spectra but does not make exact degeneracies smoothly differentiable.",
355];
356
357static LINALG_AD_SUPPORT: [LinalgAdSupport; LinalgAdOpKind::COUNT] = [
358 support_entry(
359 LinalgAdOpKind::Cholesky,
360 LinalgAdRuleSupport::SupportedViaLinearize,
361 LinalgAdRuleSupport::Unsupported,
362 mode(
363 LinalgAdRuleSupport::SupportedViaLinearize,
364 LinalgAdRoute::LinearizeThenTranspose,
365 ),
366 LinalgAdRuleSupport::Unsupported,
367 &CHOLESKY_OUTPUTS,
368 &DECOMPOSITION_CAVEATS,
369 ),
370 support_entry(
371 LinalgAdOpKind::Lu,
372 LinalgAdRuleSupport::PartiallySupported,
373 LinalgAdRuleSupport::Unsupported,
374 mode(
375 LinalgAdRuleSupport::PartiallySupported,
376 LinalgAdRoute::LinearizeThenTranspose,
377 ),
378 LinalgAdRuleSupport::Unsupported,
379 &LU_OUTPUTS,
380 &DECOMPOSITION_CAVEATS,
381 ),
382 support_entry(
383 LinalgAdOpKind::LuFactor,
384 LinalgAdRuleSupport::Unsupported,
385 LinalgAdRuleSupport::Unsupported,
386 mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
387 LinalgAdRuleSupport::Unsupported,
388 &LU_FACTOR_OUTPUTS,
389 &[],
390 ),
391 support_entry(
392 LinalgAdOpKind::SignDetFromLuFactor,
393 LinalgAdRuleSupport::SupportedViaLinearize,
394 LinalgAdRuleSupport::Unsupported,
395 mode(
396 LinalgAdRuleSupport::SupportedViaLinearize,
397 LinalgAdRoute::LinearizeThenTranspose,
398 ),
399 LinalgAdRuleSupport::Unsupported,
400 &SIGNDET_FROM_LU_FACTOR_OUTPUTS,
401 &[],
402 ),
403 support_entry(
404 LinalgAdOpKind::LogAbsDetFromLuFactor,
405 LinalgAdRuleSupport::SupportedViaLinearize,
406 LinalgAdRuleSupport::Unsupported,
407 mode(
408 LinalgAdRuleSupport::SupportedViaLinearize,
409 LinalgAdRoute::LinearizeThenTranspose,
410 ),
411 LinalgAdRuleSupport::Unsupported,
412 &LOGABSDET_FROM_LU_FACTOR_OUTPUTS,
413 &[],
414 ),
415 support_entry(
416 LinalgAdOpKind::LuSolvePrepared,
417 LinalgAdRuleSupport::SupportedViaLinearize,
418 LinalgAdRuleSupport::PartiallySupported,
419 mode(
420 LinalgAdRuleSupport::SupportedViaLinearize,
421 LinalgAdRoute::LinearizeThenCustomLinearTranspose,
422 ),
423 LinalgAdRuleSupport::PartiallySupported,
424 &SOLUTION_OUTPUTS,
425 &[],
426 ),
427 support_entry(
428 LinalgAdOpKind::FullPivLu,
429 LinalgAdRuleSupport::SupportedViaLinearize,
430 LinalgAdRuleSupport::Unsupported,
431 mode(
432 LinalgAdRuleSupport::SupportedViaLinearize,
433 LinalgAdRoute::LinearizeThenTranspose,
434 ),
435 LinalgAdRuleSupport::Unsupported,
436 &FULL_PIV_LU_OUTPUTS,
437 &DECOMPOSITION_CAVEATS,
438 ),
439 support_entry(
440 LinalgAdOpKind::FullPivLuSolve,
441 LinalgAdRuleSupport::SupportedViaLinearize,
442 LinalgAdRuleSupport::Supported,
443 mode(
444 LinalgAdRuleSupport::SupportedViaLinearize,
445 LinalgAdRoute::LinearizeThenCustomLinearTranspose,
446 ),
447 LinalgAdRuleSupport::Supported,
448 &FULL_PIV_LU_SOLVE_OUTPUTS,
449 &[],
450 ),
451 support_entry(
452 LinalgAdOpKind::Svd,
453 LinalgAdRuleSupport::SupportedViaLinearize,
454 LinalgAdRuleSupport::Unsupported,
455 mode(
456 LinalgAdRuleSupport::SupportedViaLinearize,
457 LinalgAdRoute::LinearizeThenTranspose,
458 ),
459 LinalgAdRuleSupport::Unsupported,
460 &SVD_OUTPUTS,
461 &DECOMPOSITION_CAVEATS,
462 ),
463 support_entry(
464 LinalgAdOpKind::SvdVals,
465 LinalgAdRuleSupport::SupportedViaLinearize,
466 LinalgAdRuleSupport::Unsupported,
467 mode(
468 LinalgAdRuleSupport::SupportedViaLinearize,
469 LinalgAdRoute::LinearizeThenTranspose,
470 ),
471 LinalgAdRuleSupport::Unsupported,
472 &SVD_VALS_OUTPUTS,
473 &DECOMPOSITION_CAVEATS,
474 ),
475 support_entry(
476 LinalgAdOpKind::Qr,
477 LinalgAdRuleSupport::SupportedViaLinearize,
478 LinalgAdRuleSupport::Unsupported,
479 mode(
480 LinalgAdRuleSupport::SupportedViaLinearize,
481 LinalgAdRoute::LinearizeThenTranspose,
482 ),
483 LinalgAdRuleSupport::Unsupported,
484 &QR_OUTPUTS,
485 &DECOMPOSITION_CAVEATS,
486 ),
487 support_entry(
488 LinalgAdOpKind::Eigh,
489 LinalgAdRuleSupport::SupportedViaLinearize,
490 LinalgAdRuleSupport::Unsupported,
491 mode(
492 LinalgAdRuleSupport::SupportedViaLinearize,
493 LinalgAdRoute::LinearizeThenTranspose,
494 ),
495 LinalgAdRuleSupport::Unsupported,
496 &EIGH_OUTPUTS,
497 &DECOMPOSITION_CAVEATS,
498 ),
499 support_entry(
500 LinalgAdOpKind::EighVals,
501 LinalgAdRuleSupport::SupportedViaLinearize,
502 LinalgAdRuleSupport::Unsupported,
503 mode(
504 LinalgAdRuleSupport::SupportedViaLinearize,
505 LinalgAdRoute::LinearizeThenTranspose,
506 ),
507 LinalgAdRuleSupport::Unsupported,
508 &EIGH_VALS_OUTPUTS,
509 &DECOMPOSITION_CAVEATS,
510 ),
511 support_entry(
512 LinalgAdOpKind::Eig,
513 LinalgAdRuleSupport::PartiallySupported,
514 LinalgAdRuleSupport::Unsupported,
515 mode(
516 LinalgAdRuleSupport::PartiallySupported,
517 LinalgAdRoute::LinearizeThenTranspose,
518 ),
519 LinalgAdRuleSupport::Unsupported,
520 &EIG_OUTPUTS,
521 &DECOMPOSITION_CAVEATS,
522 ),
523 support_entry(
524 LinalgAdOpKind::EigVals,
525 LinalgAdRuleSupport::SupportedViaLinearize,
526 LinalgAdRuleSupport::Unsupported,
527 mode(
528 LinalgAdRuleSupport::SupportedViaLinearize,
529 LinalgAdRoute::LinearizeThenTranspose,
530 ),
531 LinalgAdRuleSupport::Unsupported,
532 &EIG_VALS_OUTPUTS,
533 &DECOMPOSITION_CAVEATS,
534 ),
535 support_entry(
536 LinalgAdOpKind::TriangularSolve,
537 LinalgAdRuleSupport::SupportedViaLinearize,
538 LinalgAdRuleSupport::Supported,
539 mode(
540 LinalgAdRuleSupport::SupportedViaLinearize,
541 LinalgAdRoute::LinearizeThenCustomLinearTranspose,
542 ),
543 LinalgAdRuleSupport::Supported,
544 &SOLUTION_OUTPUTS,
545 &[],
546 ),
547 support_entry(
551 LinalgAdOpKind::SvdFull,
552 LinalgAdRuleSupport::Unsupported,
553 LinalgAdRuleSupport::Unsupported,
554 mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
555 LinalgAdRuleSupport::Unsupported,
556 &SVD_FULL_OUTPUTS,
557 &[],
558 ),
559];
560
561pub fn all_linalg_ad_support() -> &'static [LinalgAdSupport; LinalgAdOpKind::COUNT] {
570 &LINALG_AD_SUPPORT
571}
572
573pub fn linalg_ad_support(kind: LinalgAdOpKind) -> &'static LinalgAdSupport {
584 &LINALG_AD_SUPPORT[kind.as_index()]
585}
586
587#[cfg(test)]
588pub(crate) fn linalg_ad_support_for_op(op: LinalgOp) -> &'static LinalgAdSupport {
589 linalg_ad_support(LinalgAdOpKind::from_linalg_op(op))
590}
591
592#[cfg(test)]
593mod tests;