diff --git a/matrix/SquareMatrix.hpp b/matrix/SquareMatrix.hpp index fe98cbefb0..912a27a31a 100644 --- a/matrix/SquareMatrix.hpp +++ b/matrix/SquareMatrix.hpp @@ -232,9 +232,8 @@ bool inv(const SquareMatrix & A, SquareMatrix & inv) // divide by the factor // on current // term to be solved - if(fabsf(U(i,i)) < 1e-8f) { - return false; - } + // + // we know that U(i, i) != 0 from above P(i, c) /= U(i, i); } } diff --git a/test/inverse.cpp b/test/inverse.cpp index 0f7ef054ad..3c157388c0 100644 --- a/test/inverse.cpp +++ b/test/inverse.cpp @@ -82,7 +82,30 @@ int main() SquareMatrix A3(data3); SquareMatrix A3_I = inv(A3); SquareMatrix A3_I_check(data3_check); - TEST((A3_I - A3_I_check).abs().max() < 1e-5); + TEST(isEqual(inv(A3), A3_I_check)); + TEST(isEqual(A3_I, A3_I_check)); + TEST(A3.I(A3_I)); + TEST(isEqual(A3_I, A3_I_check)); + + // cover singular matrices + A3(0, 0) = 0; + A3(0, 1) = 0; + A3(0, 2) = 0; + A3_I = inv(A3); + SquareMatrix Z3 = zeros(); + TEST(!A3.I(A3_I)); + TEST(!Z3.I(A3_I)); + TEST(isEqual(A3_I, Z3)); + TEST(isEqual(A3.I(), Z3)); + + // cover NaN + A3(0, 0) = NAN; + A3(0, 1) = 0; + A3(0, 2) = 0; + A3_I = inv(A3); + TEST(isEqual(A3_I, Z3)); + TEST(isEqual(A3.I(), Z3)); + return 0; }