- replace ndarray-linalg matrix inverse by a function that does not require blas
This commit is contained in:
@@ -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,
|
||||
|
||||
+93
-3
@@ -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<f64>) -> Array2<f64> {
|
||||
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<usize> = (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<D: Dimension> Transform<D> {
|
||||
/// get the inverse transform
|
||||
pub fn inverse(&self) -> Result<Self, Error> {
|
||||
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<Ix2> {
|
||||
|
||||
#[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<dyn std::error::Error>> {
|
||||
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<dyn std::error::Error>> {
|
||||
|
||||
Reference in New Issue
Block a user