Add determinant and sign-log-determinant functions to mlx.core.linalg (#3416)

Co-authored-by: Lucas Fernandes Martins <[email protected]>
This commit is contained in:
Abhilash Shankarampeta
2026-05-05 09:06:23 +09:00
committed by GitHub
co-authored by Lucas Fernandes Martins
parent e8ebdebeeb
commit 0938db7e54
7 changed files with 568 additions and 4 deletions
+65
View File
@@ -637,3 +637,68 @@ TEST_CASE("test solve_triangluar") {
expected = array({-3., 2., 3.});
CHECK(allclose(expected, result).item<bool>());
}
TEST_CASE("test det") {
// 1x1 fast path
{
array a = array({5.0f}, {1, 1});
auto d = det(a, Device::cpu);
CHECK_EQ(d.item<float>(), doctest::Approx(5.0f));
}
// 2x2 fast path: det([[1,2],[3,4]]) = -2
{
array a = array({1.0f, 2.0f, 3.0f, 4.0f}, {2, 2});
auto d = det(a, Device::cpu);
CHECK_EQ(d.item<float>(), doctest::Approx(-2.0f));
}
// 3x3 fast path: det([[1,2,3],[0,1,4],[5,6,0]]) = 1
{
array a =
array({1.0f, 2.0f, 3.0f, 0.0f, 1.0f, 4.0f, 5.0f, 6.0f, 0.0f}, {3, 3});
auto d = det(a, Device::cpu);
CHECK_EQ(d.item<float>(), doctest::Approx(1.0f));
}
// 4x4 LU path: identity matrix det = 1
{
array a = eye(4);
auto d = det(a, Device::cpu);
CHECK_EQ(d.item<float>(), doctest::Approx(1.0f));
}
// Non-square should throw
CHECK_THROWS(
det(array({1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f}, {2, 3}), Device::cpu));
// 1D should throw
CHECK_THROWS(det(array({1.0f, 2.0f}), Device::cpu));
}
TEST_CASE("test slogdet") {
// 2x2: det = -2, so sign = -1, logabsdet = log(2)
{
array a = array({1.0f, 2.0f, 3.0f, 4.0f}, {2, 2});
auto [s, logabs] = slogdet(a, Device::cpu);
CHECK_EQ(s.item<float>(), doctest::Approx(-1.0f));
CHECK_EQ(logabs.item<float>(), doctest::Approx(std::log(2.0f)));
}
// Identity: sign = 1, logabsdet = 0
{
array a = eye(4);
auto [s, logabs] = slogdet(a, Device::cpu);
CHECK_EQ(s.item<float>(), doctest::Approx(1.0f));
CHECK_EQ(logabs.item<float>(), doctest::Approx(0.0f));
}
// Singular: sign = 0, logabsdet = -inf
{
array a = array({1.0f, 2.0f, 2.0f, 4.0f}, {2, 2});
auto [s, logabs] = slogdet(a, Device::cpu);
CHECK_EQ(s.item<float>(), 0.0f);
CHECK(std::isinf(logabs.item<float>()));
CHECK(logabs.item<float>() < 0);
}
}