use torsh_core::{Result as TorshResult, TorshError};
use torsh_tensor::{
creation::{eye, ones},
Tensor,
};
pub fn lu(tensor: &Tensor) -> TorshResult<(Tensor, Tensor, Tensor)> {
if tensor.shape().ndim() != 2 {
return Err(TorshError::invalid_argument_with_context(
"LU decomposition requires 2D tensor",
"lu",
));
}
let n = tensor.shape().dims()[0];
let m = tensor.shape().dims()[1];
if n != m {
return Err(TorshError::invalid_argument_with_context(
"LU decomposition requires square matrix",
"lu",
));
}
let p = eye(n)?;
let l = eye(n)?;
let u = tensor.clone();
Ok((p, l, u))
}
pub fn qr(tensor: &Tensor, reduced: bool) -> TorshResult<(Tensor, Tensor)> {
if tensor.shape().ndim() != 2 {
return Err(TorshError::invalid_argument_with_context(
"QR decomposition requires 2D tensor",
"qr",
));
}
let (q, r) = torsh_linalg::decomposition::qr(tensor)?;
if reduced {
Ok((q, r))
} else {
let shape = tensor.shape();
let dims = shape.dims();
let m = dims[0];
let n = dims[1];
let k = m.min(n);
if k == m {
Ok((q, r))
} else {
let mut q_data = vec![0.0f32; m * m];
let q_vec = q.to_vec()?;
for i in 0..m {
for j in 0..k {
q_data[i * m + j] = q_vec[i * k + j];
}
}
for i in k..m {
q_data[i * m + i] = 1.0;
}
let q_full = Tensor::from_data(q_data, vec![m, m], tensor.device())?;
Ok((q_full, r))
}
}
}
pub fn cholesky(tensor: &Tensor, upper: bool) -> TorshResult<Tensor> {
if tensor.shape().ndim() != 2 {
return Err(TorshError::invalid_argument_with_context(
"Cholesky decomposition requires 2D tensor",
"cholesky",
));
}
let binding = tensor.shape();
let dims = binding.dims();
if dims[0] != dims[1] {
return Err(TorshError::invalid_argument_with_context(
"Cholesky decomposition requires square matrix",
"cholesky",
));
}
torsh_linalg::decomposition::cholesky(tensor, upper)
}
pub fn svd(tensor: &Tensor, full_matrices: bool) -> TorshResult<(Tensor, Tensor, Tensor)> {
if tensor.shape().ndim() != 2 {
return Err(TorshError::invalid_argument_with_context(
"SVD requires 2D tensor",
"svd",
));
}
torsh_linalg::decomposition::svd(tensor, full_matrices)
}
pub fn eig(tensor: &Tensor) -> TorshResult<(Tensor, Tensor)> {
use scirs2_core::ndarray::Array2;
if tensor.shape().ndim() != 2 {
return Err(TorshError::invalid_argument_with_context(
"Eigenvalue decomposition requires 2D tensor",
"eig",
));
}
let shape = tensor.shape();
let dims = shape.dims();
if dims[0] != dims[1] {
return Err(TorshError::invalid_argument_with_context(
"Eigenvalue decomposition requires square matrix",
"eig",
));
}
let n = dims[0];
let device = tensor.device();
let data: Vec<f64> = tensor.to_vec()?.into_iter().map(|x| x as f64).collect();
let decomposed: Vec<f64>;
let mut pairs: Vec<(f64, Vec<f64>)> = if matrix_is_symmetric(&data, n) {
let symmetric = symmetrize(&data, n);
let a = Array2::from_shape_vec((n, n), symmetric.clone()).map_err(|e| {
TorshError::ComputeError(format!("eig: matrix construction failed: {e}"))
})?;
decomposed = symmetric;
let (values, vectors) = scirs2_linalg::eigh(&a.view(), None).map_err(|e| {
TorshError::ComputeError(format!("eig: symmetric eigensolver failed: {e}"))
})?;
if values.len() != n {
return Err(TorshError::ComputeError(format!(
"eig: solver returned {} eigenvalues for an {n}x{n} matrix",
values.len()
)));
}
(0..n)
.map(|j| {
let column = normalized((0..n).map(|i| vectors[[i, j]]).collect());
(values[j], column)
})
.collect()
} else {
let a = Array2::from_shape_vec((n, n), data.clone()).map_err(|e| {
TorshError::ComputeError(format!("eig: matrix construction failed: {e}"))
})?;
decomposed = data.clone();
let (values, vectors) = scirs2_linalg::eig(&a.view(), None)
.map_err(|e| TorshError::ComputeError(format!("eig: eigensolver failed: {e}")))?;
if values.len() != n {
return Err(TorshError::ComputeError(format!(
"eig: solver returned {} eigenvalues for an {n}x{n} matrix",
values.len()
)));
}
let scale = data.iter().fold(1.0f64, |m, &x| m.max(x.abs()));
for k in 0..n {
if values[k].im.abs() > 1e-6 * scale {
return Err(TorshError::ComputeError(format!(
"eig: matrix has a complex eigenvalue ({:.6}{:+.6}i); a real \
eigendecomposition does not exist for this matrix",
values[k].re, values[k].im
)));
}
}
(0..n)
.map(|j| {
let re: Vec<f64> = (0..n).map(|i| vectors[[i, j]].re).collect();
let im: Vec<f64> = (0..n).map(|i| vectors[[i, j]].im).collect();
let re_norm = re.iter().map(|x| x * x).sum::<f64>().sqrt();
let im_norm = im.iter().map(|x| x * x).sum::<f64>().sqrt();
let column = normalized(if re_norm >= im_norm { re } else { im });
(values[j].re, column)
})
.collect()
};
pairs.sort_by(|a, b| b.0.total_cmp(&a.0));
let residual = max_relative_residual(&decomposed, n, &pairs);
if residual > 1e-3 {
return Err(TorshError::ComputeError(format!(
"eig: failed to compute an accurate eigendecomposition (max relative \
residual {residual:.3e}); the matrix is likely defective / \
non-diagonalizable"
)));
}
let mut eigenvalue_data = Vec::with_capacity(n);
let mut eigenvector_data = vec![0.0f32; n * n];
for (col, (lambda, vector)) in pairs.iter().enumerate() {
eigenvalue_data.push(*lambda as f32);
for (row, &value) in vector.iter().enumerate() {
eigenvector_data[row * n + col] = value as f32;
}
}
let eigenvalues = Tensor::from_data(eigenvalue_data, vec![n], device)?;
let eigenvectors = Tensor::from_data(eigenvector_data, vec![n, n], device)?;
Ok((eigenvalues, eigenvectors))
}
fn matrix_is_symmetric(data: &[f64], n: usize) -> bool {
for i in 0..n {
for j in (i + 1)..n {
let upper = data[i * n + j];
let lower = data[j * n + i];
let scale = upper.abs().max(lower.abs()).max(1.0);
if (upper - lower).abs() > 1e-6 * scale {
return false;
}
}
}
true
}
fn symmetrize(data: &[f64], n: usize) -> Vec<f64> {
let mut out = vec![0.0f64; n * n];
for i in 0..n {
for j in 0..n {
out[i * n + j] = 0.5 * (data[i * n + j] + data[j * n + i]);
}
}
out
}
fn normalized(mut v: Vec<f64>) -> Vec<f64> {
let norm = v.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm > 1e-300 {
for x in v.iter_mut() {
*x /= norm;
}
}
v
}
fn max_relative_residual(matrix: &[f64], n: usize, pairs: &[(f64, Vec<f64>)]) -> f64 {
let mut worst = 0.0f64;
for pair in pairs {
let lambda = pair.0;
let v = &pair.1;
let mut residual_sq = 0.0f64;
for i in 0..n {
let mut av_i = 0.0f64;
for j in 0..n {
av_i += matrix[i * n + j] * v[j];
}
let diff = av_i - lambda * v[i];
residual_sq += diff * diff;
}
let relative = residual_sq.sqrt() / (lambda.abs() + 1.0);
if relative > worst {
worst = relative;
}
}
worst
}
pub fn svd_lowrank(
tensor: &Tensor,
rank: Option<usize>,
niter: Option<usize>,
) -> TorshResult<(Tensor, Tensor, Tensor)> {
if tensor.shape().ndim() != 2 {
return Err(TorshError::invalid_argument_with_context(
"SVD requires 2D tensor",
"svd_lowrank",
));
}
let (m, n) = (tensor.shape().dims()[0], tensor.shape().dims()[1]);
let k = rank.unwrap_or(m.min(n).min(10));
let _niter = niter.unwrap_or(2);
use torsh_tensor::creation::randn;
let u = randn(&[m, k])?;
let s = ones(&[k])?;
let v = randn(&[k, n])?;
Ok((u, s, v))
}
pub fn pca_lowrank(
tensor: &Tensor,
rank: Option<usize>,
center: bool,
) -> TorshResult<(Tensor, Tensor, Tensor)> {
if tensor.shape().ndim() != 2 {
return Err(TorshError::invalid_argument_with_context(
"PCA requires 2D tensor",
"pca_lowrank",
));
}
let (m, n) = (tensor.shape().dims()[0], tensor.shape().dims()[1]);
let k = rank.unwrap_or(m.min(n));
let mut data_tensor = tensor.clone();
if center {
data_tensor = tensor.clone();
}
let (u, s, v) = svd_lowrank(&data_tensor, Some(k), None)?;
Ok((u, s, v.transpose(-2, -1)?))
}
#[cfg(test)]
mod tests {
use super::*;
use torsh_tensor::creation::eye;
#[test]
fn test_eig_identity_returns_ones() {
let identity = eye::<f32>(4).unwrap();
let (eigenvalues, eigenvectors) = eig(&identity).unwrap();
assert_eq!(eigenvalues.shape().dims(), &[4]);
assert_eq!(eigenvectors.shape().dims(), &[4, 4]);
for &lambda in eigenvalues.to_vec().unwrap().iter() {
assert!(
(lambda - 1.0).abs() < 1e-5,
"identity eigenvalue should be 1.0, got {lambda}"
);
}
}
#[test]
fn test_eig_diagonal_exact() {
let diag = Tensor::from_data(
vec![3.0, 0.0, 0.0, 0.0, -2.0, 0.0, 0.0, 0.0, 5.0],
vec![3, 3],
torsh_core::device::DeviceType::Cpu,
)
.unwrap();
let (eigenvalues, _vectors) = eig(&diag).unwrap();
let mut values = eigenvalues.to_vec().unwrap();
values.sort_by(|a, b| a.partial_cmp(b).unwrap());
let expected = [-2.0_f32, 3.0, 5.0];
for (got, want) in values.iter().zip(expected.iter()) {
assert!(
(got - want).abs() < 1e-5,
"diagonal eigenvalue mismatch: got {got}, want {want}"
);
}
}
#[test]
fn test_eig_non_square_errors() {
let rect = Tensor::from_data(
vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0],
vec![2, 3],
torsh_core::device::DeviceType::Cpu,
)
.unwrap();
assert!(eig(&rect).is_err());
}
#[test]
fn test_eig_dominant_eigenpair_residual() {
let a = Tensor::from_data(
vec![2.0, 1.0, 1.0, 2.0],
vec![2, 2],
torsh_core::device::DeviceType::Cpu,
)
.unwrap();
let (eigenvalues, eigenvectors) = eig(&a).unwrap();
let lambda = eigenvalues.to_vec().unwrap()[0];
let vecs = eigenvectors.to_vec().unwrap();
let v = [vecs[0], vecs[2]];
let av0 = 2.0 * v[0] + 1.0 * v[1];
let av1 = 1.0 * v[0] + 2.0 * v[1];
assert!(
(av0 - lambda * v[0]).abs() < 1e-3 && (av1 - lambda * v[1]).abs() < 1e-3,
"A v should equal lambda v (lambda={lambda}, v=[{}, {}])",
v[0],
v[1]
);
assert!(
(lambda - 3.0).abs() < 1e-3,
"dominant eigenvalue should be 3.0, got {lambda}"
);
}
fn assert_eigenpairs_valid(
a: &Tensor,
eigenvalues: &Tensor,
eigenvectors: &Tensor,
n: usize,
tol: f32,
) {
let a_data = a.to_vec().unwrap();
let vals = eigenvalues.to_vec().unwrap();
let vecs = eigenvectors.to_vec().unwrap();
for col in 0..n {
let lambda = vals[col];
let v: Vec<f32> = (0..n).map(|row| vecs[row * n + col]).collect();
let v_norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
v_norm > 0.5,
"eigenvector {col} is degenerate (norm {v_norm})"
);
for row in 0..n {
let mut av = 0.0f32;
for k in 0..n {
av += a_data[row * n + k] * v[k];
}
assert!(
(av - lambda * v[row]).abs() < tol,
"A v != lambda v at row {row}, col {col}: Av={av}, lambda*v={}",
lambda * v[row]
);
}
}
}
#[test]
fn test_eig_repeated_eigenvalue() {
let a = Tensor::from_data(
vec![3.0, 1.0, 1.0, 1.0, 3.0, 1.0, 1.0, 1.0, 3.0],
vec![3, 3],
torsh_core::device::DeviceType::Cpu,
)
.unwrap();
let (eigenvalues, eigenvectors) = eig(&a).unwrap();
assert_eq!(eigenvalues.shape().dims(), &[3]);
assert_eq!(eigenvectors.shape().dims(), &[3, 3]);
let mut values = eigenvalues.to_vec().unwrap();
values.sort_by(|x, y| x.partial_cmp(y).unwrap());
let expected = [2.0_f32, 2.0, 5.0];
for (got, want) in values.iter().zip(expected.iter()) {
assert!(
(got - want).abs() < 1e-3,
"repeated-eigenvalue spectrum mismatch: got {got}, want {want}"
);
}
assert_eigenpairs_valid(&a, &eigenvalues, &eigenvectors, 3, 1e-3);
}
#[test]
fn test_eig_rank_deficient_zero_eigenvalue() {
let a = Tensor::from_data(
vec![1.0; 9],
vec![3, 3],
torsh_core::device::DeviceType::Cpu,
)
.unwrap();
let (eigenvalues, eigenvectors) = eig(&a).unwrap();
let mut values = eigenvalues.to_vec().unwrap();
values.sort_by(|x, y| x.partial_cmp(y).unwrap());
let expected = [0.0_f32, 0.0, 3.0];
for (got, want) in values.iter().zip(expected.iter()) {
assert!(
(got - want).abs() < 1e-3,
"rank-deficient spectrum mismatch: got {got}, want {want}"
);
}
assert_eigenpairs_valid(&a, &eigenvalues, &eigenvectors, 3, 1e-3);
}
}