Files
tensor/linalg/csvd_test.go
T

192 lines
5.5 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/core"
"testing"
)
// csvdSample builds a deterministic complex m×n matrix.
func csvdSample(m, n int) *core.Array {
vals := make([]complex128, m*n)
for i := range m * n {
vals[i] = complex(math.Sin(float64(3*i+1)), math.Cos(float64(2*i+1)))
}
a, _ := core.FromComplexes(vals, m, n)
return a
}
// csvdPlane returns the m×m unitary that rotates coordinates i, j by
// angle theta with phase phi, the building block for test unitaries
// with known spectra.
func csvdPlane(m, i, j int, theta, phi float64) []complex128 {
q := make([]complex128, m*m)
for k := range m {
q[k*m+k] = 1
}
q[i*m+i] = complex(math.Cos(theta), 0)
q[j*m+j] = complex(math.Cos(theta), 0)
q[i*m+j] = complex(math.Sin(theta)*math.Cos(phi), math.Sin(theta)*math.Sin(phi))
q[j*m+i] = complex(-math.Sin(theta)*math.Cos(phi), math.Sin(theta)*math.Sin(phi))
return q
}
// TestSVDComplexContracts pins the decomposition contract on a spread
// of deterministic shapes: reconstruction, unitarity of both factors,
// descending non-negative singular values.
func TestSVDComplexContracts(t *testing.T) {
cases := []struct{ m, n int }{{5, 3}, {4, 4}, {3, 1}, {2, 2}, {1, 1}, {6, 2}}
for _, tc := range cases {
a := csvdSample(tc.m, tc.n)
u, sigma, vh, err := SVDComplex(a)
if err != nil {
t.Fatalf("%dx%d: SVDComplex: %v", tc.m, tc.n, err)
}
r := min(tc.m, tc.n)
scale := 0.0
for i := range a.Len() {
scale = math.Max(scale, cmplx.Abs(a.ComplexAt(i)))
}
// Reconstruction.
recon := 0.0
for i := range tc.m {
for j := range tc.n {
s := complex(0, 0)
for k := range r {
s += u.ComplexAt(i*r+k) * complex(sigma.FloatAt(k), 0) * vh.ComplexAt(k*tc.n+j)
}
recon = math.Max(recon, cmplx.Abs(s-a.ComplexAt(i*tc.n+j)))
}
}
if recon > 1e-10*math.Max(1, scale) {
t.Fatalf("%dx%d: reconstruction error %.3g", tc.m, tc.n, recon)
}
// Unitarity of U's columns and of V.
gram := func(get func(i, j int) complex128, rows, cols int) float64 {
worst := 0.0
for i := range cols {
for j := range cols {
s := complex(0, 0)
for l := range rows {
s += cmplx.Conj(get(l, i)) * get(l, j)
}
want := 0.0
if i == j {
want = 1
}
worst = math.Max(worst, math.Abs(cmplx.Abs(s)-want))
}
}
return worst
}
if g := gram(func(i, j int) complex128 { return u.ComplexAt(i*r + j) }, tc.m, r); g > 1e-10 {
t.Fatalf("%dx%d: U not orthonormal, error %.3g", tc.m, tc.n, g)
}
if g := gram(func(i, j int) complex128 { return vh.ComplexAt(i*tc.n + j) }, tc.n, tc.n); g > 1e-10 {
t.Fatalf("%dx%d: Vᴴ not unitary, error %.3g", tc.m, tc.n, g)
}
for k := range sigma.Len() {
if sigma.FloatAt(k) < 0 {
t.Fatalf("%dx%d: negative singular value %g", tc.m, tc.n, sigma.FloatAt(k))
}
if k > 0 && sigma.FloatAt(k) > sigma.FloatAt(k-1)+1e-12 {
t.Fatalf("%dx%d: singular values not descending", tc.m, tc.n)
}
}
}
}
// TestSVDComplexIllConditioned is the reason the direct route exists:
// with a known spectrum 1, 1e-8, 1e-16 the small singular values keep
// their relative accuracy, where the squared-condition AᴴA route would
// lose half the digits.
func TestSVDComplexIllConditioned(t *testing.T) {
const m, n = 3, 3
sigma := []float64{1, 1e-8, 1e-16}
// A = P·diag(σ)·Qᴴ for two deterministic complex unitaries P, Q.
p := csvdPlane(m, 0, 1, 0.7, 1.1)
q := csvdPlane(n, 1, 2, 1.3, 0.4)
_ = q
pq := csvdPlane(m, 1, 2, 0.5, 2.2)
// Compose P = p·pq.
pMat := make([]complex128, m*m)
for i := range m {
for j := range m {
s := complex(0, 0)
for k := range m {
s += p[i*m+k] * pq[k*m+j]
}
pMat[i*m+j] = s
}
}
vals := make([]complex128, m*n)
for i := range m {
for j := range n {
s := complex(0, 0)
for k := range m {
s += pMat[i*m+k] * complex(sigma[k], 0) * cmplx.Conj(q[j*n+k])
}
vals[i*n+j] = s
}
}
a, err := core.FromComplexes(vals, m, n)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
_, sigmaOut, _, err := SVDComplex(a)
if err != nil {
t.Fatalf("SVDComplex: %v", err)
}
for k := range n {
want := sigma[k]
got := sigmaOut.FloatAt(k)
if k < 2 {
if rel := math.Abs(got-want) / want; rel > 1e-9 {
t.Fatalf("σ%d = %.17g, want %.17g (relative error %.3g)", k, got, want, rel)
}
} else {
// At the round-off floor the honest guarantee is absolute:
// the direct route pins σ to eps·σ_max, the squared route
// could not.
if math.Abs(got-want) > 1e-15*sigma[0] {
t.Fatalf("σ%d = %.17g, want %.17g (absolute error %.3g)",
k, got, want, math.Abs(got-want))
}
}
}
}
// TestSVDComplexMatchesReal cross-checks the complex solver against
// the independent real SVD on a real matrix embedded in complex.
func TestSVDComplexMatchesReal(t *testing.T) {
realPart := mustFloats(t, []float64{
3, 0, 1,
1, 2, 1,
1, 1, 2,
0, 1, 4,
}, 4, 3)
uR, sigmaR, _, err := SVD(realPart)
if err != nil {
t.Fatalf("SVD: %v", err)
}
_ = uR
vals := make([]complex128, 12)
for i := range 12 {
vals[i] = complex(realPart.FloatAt(i), 0)
}
a, _ := core.FromComplexes(vals, 4, 3)
_, sigmaC, _, err := SVDComplex(a)
if err != nil {
t.Fatalf("SVDComplex: %v", err)
}
for k := range 3 {
if math.Abs(sigmaC.FloatAt(k)-sigmaR.FloatAt(k)) > 1e-10*math.Max(1, sigmaR.FloatAt(k)) {
t.Fatalf("σ%d: complex %.12g, real %.12g", k, sigmaC.FloatAt(k), sigmaR.FloatAt(k))
}
}
}