Files
tensor/linalg/decomp3_test.go

346 lines
10 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"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
func mustComplexes(t *testing.T, vals []complex128, shape ...int) *core.Array {
t.Helper()
a, err := core.FromComplexes(vals, shape...)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
return a
}
// matmulComplex multiplies flat complex matrices a (m×k) by b (k×n).
func matmulComplex(a, b []complex128, m, k, n int) []complex128 {
out := make([]complex128, m*n)
for i := range m {
for p := range k {
aip := a[i*k+p]
for j := range n {
out[i*n+j] += aip * b[p*n+j]
}
}
}
return out
}
// flatNorm returns the Frobenius norm of a flat complex matrix.
func flatNorm(a []complex128) float64 {
s := 0.0
for _, z := range a {
s += real(z)*real(z) + imag(z)*imag(z)
}
return math.Sqrt(s)
}
// TestEigenComplexPauli checks the Hermitian solver on matrices whose
// spectrum is known exactly: a·I + b·σx + c·σy + d·σz has eigenvalues
// a ± √(b²+c²+d²).
func TestEigenComplexPauli(t *testing.T) {
// [[2, 1−i],[1+i, 3]] = 2.5·I + 1·σx + 1·σy + 0.5·σz, so the
// eigenvalues are 2.5 ± 1.5.
h := mustComplexes(t, []complex128{
2, complex(1, -1),
complex(1, 1), 3,
}, 2, 2)
values, vectors, err := EigenComplex(h)
if err != nil {
t.Fatalf("EigenComplex: %v", err)
}
want := []float64{1, 4}
for i := range 2 {
if math.Abs(values.FloatAt(i)-want[i]) > 1e-12 {
t.Fatalf("eigenvalue[%d] = %v, want %v", i, values.FloatAt(i), want[i])
}
}
// Eigenvector residuals ‖H·v − λ·v‖, column by column.
hFlat := []complex128{2, complex(1, -1), complex(1, 1), 3}
for j := range 2 {
col := []complex128{vectors.ComplexAt(0*2 + j), vectors.ComplexAt(1*2 + j)}
hv := matmulComplex(hFlat, col, 2, 2, 1)
for i := range 2 {
res := hv[i] - complex(values.FloatAt(j), 0)*col[i]
if cmplx.Abs(res) > 1e-12 {
t.Fatalf("residual ‖Hv−λv‖[%d,%d] = %v", i, j, cmplx.Abs(res))
}
}
}
// Unitarity: Vᴴ·V = I.
for i := range 2 {
for j := range 2 {
s := complex(0, 0)
for k := range 2 {
s += cmplx.Conj(vectors.ComplexAt(k*2+i)) * vectors.ComplexAt(k*2+j)
}
want := 0.0
if i == j {
want = 1
}
if cmplx.Abs(s-complex(want, 0)) > 1e-12 {
t.Fatalf("VᴴV[%d,%d] = %v, want %v", i, j, s, want)
}
}
}
}
// TestEigenComplexLarger runs the solver on a 5×5 Hermitian matrix and
// verifies every Ritz pair by residual and every column by
// orthogonality, the properties users actually consume.
func TestEigenComplexLarger(t *testing.T) {
const n = 5
// H = B + Bᴴ for a pseudorandom complex B, Hermitian by
// construction.
flat := make([]complex128, n*n)
seed := uint64(88172645463325252)
next := func() complex128 {
seed ^= seed << 13
seed ^= seed >> 7
seed ^= seed << 17
return complex(float64(int64(seed%2000)-1000)/1000, float64(int64(seed%2000)-1000)/1000)
}
for i := range n * n {
flat[i] = next()
}
h := make([]complex128, n*n)
for i := range n {
for j := range n {
h[i*n+j] = flat[i*n+j] + cmplx.Conj(flat[j*n+i])
}
}
hArr := mustComplexes(t, h, n, n)
values, vectors, err := EigenComplex(hArr)
if err != nil {
t.Fatalf("EigenComplex: %v", err)
}
// Ascending order.
for i := 1; i < n; i++ {
if values.FloatAt(i) < values.FloatAt(i-1) {
t.Fatalf("eigenvalues not ascending: %v then %v", values.FloatAt(i-1), values.FloatAt(i))
}
}
for j := range n {
// Residual column: H·v_j − λ_j·v_j.
col := make([]complex128, n)
for i := range n {
col[i] = vectors.ComplexAt(i*n + j)
}
hv := matmulComplex(h, col, n, n, 1)
for i := range n {
res := hv[i] - complex(values.FloatAt(j), 0)*col[i]
if cmplx.Abs(res) > 1e-10*(1+math.Abs(values.FloatAt(j))) {
t.Fatalf("residual [%d,%d] = %v", i, j, cmplx.Abs(res))
}
}
}
}
// TestEigenComplexRejectsInvalid pins the input contract.
func TestEigenComplexRejectsInvalid(t *testing.T) {
real := mustFloats(t, []float64{1, 0, 0, 1}, 2, 2)
if _, _, err := EigenComplex(real); err == nil {
t.Fatal("expected an error for a real input")
}
nonsq := mustComplexes(t, []complex128{1, 0, 0, 1, 0, 0}, 2, 3)
if _, _, err := EigenComplex(nonsq); err == nil {
t.Fatal("expected an error for a non-square matrix")
}
asym := mustComplexes(t, []complex128{1, 2, 0, 1}, 2, 2)
if _, _, err := EigenComplex(asym); err == nil {
t.Fatal("expected an error for a non-Hermitian matrix")
}
}
// TestSVDComplexKnown checks the decomposition on a rank-1 outer
// product with exactly known singular values, plus the reconstruction,
// orthogonality and ordering contracts on tall and wide inputs.
func TestSVDComplexKnown(t *testing.T) {
// A = u·vᵀ with ‖u‖=1, ‖v‖=√2, so σ = {√2, 0}.
rt2 := 1 / math.Sqrt2
u := []complex128{complex(rt2, 0), complex(0, rt2)}
v := []float64{1, 1}
a := make([]complex128, 4)
for i := range 2 {
for j := range 2 {
a[i*2+j] = u[i] * complex(v[j], 0)
}
}
uOut, sigma, vH, err := SVDComplex(mustComplexes(t, a, 2, 2))
if err != nil {
t.Fatalf("SVDComplex: %v", err)
}
if math.Abs(sigma.FloatAt(0)-math.Sqrt2) > 1e-12 {
t.Fatalf("σ₁ = %v, want √2", sigma.FloatAt(0))
}
if sigma.FloatAt(1) > 1e-12 {
t.Fatalf("σ₂ = %v, want 0", sigma.FloatAt(1))
}
// Reconstruction A ≈ U·Σ·Vᴴ.
recon := make([]complex128, 4)
for i := range 2 {
for j := range 2 {
s := complex(0, 0)
for k := range 2 {
s += uOut.ComplexAt(i*2+k) * complex(sigma.FloatAt(k), 0) *
vH.ComplexAt(k*2+j)
}
recon[i*2+j] = s
}
}
if d := flatNorm(subComplex(a, recon)); d > 1e-12 {
t.Fatalf("reconstruction error %v", d)
}
}
// TestSVDComplexTallAndWide checks the tall and the wide path on the
// same content: reconstruction, orthogonality of both factors and
// descending singular values.
func TestSVDComplexTallAndWide(t *testing.T) {
build := func(m, n int) *core.Array {
flat := make([]complex128, m*n)
seed := uint64(11400714819323198485)
for i := range m * n {
seed ^= seed << 13
seed ^= seed >> 7
seed ^= seed << 17
flat[i] = complex(float64(int64(seed%400)-200)/100, float64(int64(seed%400)-200)/100)
}
return mustComplexes(t, flat, m, n)
}
check := func(t *testing.T, a *core.Array) {
m, n := a.Shape()[0], a.Shape()[1]
u, sigma, vH, err := SVDComplex(a)
if err != nil {
t.Fatalf("SVDComplex(%dx%d): %v", m, n, err)
}
// Shapes mirror the real SVD: U (m, min), Σ (min,), Vᴴ (m, n).
rank := min(m, n)
if u.Shape()[0] != m || u.Shape()[1] != rank {
t.Fatalf("U shape %s, want [%d %d]", base.ShapeText(u.Shape()), m, rank)
}
if sigma.Len() != rank || vH.Shape()[0] != rank || vH.Shape()[1] != n {
t.Fatalf("sigma len %d, Vᴴ shape %s", sigma.Len(), base.ShapeText(vH.Shape()))
}
for i := 1; i < rank; i++ {
if sigma.FloatAt(i) > sigma.FloatAt(i-1)+1e-12 {
t.Fatalf("singular values not descending: %v then %v", sigma.FloatAt(i-1), sigma.FloatAt(i))
}
}
// U·Σ·Vᴴ.
recon := make([]complex128, m*n)
for i := range m {
for j := range n {
s := complex(0, 0)
for k := range rank {
s += u.ComplexAt(i*rank+k) * complex(sigma.FloatAt(k), 0) *
vH.ComplexAt(k*n+j)
}
recon[i*n+j] = s
}
}
aFlat := make([]complex128, m*n)
for i := range m * n {
aFlat[i] = a.ComplexAt(i)
}
if d := flatNorm(subComplex(aFlat, recon)); d > 1e-9*float64(m) {
t.Fatalf("%dx%d reconstruction error %v", m, n, d)
}
// Orthogonality of U's columns and of Vᴴᴴ (i.e. VᴴV).
for i := range rank {
for j := range rank {
su := complex(0, 0)
for k := range m {
su += cmplx.Conj(u.ComplexAt(k*rank+i)) * u.ComplexAt(k*rank+j)
}
// Vᴴ has orthonormal ROWS in every convention.
sv2 := complex(0, 0)
for k := range n {
sv2 += vH.ComplexAt(i*n+k) * cmplx.Conj(vH.ComplexAt(j*n+k))
}
want := 0.0
if i == j {
want = 1
}
if cmplx.Abs(su-complex(want, 0)) > 1e-9 {
t.Fatalf("UᴴU[%d,%d] = %v", i, j, su)
}
if cmplx.Abs(sv2-complex(want, 0)) > 1e-9 {
t.Fatalf("(VᴴVᴴ*)[%d,%d] = %v", i, j, sv2)
}
}
}
}
t.Run("tall", func(t *testing.T) { check(t, build(6, 4)) })
t.Run("wide", func(t *testing.T) { check(t, build(4, 6)) })
}
// subComplex subtracts two flat complex matrices of equal length.
func subComplex(a, b []complex128) []complex128 {
out := make([]complex128, len(a))
for i := range a {
out[i] = a[i] - b[i]
}
return out
}
// TestEigenComplexNearDiagonalConverges pins the Jacobi sweep against a
// matrix whose every off-diagonal entry sits just under the per-entry
// skip level while the aggregate off-norm stays above the convergence
// threshold: the skip must not freeze the sweep above its own
// convergence test, which used to exhaust the passes and report the
// nearly diagonal matrix as unconverged.
func TestEigenComplexNearDiagonalConverges(t *testing.T) {
const n = 30
cv := make([]complex128, n*n)
for i := range n {
cv[i*n+i] = 1
}
for i := range n {
for j := i + 1; j < n; j++ {
cv[i*n+j] = complex(0.9e-13, 0)
cv[j*n+i] = complex(0.9e-13, 0)
}
}
a, err := core.FromComplexes(cv, n, n)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
vals, _, err := EigenComplex(a)
if err != nil {
t.Fatalf("EigenComplex: %v", err)
}
for i := range n {
if math.Abs(vals.FloatAt(i)-1) > 1e-11 {
t.Fatalf("eigenvalue %d is %.12g, want 1", i, vals.FloatAt(i))
}
}
}
// TestEigenComplexRefusesNonFinite pins that a poisoned matrix never
// reads as Hermitian: the mirror comparison cannot see a NaN
// difference, so the entry is refused outright, the way the sparse
// sibling's Hermitian check refuses it.
func TestEigenComplexRefusesNonFinite(t *testing.T) {
cv := []complex128{complex(math.NaN(), 0), 0, 0, 1}
a := mustComplexes(t, cv, 2, 2)
if _, _, err := EigenComplex(a); err == nil || !strings.Contains(err.Error(), "not finite") {
t.Fatalf("EigenComplex(NaN): %v", err)
}
cv2 := []complex128{complex(0, math.Inf(1)), 0, 0, 1}
b := mustComplexes(t, cv2, 2, 2)
if _, _, err := EigenComplex(b); err == nil || !strings.Contains(err.Error(), "not finite") {
t.Fatalf("EigenComplex(Inf): %v", err)
}
}