diff --git a/Cargo.toml b/Cargo.toml index 1bffb81..43fb81d 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -18,7 +18,6 @@ exclude = ["/tests"] algos = "0.6" itertools = "0.15" ndarray = { version = "0.17", features = ["rayon"] } -ndarray-linalg = { version = "0.18", features = ["openblas-system"] } ndarray-npy = { version = "0.10.0", features = ["npz"] } ndrustfft = "0.6" num = "0.4" diff --git a/src/error.rs b/src/error.rs index c013d2c..bac24cd 100644 --- a/src/error.rs +++ b/src/error.rs @@ -9,8 +9,6 @@ pub enum Error { #[error(transparent)] ShapeError(#[from] ndarray::ShapeError), #[error(transparent)] - LinAlg(#[from] ndarray_linalg::error::LinalgError), - #[error(transparent)] NpyError(#[from] ndarray_npy::WriteNpzError), #[error("number of dimensions is not defined")] NumberOfDimensionsNotDefined, diff --git a/src/transform.rs b/src/transform.rs index e574d10..e4895ae 100644 --- a/src/transform.rs +++ b/src/transform.rs @@ -4,7 +4,6 @@ use crate::metric::FixedMu; use crate::register::{Registration, RegistrationResult, RegistrationStep}; use itertools::Itertools; use ndarray::{Array, Array2, ArrayD, AsArray, Dimension, Ix2, s}; -use ndarray_linalg::Inverse; use num::cast::AsPrimitive; use serde::{Deserialize, Serialize}; use serde_yaml::{from_reader, to_writer}; @@ -44,6 +43,63 @@ where .collect() } +/// Calculate the inverse of a square matrix using LU factorization with partial pivoting. +fn matrix_inverse(matrix: Array2) -> Array2 { + let n = matrix.nrows(); + debug_assert_eq!(n, matrix.ncols(), "matrix must be square"); + + // PA = LU factorization (Doolittle with partial pivoting) + let mut lu = matrix; + let mut perm: Vec = (0..n).collect(); + + for k in 0..n { + let mut pivot_row = k; + for i in (k + 1)..n { + if lu[[i, k]].abs() > lu[[pivot_row, k]].abs() { + pivot_row = i; + } + } + assert_ne!(lu[[pivot_row, k]], 0.0, "matrix is singular"); + if pivot_row != k { + for j in 0..n { + lu.swap([k, j], [pivot_row, j]); + } + perm.swap(k, pivot_row); + } + + for i in (k + 1)..n { + let factor = lu[[i, k]] / lu[[k, k]]; + lu[[i, k]] = factor; + for j in (k + 1)..n { + lu[[i, j]] -= factor * lu[[k, j]]; + } + } + } + + // Solve A * inv = I column by column: first L*y = P*e_col, then U*x = y. + let mut inv = Array2::zeros((n, n)); + for col in 0..n { + let mut y = vec![0.0; n]; + for i in 0..n { + let rhs = if perm[i] == col { 1.0 } else { 0.0 }; + let mut sum = rhs; + for j in 0..i { + sum -= lu[[i, j]] * y[j]; + } + y[i] = sum; + } + for i in (0..n).rev() { + let mut sum = y[i]; + for j in (i + 1)..n { + sum -= lu[[i, j]] * inv[[j, col]]; + } + inv[[i, col]] = sum / lu[[i, i]]; + } + } + + inv +} + /// a struct describing the transform /// generic parameter N = # image dimensions #[derive(Clone, Debug, Deserialize, Serialize, PartialEq)] @@ -428,7 +484,7 @@ impl Transform { /// get the inverse transform pub fn inverse(&self) -> Result { let matrix = self.matrix(); - let inverse = matrix.inv()?; + let inverse = matrix_inverse(matrix); let parameters = inverse .slice(s![..self.ndim, ..self.ndim]) .iter() @@ -519,6 +575,7 @@ impl Transform { #[cfg(test)] mod tests { + use super::matrix_inverse; use crate::julia_image; use crate::transform::Transform; use itertools::Itertools; @@ -528,7 +585,40 @@ mod tests { use std::path::Path; use tiff::decoder::{Decoder, DecodingResult}; use tiff::tags::Tag; - use tiffwrite::IJTiffFile; + + #[test] + fn matrix_inverse_test() -> Result<(), Box> { + use ndarray::array; + + let a = array![[4.0, 7.0], [2.0, 6.0]]; + let expected = array![[0.6, -0.7], [-0.2, 0.4]]; + let got = matrix_inverse(a.clone()); + assert!( + got.iter() + .zip_eq(expected.iter()) + .all(|(x, y)| (x - y).abs() < 1e-12) + ); + + let b = array![[0.0, 2.0], [3.0, 4.0]]; + let expected = array![[-2.0 / 3.0, 1.0 / 3.0], [0.5, 0.0]]; + let got = matrix_inverse(b.clone()); + assert!( + got.iter() + .zip_eq(expected.iter()) + .all(|(x, y)| (x - y).abs() < 1e-12) + ); + + let c = array![[1.0, 2.0, 3.0], [0.0, 1.0, 4.0], [5.0, 6.0, 0.0]]; + let reference = array![[-24.0, 18.0, 5.0], [20.0, -15.0, -4.0], [-5.0, 4.0, 1.0]]; + let got = matrix_inverse(c); + assert!( + got.iter() + .zip_eq(reference.iter()) + .all(|(x, y)| (x - y).abs() < 1e-10) + ); + + Ok(()) + } #[test] fn interpolate() -> Result<(), Box> {