308 lines
8.7 KiB
Go
308 lines
8.7 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/core"
|
||
"testing"
|
||
)
|
||
|
||
// eigenResidual returns ‖A·v_j − λ_j·v_j‖∞ for every eigenpair j of a
|
||
// square matrix a.
|
||
func eigenResidual(a *core.Array, values, vectors *core.Array) []float64 {
|
||
n := a.Shape()[0]
|
||
out := make([]float64, n)
|
||
for j := range n {
|
||
lam := values.ComplexAt(j)
|
||
worst := 0.0
|
||
for i := range n {
|
||
acc := complex(0, 0)
|
||
for k := range n {
|
||
acc += a.ComplexAt(i*n+k) * vectors.ComplexAt(k*n+j)
|
||
}
|
||
acc -= lam * vectors.ComplexAt(i*n+j)
|
||
if m := absComplex(acc); m > worst {
|
||
worst = m
|
||
}
|
||
}
|
||
out[j] = worst
|
||
}
|
||
return out
|
||
}
|
||
|
||
// checkEigenPairs asserts every eigenpair satisfies A·v = λ·v to
|
||
// rounding level and every vector has unit norm.
|
||
func checkEigenPairs(t *testing.T, a *core.Array, values, vectors *core.Array, tol float64) {
|
||
t.Helper()
|
||
n := a.Shape()[0]
|
||
for _, r := range eigenResidual(a, values, vectors) {
|
||
if r > tol {
|
||
t.Fatalf("eigenpair residual %g, want ≤ %g", r, tol)
|
||
}
|
||
}
|
||
for j := range n {
|
||
norm := 0.0
|
||
for i := range n {
|
||
z := vectors.ComplexAt(i*n + j)
|
||
norm += real(z)*real(z) + imag(z)*imag(z)
|
||
}
|
||
if math.Abs(norm-1) > 1e-12 {
|
||
t.Fatalf("vector %d has norm² %g, want 1", j, norm)
|
||
}
|
||
}
|
||
}
|
||
|
||
// matchComplexSpectrum asserts every expected value appears in the
|
||
// computed spectrum and the magnitudes come out descending; within a
|
||
// magnitude tie (conjugate pairs, ±λ) the computed values are equal
|
||
// only to rounding, so the order inside the tie is not asserted.
|
||
func matchComplexSpectrum(t *testing.T, values *core.Array, want []complex128, tol float64) {
|
||
t.Helper()
|
||
used := make([]bool, len(want))
|
||
for i := range values.Len() {
|
||
got := values.ComplexAt(i)
|
||
found := false
|
||
for k, w := range want {
|
||
if !used[k] && absComplex(got-w) <= tol {
|
||
used[k] = true
|
||
found = true
|
||
break
|
||
}
|
||
}
|
||
if !found {
|
||
t.Fatalf("value[%d] = %v has no expected match within %g", i, got, tol)
|
||
}
|
||
if i > 0 && absComplex(values.ComplexAt(i-1)) < absComplex(got) {
|
||
t.Fatalf("magnitude order broken at %d: |%v| < |%v|", i, values.ComplexAt(i-1), got)
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestEigenGeneralPauli checks the σx matrix: the eigenvalues are ±1.
|
||
func TestEigenGeneralPauli(t *testing.T) {
|
||
a := mustFloats(t, []float64{0, 1, 1, 0}, 2, 2)
|
||
values, vectors, err := EigenGeneral(a)
|
||
if err != nil {
|
||
t.Fatalf("EigenGeneral: %v", err)
|
||
}
|
||
matchComplexSpectrum(t, values, []complex128{1, -1}, 1e-12)
|
||
checkEigenPairs(t, a, values, vectors, 1e-14)
|
||
}
|
||
|
||
// TestEigenGeneralRotation checks the rotation generator [[0, −θ],
|
||
// [θ, 0]]: the eigenvalues are the purely imaginary pair ±iθ.
|
||
func TestEigenGeneralRotation(t *testing.T) {
|
||
const theta = 0.7
|
||
a := mustFloats(t, []float64{0, -theta, theta, 0}, 2, 2)
|
||
values, vectors, err := EigenGeneral(a)
|
||
if err != nil {
|
||
t.Fatalf("EigenGeneral: %v", err)
|
||
}
|
||
matchComplexSpectrum(t, values, []complex128{complex(0, theta), complex(0, -theta)}, 1e-12)
|
||
checkEigenPairs(t, a, values, vectors, 1e-14)
|
||
}
|
||
|
||
// TestEigenGeneralKnownSpectrum checks a 6×6 real nonsymmetric matrix
|
||
// assembled as V·D·V⁻¹, whose spectrum is known by construction:
|
||
// 5, −3, the conjugate pair 2 ± 1.5i and the conjugate pair
|
||
// 0.5 ± 4i. The expected order follows the descending-magnitude sort
|
||
// with the documented tiebreaks: 5, 0.5+4i, 0.5−4i, −3, 2+1.5i,
|
||
// 2−1.5i. The trace and determinant cross-check the spectrum through
|
||
// two independent invariants.
|
||
func TestEigenGeneralKnownSpectrum(t *testing.T) {
|
||
d := mustFloats(t, []float64{
|
||
5, 0, 0, 0, 0, 0,
|
||
0, -3, 0, 0, 0, 0,
|
||
0, 0, 2, -1.5, 0, 0,
|
||
0, 0, 1.5, 2, 0, 0,
|
||
0, 0, 0, 0, 0.5, -4,
|
||
0, 0, 0, 0, 4, 0.5,
|
||
}, 6, 6)
|
||
v := mustFloats(t, []float64{
|
||
1, 0.2, -0.2, 0, 0.2, 0,
|
||
0.2, 1, 0, 0.2, -0.2, 0.2,
|
||
-0.2, 0, 1, 0.2, 0, -0.2,
|
||
0, 0.2, 0.2, 1, -0.2, 0,
|
||
0.2, -0.2, 0, -0.2, 1, 0.2,
|
||
0, 0.2, -0.2, 0, 0.2, 1,
|
||
}, 6, 6)
|
||
vInv, err := Inv(v)
|
||
if err != nil {
|
||
t.Fatalf("Inv: %v", err)
|
||
}
|
||
dv, err := core.MatMul2D(d, vInv)
|
||
if err != nil {
|
||
t.Fatalf("MatMul2D: %v", err)
|
||
}
|
||
a, err := core.MatMul2D(v, dv)
|
||
if err != nil {
|
||
t.Fatalf("MatMul2D: %v", err)
|
||
}
|
||
values, vectors, err := EigenGeneral(a)
|
||
if err != nil {
|
||
t.Fatalf("EigenGeneral: %v", err)
|
||
}
|
||
want := []complex128{
|
||
5,
|
||
complex(0.5, 4), complex(0.5, -4),
|
||
-3,
|
||
complex(2, 1.5), complex(2, -1.5),
|
||
}
|
||
matchComplexSpectrum(t, values, want, 1e-8)
|
||
checkEigenPairs(t, a, values, vectors, 1e-9)
|
||
|
||
sum := complex(0, 0)
|
||
prod := complex(1, 0)
|
||
for i := range 6 {
|
||
sum += values.ComplexAt(i)
|
||
prod *= values.ComplexAt(i)
|
||
}
|
||
if absComplex(sum-complex(7, 0)) > 1e-8 {
|
||
t.Fatalf("Σλ = %v, want 7", sum)
|
||
}
|
||
det, err := Det(a)
|
||
if err != nil {
|
||
t.Fatalf("Det: %v", err)
|
||
}
|
||
// det = 5·(−3)·(2²+1.5²)·(0.5²+4²) = −1523.4375.
|
||
if absComplex(prod-complex(det, 0)) > 1e-6*math.Abs(det) {
|
||
t.Fatalf("Πλ = %v, det = %g", prod, det)
|
||
}
|
||
}
|
||
|
||
// TestEigenGeneralHermitian checks a complex input: [[2, i], [−i, 2]]
|
||
// has eigenvalues 1 and 3 (trace 4, determinant 3).
|
||
func TestEigenGeneralHermitian(t *testing.T) {
|
||
a := mustComplexes(t, []complex128{
|
||
2, complex(0, 1),
|
||
complex(0, -1), 2,
|
||
}, 2, 2)
|
||
values, vectors, err := EigenGeneral(a)
|
||
if err != nil {
|
||
t.Fatalf("EigenGeneral: %v", err)
|
||
}
|
||
want := []complex128{3, 1}
|
||
for i := range 2 {
|
||
if absComplex(values.ComplexAt(i)-want[i]) > 1e-12 {
|
||
t.Fatalf("value[%d] = %v, want %v", i, values.ComplexAt(i), want[i])
|
||
}
|
||
}
|
||
checkEigenPairs(t, a, values, vectors, 1e-14)
|
||
}
|
||
|
||
// TestEigenGeneralSymmetricCrossCheck pins the general path against
|
||
// the dedicated symmetric solver on the same matrix.
|
||
func TestEigenGeneralSymmetricCrossCheck(t *testing.T) {
|
||
vals := []float64{
|
||
4, 1, 0,
|
||
1, 3, 2,
|
||
0, 2, 5,
|
||
}
|
||
a := mustFloats(t, vals, 3, 3)
|
||
gen, _, err := EigenGeneral(a)
|
||
if err != nil {
|
||
t.Fatalf("EigenGeneral: %v", err)
|
||
}
|
||
sym, _, err := Eigen(a)
|
||
if err != nil {
|
||
t.Fatalf("Eigen: %v", err)
|
||
}
|
||
for i := range 3 {
|
||
best := math.MaxFloat64
|
||
for k := range 3 {
|
||
if d := absComplex(gen.ComplexAt(i) - complex(sym.FloatAt(k), 0)); d < best {
|
||
best = d
|
||
}
|
||
}
|
||
if best > 1e-12 {
|
||
t.Fatalf("general value[%d] = %v has no symmetric match within 1e-12", i, gen.ComplexAt(i))
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestEigenGeneralCyclicPermutation checks the matrix that made the
|
||
// shifted QR iteration famous: the cyclic permutation, where the
|
||
// spectrum is the full set of n-th roots of unity and the plain
|
||
// Wilkinson shift cycles until the exceptional shift breaks the
|
||
// symmetry.
|
||
func TestEigenGeneralCyclicPermutation(t *testing.T) {
|
||
const n = 6
|
||
vals := make([]float64, n*n)
|
||
for i := range n {
|
||
vals[i*n+(i+1)%n] = 1
|
||
}
|
||
a := mustFloats(t, vals, n, n)
|
||
values, vectors, err := EigenGeneral(a)
|
||
if err != nil {
|
||
t.Fatalf("EigenGeneral: %v", err)
|
||
}
|
||
want := make([]complex128, n)
|
||
for k := range n {
|
||
want[k] = cmplx.Exp(complex(0, 2*math.Pi*float64(k)/float64(n)))
|
||
}
|
||
matchComplexSpectrum(t, values, want, 1e-8)
|
||
checkEigenPairs(t, a, values, vectors, 1e-9)
|
||
}
|
||
|
||
// TestEigenGeneralSingleCell covers the trivial 1×1 case.
|
||
func TestEigenGeneralSingleCell(t *testing.T) {
|
||
a := mustFloats(t, []float64{2.5}, 1, 1)
|
||
values, vectors, err := EigenGeneral(a)
|
||
if err != nil {
|
||
t.Fatalf("EigenGeneral: %v", err)
|
||
}
|
||
if absComplex(values.ComplexAt(0)-complex(2.5, 0)) > 1e-15 {
|
||
t.Fatalf("value = %v, want 2.5", values.ComplexAt(0))
|
||
}
|
||
if absComplex(vectors.ComplexAt(0)-complex(1, 0)) > 1e-15 {
|
||
t.Fatalf("vector = %v, want 1", vectors.ComplexAt(0))
|
||
}
|
||
}
|
||
|
||
// TestEigenGeneralTinyRotation pins the purely relative deflation
|
||
// floors: a rotation scaled by 1e-20 keeps its complex conjugate
|
||
// spectrum instead of deflating to the real diagonal (the old
|
||
// max(1, scale) floor treated the whole matrix as rounding noise).
|
||
func TestEigenGeneralTinyRotation(t *testing.T) {
|
||
const scale = 1e-20
|
||
theta := math.Pi / 5
|
||
c, s := math.Cos(theta), math.Sin(theta)
|
||
a := mustFloats(t, []float64{
|
||
scale * c, -scale * s,
|
||
scale * s, scale * c,
|
||
}, 2, 2)
|
||
values, vectors, err := EigenGeneral(a)
|
||
if err != nil {
|
||
t.Fatalf("EigenGeneral: %v", err)
|
||
}
|
||
want := []complex128{
|
||
scale * cmplx.Exp(complex(0, theta)),
|
||
scale * cmplx.Exp(complex(0, -theta)),
|
||
}
|
||
for j := range 2 {
|
||
got := values.ComplexAt(j)
|
||
best := math.Inf(1)
|
||
for _, w := range want {
|
||
if m := absComplex(got - w); m < best {
|
||
best = m
|
||
}
|
||
}
|
||
if best > 1e-26 {
|
||
t.Fatalf("value[%d] = %v has no expected match within %g", j, got, 1e-26)
|
||
}
|
||
}
|
||
checkEigenPairs(t, a, values, vectors, 1e-25)
|
||
}
|
||
|
||
func TestEigenGeneralErrors(t *testing.T) {
|
||
if _, _, err := EigenGeneral(mustFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3)); err == nil {
|
||
t.Fatal("non-square matrix: want an error")
|
||
}
|
||
if _, _, err := EigenGeneral(mustFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 2, 2)); err == nil {
|
||
t.Fatal("3-D input: want an error")
|
||
}
|
||
}
|