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))
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|