// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package linalg import ( "math" "sourcedock.dev/petrbalvin/tensor/internal/core" "strings" "testing" ) func approx(t *testing.T, name string, got, want float64) { t.Helper() if math.Abs(got-want) > 1e-9 { t.Fatalf("%s: got %v, want %v", name, got, want) } } func TestDet(t *testing.T) { m := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) d, err := Det(m) if err != nil { t.Fatalf("Det: %v", err) } approx(t, "Det", d, -2) id, _ := core.Identity(core.Float, 3) d, err = Det(id) if err != nil { t.Fatalf("Det identity: %v", err) } approx(t, "Det identity", d, 1) // A singular matrix yields 0, not an error. singular := mustFromFloats(t, []float64{1, 2, 2, 4}, 2, 2) d, err = Det(singular) if err != nil { t.Fatalf("Det singular: %v", err) } approx(t, "Det singular", d, 0) // int matrices convert. im := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) d, err = Det(im) if err != nil { t.Fatalf("Det int: %v", err) } approx(t, "Det int", d, -2) nonSquare := mustFromFloats(t, []float64{1, 2, 3}, 1, 3) if _, err := Det(nonSquare); err == nil || !strings.Contains(err.Error(), "square 2-D matrix") { t.Fatalf("Det shape: %v", err) } } func TestSolve(t *testing.T) { // 3x + y = 9, x + 2y = 8 so x = 2, y = 3. a := mustFromFloats(t, []float64{3, 1, 1, 2}, 2, 2) b := mustFromFloats(t, []float64{9, 8}, 2) x, err := Solve(a, b) if err != nil { t.Fatalf("Solve: %v", err) } if x.Dtype() != core.Float || x.Shape()[0] != 2 { t.Fatalf("Solve shape: %s", x) } v0, _ := core.FloatAt(x, 0) v1, _ := core.FloatAt(x, 1) approx(t, "Solve x", v0, 2) approx(t, "Solve y", v1, 3) // Matrix right-hand side solves column by column: rows [9,1] and [8,2] // give the columns [9,8] and [1,2]. bm := mustFromFloats(t, []float64{9, 1, 8, 2}, 2, 2) xm, err := Solve(a, bm) if err != nil { t.Fatalf("Solve matrix: %v", err) } // Second column: 3x+y=1, x+2y=2 so x=0, y=1. c00, _ := core.FloatAt(xm, 0, 0) c01, _ := core.FloatAt(xm, 0, 1) c10, _ := core.FloatAt(xm, 1, 0) c11, _ := core.FloatAt(xm, 1, 1) approx(t, "Solve matrix 00", c00, 2) approx(t, "Solve matrix 01", c01, 0) approx(t, "Solve matrix 10", c10, 3) approx(t, "Solve matrix 11", c11, 1) // int operands convert on both sides. ia := mustFromInts(t, []int64{3, 1, 1, 2}, 2, 2) ib := mustFromInts(t, []int64{9, 8}, 2) ix, err := Solve(ia, ib) if err != nil || ix.Dtype() != core.Float { t.Fatalf("Solve int: %s %v", ix, err) } iv, _ := core.FloatAt(ix, 0) approx(t, "Solve int x", iv, 2) singular := mustFromFloats(t, []float64{1, 2, 2, 4}, 2, 2) if _, err := Solve(singular, b); err == nil || !strings.Contains(err.Error(), "singular") { t.Fatalf("Solve singular: %v", err) } wrong := mustFromFloats(t, []float64{1, 2, 3}, 3) if _, err := Solve(a, wrong); err == nil || !strings.Contains(err.Error(), "b must be") { t.Fatalf("Solve b shape: %v", err) } if _, err := Solve(wrong, a); err == nil || !strings.Contains(err.Error(), "square 2-D matrix") { t.Fatalf("Solve shape: %v", err) } } func TestInv(t *testing.T) { m := mustFromFloats(t, []float64{4, 7, 2, 6}, 2, 2) inv, err := Inv(m) if err != nil { t.Fatalf("Inv: %v", err) } // 1/10 · [[6, -7], [-2, 4]] v00, _ := core.FloatAt(inv, 0, 0) v01, _ := core.FloatAt(inv, 0, 1) v10, _ := core.FloatAt(inv, 1, 0) v11, _ := core.FloatAt(inv, 1, 1) approx(t, "Inv 00", v00, 0.6) approx(t, "Inv 01", v01, -0.7) approx(t, "Inv 10", v10, -0.2) approx(t, "Inv 11", v11, 0.4) // A·A⁻¹ is the identity. prod, err := core.MatMul2D(m, inv) if err != nil { t.Fatalf("Inv check: %v", err) } for i := range 2 { for j := range 2 { want := 0.0 if i == j { want = 1 } got, _ := core.FloatAt(prod, i, j) approx(t, "Inv product", got, want) } } // A permutation matrix exercises the pivot path. p := mustFromFloats(t, []float64{0, 1, 1, 0}, 2, 2) pinv, err := Inv(p) if err != nil { t.Fatalf("Inv permutation: %v", err) } g00, _ := core.FloatAt(pinv, 0, 0) g01, _ := core.FloatAt(pinv, 0, 1) approx(t, "Inv permutation 00", g00, 0) approx(t, "Inv permutation 01", g01, 1) singular := mustFromFloats(t, []float64{1, 2, 2, 4}, 2, 2) if _, err := Inv(singular); err == nil || !strings.Contains(err.Error(), "singular") { t.Fatalf("Inv singular: %v", err) } if _, err := Inv(mustFromFloats(t, []float64{1, 2, 3}, 1, 3)); err == nil { t.Fatalf("Inv shape must error") } } // TestSolveRealMatrixComplexRHS pins the promotion contract: a real // system with a complex right-hand side promotes the whole solve to // complex128 instead of erroring. func TestSolveRealMatrixComplexRHS(t *testing.T) { a := mustFromFloats(t, []float64{2, 0, 0, 4}, 2, 2) b, _ := core.FromComplexes([]complex128{2 + 4i, 8}, 2) x, err := Solve(a, b) if err != nil { t.Fatalf("Solve real x complex: %v", err) } if x.Dtype() != core.Complex { t.Fatalf("Solve promote dtype: %s", x.Dtype()) } // A is diagonal: x = b / diag = [1+2i, 2]. want := []complex128{1 + 2i, 2} for i := range want { if v, _ := core.ComplexAt(x, i); v != want[i] { t.Fatalf("Solve [%d]: got %v, want %v", i, v, want[i]) } } }