Files
tensor/linalg/decomp3_test.go
T

294 lines
8.4 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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"
)
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
}