Files
tensor/linalg/expm_test.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

324 lines
8.3 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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)
}
}
}