294 lines
8.4 KiB
Go
294 lines
8.4 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/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
|
|||
|
|
}
|