- replace ndarray-linalg matrix inverse by a function that does not require blas

This commit is contained in:
w.pomp
2026-07-31 15:28:41 +02:00
parent fdac4e04c4
commit db2b72c3bf
3 changed files with 93 additions and 6 deletions
-1
View File
@@ -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"
-2
View File
@@ -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
View File
@@ -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>> {