324 lines
8.3 KiB
Go
324 lines
8.3 KiB
Go
// 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)
|
||
}
|
||
}
|
||
}
|