188 lines
5.1 KiB
Go
188 lines
5.1 KiB
Go
// 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])
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|