313 lines
9.2 KiB
Go
313 lines
9.2 KiB
Go
// 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
|
||
}
|
||
|
||
// 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)
|
||
}
|
||
}
|