Files
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

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