use std::sync::Arc;
use tenferro_ad::error::{Error, Result};
use tenferro_ad::extension::apply_eager;
use tenferro_ad::EagerTensor;
use crate::extension::{LinalgExtensionOp, LinalgOp, DEFAULT_DECOMPOSITION_AD_EPS};
use crate::register_runtime;
pub trait EagerTensorLinalgExt {
fn svd(&self) -> Result<(EagerTensor, EagerTensor, EagerTensor)>;
fn qr(&self) -> Result<(EagerTensor, EagerTensor)>;
fn lu(&self) -> Result<(EagerTensor, EagerTensor, EagerTensor, EagerTensor)>;
fn full_piv_lu(
&self,
) -> Result<(
EagerTensor,
EagerTensor,
EagerTensor,
EagerTensor,
EagerTensor,
)>;
fn full_piv_lu_solve(&self, b: &EagerTensor) -> Result<EagerTensor>;
fn solve(&self, b: &EagerTensor) -> Result<EagerTensor>;
fn cholesky(&self) -> Result<EagerTensor>;
fn eigh(&self) -> Result<(EagerTensor, EagerTensor)>;
fn eig(&self) -> Result<(EagerTensor, EagerTensor)>;
fn triangular_solve(
&self,
b: &EagerTensor,
left_side: bool,
lower: bool,
transpose_a: bool,
unit_diagonal: bool,
) -> Result<EagerTensor>;
}
impl EagerTensorLinalgExt for EagerTensor {
fn svd(&self) -> Result<(EagerTensor, EagerTensor, EagerTensor)> {
svd(self)
}
fn qr(&self) -> Result<(EagerTensor, EagerTensor)> {
qr(self)
}
fn lu(&self) -> Result<(EagerTensor, EagerTensor, EagerTensor, EagerTensor)> {
lu(self)
}
fn full_piv_lu(
&self,
) -> Result<(
EagerTensor,
EagerTensor,
EagerTensor,
EagerTensor,
EagerTensor,
)> {
full_piv_lu(self)
}
fn full_piv_lu_solve(&self, b: &EagerTensor) -> Result<EagerTensor> {
full_piv_lu_solve(self, b)
}
fn solve(&self, b: &EagerTensor) -> Result<EagerTensor> {
solve(self, b)
}
fn cholesky(&self) -> Result<EagerTensor> {
cholesky(self)
}
fn eigh(&self) -> Result<(EagerTensor, EagerTensor)> {
eigh(self)
}
fn eig(&self) -> Result<(EagerTensor, EagerTensor)> {
eig(self)
}
fn triangular_solve(
&self,
b: &EagerTensor,
left_side: bool,
lower: bool,
transpose_a: bool,
unit_diagonal: bool,
) -> Result<EagerTensor> {
triangular_solve(self, b, left_side, lower, transpose_a, unit_diagonal)
}
}
fn apply_linalg_eager(op: LinalgOp, inputs: &[&EagerTensor]) -> Result<Vec<EagerTensor>> {
if let Some(first) = inputs.first() {
first
.runtime()
.register_extension(register_runtime)
.map_err(|err| Error::Internal(err.to_string()))?;
}
apply_eager(Arc::new(LinalgExtensionOp::new(op)), inputs)
}
pub fn svd(a: &EagerTensor) -> Result<(EagerTensor, EagerTensor, EagerTensor)> {
let mut outputs = apply_linalg_eager(
LinalgOp::Svd {
eps: DEFAULT_DECOMPOSITION_AD_EPS,
},
&[a],
)?
.into_iter();
match (
outputs.next(),
outputs.next(),
outputs.next(),
outputs.next(),
) {
(Some(u), Some(s), Some(vt), None) => Ok((u, s, vt)),
_ => Err(Error::Internal(
"svd eager op returned an unexpected number of outputs".to_string(),
)),
}
}
pub fn qr(a: &EagerTensor) -> Result<(EagerTensor, EagerTensor)> {
two_outputs(apply_linalg_eager(LinalgOp::Qr, &[a])?, "qr")
}
pub fn lu(a: &EagerTensor) -> Result<(EagerTensor, EagerTensor, EagerTensor, EagerTensor)> {
let mut outputs = apply_linalg_eager(LinalgOp::Lu, &[a])?.into_iter();
match (
outputs.next(),
outputs.next(),
outputs.next(),
outputs.next(),
outputs.next(),
) {
(Some(p), Some(l), Some(u), Some(parity), None) => Ok((p, l, u, parity)),
_ => Err(Error::Internal(
"lu eager op returned an unexpected number of outputs".to_string(),
)),
}
}
pub fn full_piv_lu(
a: &EagerTensor,
) -> Result<(
EagerTensor,
EagerTensor,
EagerTensor,
EagerTensor,
EagerTensor,
)> {
let mut outputs = apply_linalg_eager(LinalgOp::FullPivLu, &[a])?.into_iter();
match (
outputs.next(),
outputs.next(),
outputs.next(),
outputs.next(),
outputs.next(),
outputs.next(),
) {
(Some(p), Some(l), Some(u), Some(q), Some(parity), None) => Ok((p, l, u, q, parity)),
_ => Err(Error::Internal(
"full_piv_lu eager op returned an unexpected number of outputs".to_string(),
)),
}
}
pub fn full_piv_lu_solve(a: &EagerTensor, b: &EagerTensor) -> Result<EagerTensor> {
one_output(
apply_linalg_eager(LinalgOp::FullPivLuSolve { transpose_a: false }, &[a, b])?,
"full_piv_lu_solve",
)
}
pub fn solve(a: &EagerTensor, b: &EagerTensor) -> Result<EagerTensor> {
let mut factor_outputs = apply_linalg_eager(LinalgOp::LuFactor, &[a])?.into_iter();
let (packed_lu, pivots) = match (
factor_outputs.next(),
factor_outputs.next(),
factor_outputs.next(),
factor_outputs.next(),
) {
(Some(packed_lu), Some(pivots), Some(_parity), None) => (packed_lu, pivots),
_ => {
return Err(Error::Internal(
"lu_factor eager op returned an unexpected number of outputs".to_string(),
));
}
};
one_output(
apply_linalg_eager(
LinalgOp::LuSolvePrepared {
transpose_a: false,
conjugate_a: false,
},
&[a, &packed_lu, &pivots, b],
)?,
"solve",
)
}
pub fn cholesky(a: &EagerTensor) -> Result<EagerTensor> {
one_output(apply_linalg_eager(LinalgOp::Cholesky, &[a])?, "cholesky")
}
pub fn eigh(a: &EagerTensor) -> Result<(EagerTensor, EagerTensor)> {
two_outputs(
apply_linalg_eager(
LinalgOp::Eigh {
eps: DEFAULT_DECOMPOSITION_AD_EPS,
},
&[a],
)?,
"eigh",
)
}
pub fn eig(a: &EagerTensor) -> Result<(EagerTensor, EagerTensor)> {
two_outputs(
apply_linalg_eager(
LinalgOp::Eig {
input_dtype: a.dtype(),
},
&[a],
)?,
"eig",
)
}
pub fn triangular_solve(
a: &EagerTensor,
b: &EagerTensor,
left_side: bool,
lower: bool,
transpose_a: bool,
unit_diagonal: bool,
) -> Result<EagerTensor> {
one_output(
apply_linalg_eager(
LinalgOp::TriangularSolve {
left_side,
lower,
transpose_a,
unit_diagonal,
},
&[a, b],
)?,
"triangular_solve",
)
}
fn one_output(outputs: Vec<EagerTensor>, name: &str) -> Result<EagerTensor> {
let mut outputs = outputs.into_iter();
match (outputs.next(), outputs.next()) {
(Some(output), None) => Ok(output),
_ => Err(Error::Internal(format!(
"{name} eager op returned an unexpected number of outputs"
))),
}
}
fn two_outputs(outputs: Vec<EagerTensor>, name: &str) -> Result<(EagerTensor, EagerTensor)> {
let mut outputs = outputs.into_iter();
match (outputs.next(), outputs.next(), outputs.next()) {
(Some(lhs), Some(rhs), None) => Ok((lhs, rhs)),
_ => Err(Error::Internal(format!(
"{name} eager op returned an unexpected number of outputs"
))),
}
}