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::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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
175pub struct LinalgAdOutputSupport {
176 pub index: usize,
178 pub name: &'static str,
180 pub status: LinalgAdRuleSupport,
182}
183
184#[derive(Clone, Copy, Debug, PartialEq, Eq)]
195pub struct LinalgAdSupport {
196 pub kind: LinalgAdOpKind,
198 pub jvp: LinalgAdModeSupport,
200 pub vjp: LinalgAdModeSupport,
202 pub linearize_rule: LinalgAdRuleSupport,
204 pub custom_vjp_rule: LinalgAdRuleSupport,
206 pub custom_linear_transpose_rule: LinalgAdRuleSupport,
208 pub linearize: LinalgAdRuleSupport,
210 pub transpose: LinalgAdRuleSupport,
212 pub outputs: &'static [LinalgAdOutputSupport],
214 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 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
565pub fn all_linalg_ad_support() -> &'static [LinalgAdSupport; LinalgAdOpKind::COUNT] {
574 &LINALG_AD_SUPPORT
575}
576
577pub 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;