feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,323 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package linalg
|
||||
|
||||
import (
|
||||
"math"
|
||||
"math/cmplx"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// approxEqual reports whether two real arrays agree elementwise within
|
||||
// an absolute and a relative tolerance.
|
||||
func approxEqual(t *testing.T, got, want *core.Array, atol, rtol float64, label string) {
|
||||
t.Helper()
|
||||
if got.NDim() != want.NDim() {
|
||||
t.Fatalf("%s: rank %d, want %d", label, got.NDim(), want.NDim())
|
||||
}
|
||||
for d := range got.NDim() {
|
||||
if got.Shape()[d] != want.Shape()[d] {
|
||||
t.Fatalf("%s: shape %s, want %s", label, base.ShapeText(got.Shape()), base.ShapeText(want.Shape()))
|
||||
}
|
||||
}
|
||||
for i := range got.Len() {
|
||||
g, w := got.FloatAt(i), want.FloatAt(i)
|
||||
diff := math.Abs(g - w)
|
||||
if diff > atol+rtol*math.Abs(w) {
|
||||
t.Fatalf("%s: element %d = %.12g, want %.12g (|diff| %.3g)",
|
||||
label, i, g, w, diff)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestMatrixExpClosedForms checks exp(A) against forms with a known
|
||||
// analytic answer: the identity, a nilpotent matrix whose series
|
||||
// terminates, a diagonal matrix, and the rotation generator whose
|
||||
// exponential is a rotation.
|
||||
func TestMatrixExpClosedForms(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
vals []float64
|
||||
n int
|
||||
want []float64
|
||||
}{
|
||||
{
|
||||
name: "zero_gives_identity",
|
||||
vals: []float64{0, 0, 0, 0},
|
||||
n: 2,
|
||||
want: []float64{1, 0, 0, 1},
|
||||
},
|
||||
{
|
||||
name: "scalar_multiple_of_identity",
|
||||
vals: []float64{2, 0, 0, 2},
|
||||
n: 2,
|
||||
want: []float64{math.E * math.E, 0, 0, math.E * math.E},
|
||||
},
|
||||
{
|
||||
name: "nilpotent_series_terminates",
|
||||
// A = [[0,1],[0,0]] squares to zero, so exp(A) = I + A.
|
||||
vals: []float64{0, 1, 0, 0},
|
||||
n: 2,
|
||||
want: []float64{1, 1, 0, 1},
|
||||
},
|
||||
{
|
||||
name: "diagonal",
|
||||
vals: []float64{1, 0, 0, 0, 2, 0, 0, 0, 3},
|
||||
n: 3,
|
||||
want: []float64{math.Exp(1), 0, 0, 0, math.Exp(2), 0, 0, 0, math.Exp(3)},
|
||||
},
|
||||
{
|
||||
name: "jordan_block",
|
||||
// A = [[1,1],[0,1]] gives exp(A) = e·[[1,1],[0,1]].
|
||||
vals: []float64{1, 1, 0, 1},
|
||||
n: 2,
|
||||
want: []float64{math.E, math.E, 0, math.E},
|
||||
},
|
||||
{
|
||||
name: "rotation_generator",
|
||||
// exp([[0,-t],[t,0]]) = [[cos t, -sin t],[sin t, cos t]].
|
||||
vals: []float64{0, -0.7, 0.7, 0},
|
||||
n: 2,
|
||||
want: []float64{math.Cos(0.7), -math.Sin(0.7), math.Sin(0.7), math.Cos(0.7)},
|
||||
},
|
||||
{
|
||||
name: "large_norm_needs_squarings",
|
||||
// A norm above theta_13 forces the scaling path.
|
||||
vals: []float64{30, 0, 0, -30},
|
||||
n: 2,
|
||||
want: []float64{math.Exp(30), 0, 0, math.Exp(-30)},
|
||||
},
|
||||
}
|
||||
for _, tt := range cases {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
a, err := core.FromFloats(tt.vals, tt.n, tt.n)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
got, err := MatrixExp(a)
|
||||
if err != nil {
|
||||
t.Fatalf("MatrixExp: %v", err)
|
||||
}
|
||||
want, err := core.FromFloats(tt.want, tt.n, tt.n)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats want: %v", err)
|
||||
}
|
||||
approxEqual(t, got, want, 1e-12, 1e-10, tt.name)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestMatrixExpSeries cross-checks exp(A) against the Taylor series
|
||||
// sum A^k/k! on a matrix small enough for the series to converge
|
||||
// tightly, covering a non-normal input the closed forms do not.
|
||||
func TestMatrixExpSeries(t *testing.T) {
|
||||
vals := []float64{
|
||||
0.5, 0.3, -0.1,
|
||||
0.2, -0.4, 0.6,
|
||||
0.1, 0.7, 0.2,
|
||||
}
|
||||
const n = 3
|
||||
a, err := core.FromFloats(vals, n, n)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
got, err := MatrixExp(a)
|
||||
if err != nil {
|
||||
t.Fatalf("MatrixExp: %v", err)
|
||||
}
|
||||
|
||||
// Series by repeated multiplication with the running term.
|
||||
term := make([]float64, n*n)
|
||||
for i := range n {
|
||||
term[i*n+i] = 1
|
||||
}
|
||||
sum := make([]float64, n*n)
|
||||
copy(sum, term)
|
||||
am := make([]float64, n*n)
|
||||
copy(am, vals)
|
||||
for k := 1; k <= 40; k++ {
|
||||
// term = term · A / k
|
||||
next := make([]float64, n*n)
|
||||
for i := range n {
|
||||
for p := range n {
|
||||
for j := range n {
|
||||
next[i*n+j] += term[i*n+p] * am[p*n+j]
|
||||
}
|
||||
}
|
||||
}
|
||||
for i := range n * n {
|
||||
term[i] = next[i] / float64(k)
|
||||
sum[i] += term[i]
|
||||
}
|
||||
}
|
||||
want := floatsToArray(sum, []int{n, n})
|
||||
approxEqual(t, got, want, 1e-13, 1e-12, "series")
|
||||
}
|
||||
|
||||
// TestMatrixExpGroupProperty checks the defining identity
|
||||
// exp(A)·exp(-A) = I, which exercises both the solve and the squaring
|
||||
// paths on a matrix with mixed-sign, non-symmetric entries.
|
||||
func TestMatrixExpGroupProperty(t *testing.T) {
|
||||
vals := []float64{
|
||||
1.2, -0.7, 0.4,
|
||||
0.9, 0.3, -1.1,
|
||||
-0.5, 0.8, 0.6,
|
||||
}
|
||||
const n = 3
|
||||
a, err := core.FromFloats(vals, n, n)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
neg := core.MulF(a, -1)
|
||||
ea, err := MatrixExp(a)
|
||||
if err != nil {
|
||||
t.Fatalf("MatrixExp(a): %v", err)
|
||||
}
|
||||
ena, err := MatrixExp(neg)
|
||||
if err != nil {
|
||||
t.Fatalf("MatrixExp(-a): %v", err)
|
||||
}
|
||||
prod, err := core.MatMul2D(ea, ena)
|
||||
if err != nil {
|
||||
t.Fatalf("MatMul2D: %v", err)
|
||||
}
|
||||
want, err := core.Zeros(core.Float, n, n)
|
||||
if err != nil {
|
||||
t.Fatalf("Zeros: %v", err)
|
||||
}
|
||||
for i := range n {
|
||||
want.RawFloats()[i*n+i] = 1
|
||||
}
|
||||
approxEqual(t, prod, want, 1e-10, 1e-9, "exp(A)exp(-A)")
|
||||
}
|
||||
|
||||
// TestMatrixExpComplexHermitian checks the complex path against a
|
||||
// closed form: a 2×2 Hermitian A = m·I + B splits into a scalar part
|
||||
// and a traceless B with B² = r²·I, so exp(A) = e^m·(cosh r·I +
|
||||
// (sinh r)/r·B). The real-valued tests cannot reach this path.
|
||||
func TestMatrixExpComplexHermitian(t *testing.T) {
|
||||
// A = [[2, 1-i],[1+i, 3]].
|
||||
const (
|
||||
a11 = 2.0
|
||||
a22 = 3.0
|
||||
n = 2
|
||||
)
|
||||
bOff := complex(1, 1)
|
||||
in := []complex128{
|
||||
a11, cmplx.Conj(bOff),
|
||||
bOff, a22,
|
||||
}
|
||||
a, err := core.FromComplexes(in, n, n)
|
||||
if err != nil {
|
||||
t.Fatalf("FromComplexes: %v", err)
|
||||
}
|
||||
got, err := MatrixExp(a)
|
||||
if err != nil {
|
||||
t.Fatalf("MatrixExp: %v", err)
|
||||
}
|
||||
|
||||
mid := (a11 + a22) / 2
|
||||
r := math.Hypot(a11-mid, cmplx.Abs(bOff))
|
||||
em := math.Exp(mid)
|
||||
ch := math.Cosh(r)
|
||||
sh := math.Sinh(r) / r
|
||||
want := make([]complex128, n*n)
|
||||
for i := range n {
|
||||
for j := range n {
|
||||
bij := a.ComplexAt(i*n + j)
|
||||
if i == j {
|
||||
bij -= complex(mid, 0)
|
||||
}
|
||||
idn := complex(0, 0)
|
||||
if i == j {
|
||||
idn = complex(1, 0)
|
||||
}
|
||||
want[i*n+j] = complex(em, 0) * (complex(ch, 0)*idn + complex(sh, 0)*bij)
|
||||
}
|
||||
}
|
||||
|
||||
for i := range n * n {
|
||||
g, w := got.ComplexAt(i), want[i]
|
||||
if diff := cmplx.Abs(g - w); diff > 1e-10*(1+cmplx.Abs(w)) {
|
||||
t.Fatalf("element %d = %v, want %v (|diff| %.3g)", i, g, w, diff)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestMatrixExpRejectsInvalid pins the error contract: non-square and
|
||||
// zero-sized inputs are refused rather than panicking.
|
||||
func TestMatrixExpRejectsInvalid(t *testing.T) {
|
||||
tall, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
if _, err := MatrixExp(tall); err == nil {
|
||||
t.Fatal("expected an error for a non-square matrix")
|
||||
}
|
||||
empty, err := core.Zeros(core.Float, 0, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Zeros: %v", err)
|
||||
}
|
||||
if _, err := MatrixExp(empty); err == nil {
|
||||
t.Fatal("expected an error for a zero-sized matrix")
|
||||
}
|
||||
}
|
||||
|
||||
// TestMatrixExpDtypes checks that int and float32 inputs promote to
|
||||
// float64 results, matching the promotion the rest of the linalg
|
||||
// surface applies.
|
||||
func TestMatrixExpDtypes(t *testing.T) {
|
||||
ints, err := core.FromInts([]int64{1, 0, 0, 1}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromInts: %v", err)
|
||||
}
|
||||
gotInt, err := MatrixExp(ints)
|
||||
if err != nil {
|
||||
t.Fatalf("MatrixExp(int): %v", err)
|
||||
}
|
||||
if gotInt.Dtype() != core.Float {
|
||||
t.Fatalf("int input gave dtype %s, want %s", gotInt.Dtype(), core.Float)
|
||||
}
|
||||
f32, err := core.FromFloat32s([]float32{0, 0, 0, 0}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat32s: %v", err)
|
||||
}
|
||||
gotF32, err := MatrixExp(f32)
|
||||
if err != nil {
|
||||
t.Fatalf("MatrixExp(float32): %v", err)
|
||||
}
|
||||
if gotF32.Dtype() != core.Float {
|
||||
t.Fatalf("float32 input gave dtype %s, want %s", gotF32.Dtype(), core.Float)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPadeDegreeSelection pins the degree and squaring choice across
|
||||
// the whole theta ladder, including the scaled branch.
|
||||
func TestPadeDegreeSelection(t *testing.T) {
|
||||
cases := []struct {
|
||||
nA float64
|
||||
deg int
|
||||
sq int
|
||||
}{
|
||||
{0, 3, 0},
|
||||
{1e-3, 3, 0},
|
||||
{1e-1, 5, 0},
|
||||
{0.5, 7, 0},
|
||||
{1.0, 9, 0},
|
||||
{2.0, 9, 0},
|
||||
{2.097847961257068, 9, 0},
|
||||
{2.097847961257069, 13, 0},
|
||||
{5.371920351148152, 13, 0},
|
||||
{10.743840702296304, 13, 1},
|
||||
{1000, 13, 8},
|
||||
}
|
||||
for _, tt := range cases {
|
||||
deg, sq := padeDegree(tt.nA)
|
||||
if deg != tt.deg || sq != tt.sq {
|
||||
t.Fatalf("padeDegree(%g) = (%d, %d), want (%d, %d)", tt.nA, deg, sq, tt.deg, tt.sq)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user