feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,187 @@
|
||||
// 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])
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user