feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,293 @@
|
||||
// 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
|
||||
}
|
||||
Reference in New Issue
Block a user