192 lines
5.5 KiB
Go
192 lines
5.5 KiB
Go
// 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))
|
||
}
|
||
}
|
||
}
|