Files
tensor/linalg/decomp3_test.go
T
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

294 lines
8.4 KiB
Go
Raw 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"
)
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
}