#[cfg(test)]
use crate::extension::LinalgOp;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum LinalgAdRuleSupport {
Supported,
SupportedViaLinearize,
PartiallySupported,
NonDifferentiable,
Unsupported,
PendingOracle,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum LinalgAdRoute {
Unsupported,
Linearize,
LinearizeThenTranspose,
LinearizeThenCustomLinearTranspose,
CustomVjp,
CustomPreferredWithLinearizeFallback,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct LinalgAdModeSupport {
pub status: LinalgAdRuleSupport,
pub route: LinalgAdRoute,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum LinalgAdOpKind {
Cholesky,
Lu,
LuFactor,
SignDetFromLuFactor,
LogAbsDetFromLuFactor,
LuSolvePrepared,
FullPivLu,
FullPivLuSolve,
Svd,
SvdVals,
Qr,
Eigh,
EighVals,
Eig,
EigVals,
TriangularSolve,
SvdFull,
}
impl LinalgAdOpKind {
pub const COUNT: usize = 17;
pub const fn as_index(self) -> usize {
match self {
Self::Cholesky => 0,
Self::Lu => 1,
Self::LuFactor => 2,
Self::SignDetFromLuFactor => 3,
Self::LogAbsDetFromLuFactor => 4,
Self::LuSolvePrepared => 5,
Self::FullPivLu => 6,
Self::FullPivLuSolve => 7,
Self::Svd => 8,
Self::SvdVals => 9,
Self::Qr => 10,
Self::Eigh => 11,
Self::EighVals => 12,
Self::Eig => 13,
Self::EigVals => 14,
Self::TriangularSolve => 15,
Self::SvdFull => 16,
}
}
#[cfg(test)]
pub(crate) const fn from_linalg_op(op: LinalgOp) -> Self {
match op {
LinalgOp::Cholesky => Self::Cholesky,
LinalgOp::Lu => Self::Lu,
LinalgOp::LuFactor => Self::LuFactor,
LinalgOp::SignDetFromLuFactor => Self::SignDetFromLuFactor,
LinalgOp::LogAbsDetFromLuFactor => Self::LogAbsDetFromLuFactor,
LinalgOp::LuSolvePrepared { .. } => Self::LuSolvePrepared,
LinalgOp::FullPivLu => Self::FullPivLu,
LinalgOp::FullPivLuSolve { .. } => Self::FullPivLuSolve,
LinalgOp::Svd { .. } => Self::Svd,
LinalgOp::SvdFull => Self::SvdFull,
LinalgOp::SvdVals { .. } => Self::SvdVals,
LinalgOp::Qr { .. } => Self::Qr,
LinalgOp::Eigh { .. } => Self::Eigh,
LinalgOp::EighVals { .. } => Self::EighVals,
LinalgOp::Eig { .. } => Self::Eig,
LinalgOp::EigVals { .. } => Self::EigVals,
LinalgOp::TriangularSolve { .. } => Self::TriangularSolve,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct LinalgAdOutputSupport {
pub index: usize,
pub name: &'static str,
pub status: LinalgAdRuleSupport,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct LinalgAdSupport {
pub kind: LinalgAdOpKind,
pub jvp: LinalgAdModeSupport,
pub vjp: LinalgAdModeSupport,
pub linearize_rule: LinalgAdRuleSupport,
pub custom_vjp_rule: LinalgAdRuleSupport,
pub custom_linear_transpose_rule: LinalgAdRuleSupport,
pub linearize: LinalgAdRuleSupport,
pub transpose: LinalgAdRuleSupport,
pub outputs: &'static [LinalgAdOutputSupport],
pub caveats: &'static [&'static str],
}
const fn mode(status: LinalgAdRuleSupport, route: LinalgAdRoute) -> LinalgAdModeSupport {
LinalgAdModeSupport { status, route }
}
const fn jvp_route(status: LinalgAdRuleSupport) -> LinalgAdRoute {
match status {
LinalgAdRuleSupport::Unsupported
| LinalgAdRuleSupport::NonDifferentiable
| LinalgAdRuleSupport::PendingOracle => LinalgAdRoute::Unsupported,
LinalgAdRuleSupport::Supported
| LinalgAdRuleSupport::SupportedViaLinearize
| LinalgAdRuleSupport::PartiallySupported => LinalgAdRoute::Linearize,
}
}
const fn support_entry(
kind: LinalgAdOpKind,
linearize: LinalgAdRuleSupport,
transpose: LinalgAdRuleSupport,
vjp: LinalgAdModeSupport,
custom_linear_transpose_rule: LinalgAdRuleSupport,
outputs: &'static [LinalgAdOutputSupport],
caveats: &'static [&'static str],
) -> LinalgAdSupport {
LinalgAdSupport {
kind,
jvp: mode(linearize, jvp_route(linearize)),
vjp,
linearize_rule: linearize,
custom_vjp_rule: LinalgAdRuleSupport::Unsupported,
custom_linear_transpose_rule,
linearize,
transpose,
outputs,
caveats,
}
}
const fn output(
index: usize,
name: &'static str,
status: LinalgAdRuleSupport,
) -> LinalgAdOutputSupport {
LinalgAdOutputSupport {
index,
name,
status,
}
}
static CHOLESKY_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
0,
"factor",
LinalgAdRuleSupport::SupportedViaLinearize,
)];
static LU_OUTPUTS: [LinalgAdOutputSupport; 4] = [
output(0, "p", LinalgAdRuleSupport::NonDifferentiable),
output(1, "l", LinalgAdRuleSupport::SupportedViaLinearize),
output(2, "u", LinalgAdRuleSupport::SupportedViaLinearize),
output(3, "parity", LinalgAdRuleSupport::NonDifferentiable),
];
static LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 3] = [
output(0, "packed_lu", LinalgAdRuleSupport::Unsupported),
output(1, "pivots", LinalgAdRuleSupport::NonDifferentiable),
output(2, "parity", LinalgAdRuleSupport::NonDifferentiable),
];
static SIGNDET_FROM_LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
0,
"sign",
LinalgAdRuleSupport::SupportedViaLinearize,
)];
static LOGABSDET_FROM_LU_FACTOR_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
0,
"logabsdet",
LinalgAdRuleSupport::SupportedViaLinearize,
)];
static SOLUTION_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
0,
"solution",
LinalgAdRuleSupport::SupportedViaLinearize,
)];
static FULL_PIV_LU_OUTPUTS: [LinalgAdOutputSupport; 5] = [
output(0, "p", LinalgAdRuleSupport::NonDifferentiable),
output(1, "l", LinalgAdRuleSupport::SupportedViaLinearize),
output(2, "u", LinalgAdRuleSupport::SupportedViaLinearize),
output(3, "q", LinalgAdRuleSupport::NonDifferentiable),
output(4, "parity", LinalgAdRuleSupport::NonDifferentiable),
];
static FULL_PIV_LU_SOLVE_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
0,
"solution",
LinalgAdRuleSupport::SupportedViaLinearize,
)];
static SVD_OUTPUTS: [LinalgAdOutputSupport; 3] = [
output(0, "u", LinalgAdRuleSupport::SupportedViaLinearize),
output(
1,
"singular_values",
LinalgAdRuleSupport::SupportedViaLinearize,
),
output(2, "vt", LinalgAdRuleSupport::SupportedViaLinearize),
];
static SVD_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
0,
"singular_values",
LinalgAdRuleSupport::SupportedViaLinearize,
)];
static SVD_FULL_OUTPUTS: [LinalgAdOutputSupport; 3] = [
output(0, "u", LinalgAdRuleSupport::Unsupported),
output(1, "singular_values", LinalgAdRuleSupport::Unsupported),
output(2, "vt", LinalgAdRuleSupport::Unsupported),
];
static QR_OUTPUTS: [LinalgAdOutputSupport; 2] = [
output(0, "q", LinalgAdRuleSupport::SupportedViaLinearize),
output(1, "r", LinalgAdRuleSupport::SupportedViaLinearize),
];
static EIGH_OUTPUTS: [LinalgAdOutputSupport; 2] = [
output(0, "eigenvalues", LinalgAdRuleSupport::SupportedViaLinearize),
output(
1,
"eigenvectors",
LinalgAdRuleSupport::SupportedViaLinearize,
),
];
static EIGH_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
0,
"eigenvalues",
LinalgAdRuleSupport::SupportedViaLinearize,
)];
static EIG_OUTPUTS: [LinalgAdOutputSupport; 2] = [
output(0, "eigenvalues", LinalgAdRuleSupport::SupportedViaLinearize),
output(1, "eigenvectors", LinalgAdRuleSupport::Unsupported),
];
static EIG_VALS_OUTPUTS: [LinalgAdOutputSupport; 1] = [output(
0,
"eigenvalues",
LinalgAdRuleSupport::SupportedViaLinearize,
)];
static DECOMPOSITION_CAVEATS: [&str; 1] = [
"Derivative regularization handles near-degenerate spectra but does not make exact degeneracies smoothly differentiable.",
];
static LINALG_AD_SUPPORT: [LinalgAdSupport; LinalgAdOpKind::COUNT] = [
support_entry(
LinalgAdOpKind::Cholesky,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&CHOLESKY_OUTPUTS,
&DECOMPOSITION_CAVEATS,
),
support_entry(
LinalgAdOpKind::Lu,
LinalgAdRuleSupport::PartiallySupported,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::PartiallySupported,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&LU_OUTPUTS,
&DECOMPOSITION_CAVEATS,
),
support_entry(
LinalgAdOpKind::LuFactor,
LinalgAdRuleSupport::Unsupported,
LinalgAdRuleSupport::Unsupported,
mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
LinalgAdRuleSupport::Unsupported,
&LU_FACTOR_OUTPUTS,
&[],
),
support_entry(
LinalgAdOpKind::SignDetFromLuFactor,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&SIGNDET_FROM_LU_FACTOR_OUTPUTS,
&[],
),
support_entry(
LinalgAdOpKind::LogAbsDetFromLuFactor,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&LOGABSDET_FROM_LU_FACTOR_OUTPUTS,
&[],
),
support_entry(
LinalgAdOpKind::LuSolvePrepared,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::PartiallySupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenCustomLinearTranspose,
),
LinalgAdRuleSupport::PartiallySupported,
&SOLUTION_OUTPUTS,
&[],
),
support_entry(
LinalgAdOpKind::FullPivLu,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&FULL_PIV_LU_OUTPUTS,
&DECOMPOSITION_CAVEATS,
),
support_entry(
LinalgAdOpKind::FullPivLuSolve,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Supported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenCustomLinearTranspose,
),
LinalgAdRuleSupport::Supported,
&FULL_PIV_LU_SOLVE_OUTPUTS,
&[],
),
support_entry(
LinalgAdOpKind::Svd,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&SVD_OUTPUTS,
&DECOMPOSITION_CAVEATS,
),
support_entry(
LinalgAdOpKind::SvdVals,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&SVD_VALS_OUTPUTS,
&DECOMPOSITION_CAVEATS,
),
support_entry(
LinalgAdOpKind::Qr,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&QR_OUTPUTS,
&DECOMPOSITION_CAVEATS,
),
support_entry(
LinalgAdOpKind::Eigh,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&EIGH_OUTPUTS,
&DECOMPOSITION_CAVEATS,
),
support_entry(
LinalgAdOpKind::EighVals,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&EIGH_VALS_OUTPUTS,
&DECOMPOSITION_CAVEATS,
),
support_entry(
LinalgAdOpKind::Eig,
LinalgAdRuleSupport::PartiallySupported,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::PartiallySupported,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&EIG_OUTPUTS,
&DECOMPOSITION_CAVEATS,
),
support_entry(
LinalgAdOpKind::EigVals,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Unsupported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenTranspose,
),
LinalgAdRuleSupport::Unsupported,
&EIG_VALS_OUTPUTS,
&DECOMPOSITION_CAVEATS,
),
support_entry(
LinalgAdOpKind::TriangularSolve,
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRuleSupport::Supported,
mode(
LinalgAdRuleSupport::SupportedViaLinearize,
LinalgAdRoute::LinearizeThenCustomLinearTranspose,
),
LinalgAdRuleSupport::Supported,
&SOLUTION_OUTPUTS,
&[],
),
support_entry(
LinalgAdOpKind::SvdFull,
LinalgAdRuleSupport::Unsupported,
LinalgAdRuleSupport::Unsupported,
mode(LinalgAdRuleSupport::Unsupported, LinalgAdRoute::Unsupported),
LinalgAdRuleSupport::Unsupported,
&SVD_FULL_OUTPUTS,
&[],
),
];
pub fn all_linalg_ad_support() -> &'static [LinalgAdSupport; LinalgAdOpKind::COUNT] {
&LINALG_AD_SUPPORT
}
pub fn linalg_ad_support(kind: LinalgAdOpKind) -> &'static LinalgAdSupport {
&LINALG_AD_SUPPORT[kind.as_index()]
}
#[cfg(test)]
pub(crate) fn linalg_ad_support_for_op(op: LinalgOp) -> &'static LinalgAdSupport {
linalg_ad_support(LinalgAdOpKind::from_linalg_op(op))
}
#[cfg(test)]
mod tests;