Files
tensor/linalg/linalg_test.go
T

188 lines
5.1 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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])
}
}
}