Files
tensor/linalg/svdsolve_test.go
T
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

181 lines
5.6 KiB
Go
Raw 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"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
// diagonalSpectrum returns the n×n diagonal matrix with the given
// singular values, the cleanest stage for checking what truncation
// and damping do per direction.
func diagonalSpectrum(t *testing.T, sigma []float64) *core.Array {
t.Helper()
vals := make([]float64, len(sigma)*len(sigma))
for i, s := range sigma {
vals[i*len(sigma)+i] = s
}
a, err := core.FromFloats(vals, len(sigma), len(sigma))
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
return a
}
// TestSolveTruncatedFullRank checks the untruncated case against the
// plain solve on a well-conditioned system.
func TestSolveTruncatedFullRank(t *testing.T) {
a := mustFloats(t, []float64{3, 0, 1, 1, 2, 1, 1, 1, 2, 1, 4, 0, 0, 1, 1, 5}, 4, 4)
b := mustFloats(t, []float64{1, 2, 3, 4}, 4)
x, err := SolveTruncated(a, b, 4)
if err != nil {
t.Fatalf("SolveTruncated: %v", err)
}
want, err := Solve(a, b)
if err != nil {
t.Fatalf("Solve: %v", err)
}
for i := range 4 {
if math.Abs(x.FloatAt(i)-want.FloatAt(i)) > 1e-8 {
t.Fatalf("x[%d] = %.12g, want %.12g", i, x.FloatAt(i), want.FloatAt(i))
}
}
}
// TestSolveTruncatedCutsNoise shows the point of truncation: with a
// spectrum of 10, 1, 1e-8, 1e-12 and b = 1 everywhere, the full solve
// amplifies b by 1e12 while rank-2 truncation answers (0.1, 1, 0, 0).
func TestSolveTruncatedCutsNoise(t *testing.T) {
a := diagonalSpectrum(t, []float64{10, 1, 1e-8, 1e-12})
b := mustFloats(t, []float64{1, 1, 1, 1}, 4)
x, err := SolveTruncated(a, b, 2)
if err != nil {
t.Fatalf("SolveTruncated: %v", err)
}
want := []float64{0.1, 1, 0, 0}
for i := range 4 {
if math.Abs(x.FloatAt(i)-want[i]) > 1e-12 {
t.Fatalf("x[%d] = %.12g, want %.12g", i, x.FloatAt(i), want[i])
}
}
}
// TestSolveTikhonovDiagonal pins the per-direction damping factor
// σ/(σ²+λ) on the diagonal system: large directions pass nearly
// untouched, tiny directions are crushed.
func TestSolveTikhonovDiagonal(t *testing.T) {
const lambda = 1.0
a := diagonalSpectrum(t, []float64{10, 1, 1e-8, 1e-12})
b := mustFloats(t, []float64{1, 1, 1, 1}, 4)
x, err := SolveTikhonov(a, b, lambda)
if err != nil {
t.Fatalf("SolveTikhonov: %v", err)
}
for i, s := range []float64{10, 1, 1e-8, 1e-12} {
want := s / (s*s + lambda)
if math.Abs(x.FloatAt(i)-want) > 1e-12*math.Max(1, want) {
t.Fatalf("x[%d] = %.14g, want %.14g", i, x.FloatAt(i), want)
}
}
}
// TestSolveTikhonovNormalEquations cross-checks a rectangular
// problem against the normal equations (AᵀA + λI)x = Aᵀb solved by
// the plain dense solver.
func TestSolveTikhonovNormalEquations(t *testing.T) {
const lambda = 0.7
a := mustFloats(t, []float64{2, 0, 1, 1, 0, 3, 1, 1, 1, 4, 2, 0}, 4, 3)
b := mustFloats(t, []float64{1, 0, 2, 1}, 4)
x, err := SolveTikhonov(a, b, lambda)
if err != nil {
t.Fatalf("SolveTikhonov: %v", err)
}
// Normal equations, assembled row-major for Solve.
ata := make([]float64, 9)
atb := make([]float64, 3)
for i := range 3 {
for j := range 3 {
s := 0.0
for l := range 4 {
s += a.FloatAt(l*3+i) * a.FloatAt(l*3+j)
}
ata[i*3+j] = s
if i == j {
ata[i*3+j] += lambda
}
}
s := 0.0
for l := range 4 {
s += a.FloatAt(l*3+i) * b.FloatAt(l)
}
atb[i] = s
}
want, err := Solve(mustFloats(t, ata, 3, 3), mustFloats(t, atb, 3))
if err != nil {
t.Fatalf("Solve: %v", err)
}
for i := range 3 {
if math.Abs(x.FloatAt(i)-want.FloatAt(i)) > 1e-8 {
t.Fatalf("x[%d] = %.12g, want %.12g", i, x.FloatAt(i), want.FloatAt(i))
}
}
}
// TestSolveSVDMultiColumn checks a two-column right-hand side against
// two single-column solves.
func TestSolveSVDMultiColumn(t *testing.T) {
a := mustFloats(t, []float64{3, 0, 1, 1, 2, 1, 1, 1, 2, 1, 4, 0}, 4, 3)
b2 := mustFloats(t, []float64{1, 2, 3, 4, 4, 3, 2, 1}, 4, 2)
x, err := SolveTikhonov(a, b2, 0.5)
if err != nil {
t.Fatalf("SolveTikhonov: %v", err)
}
if x.NDim() != 2 || x.Shape()[0] != 3 || x.Shape()[1] != 2 {
t.Fatalf("solution shape %v, want [3 2]", x.Shape())
}
for j := range 2 {
col := make([]float64, 4)
for i := range 4 {
col[i] = b2.FloatAt(i*2 + j)
}
single, err := SolveTikhonov(a, mustFloats(t, col, 4), 0.5)
if err != nil {
t.Fatalf("SolveTikhonov column %d: %v", j, err)
}
for i := range 3 {
if math.Abs(x.FloatAt(i*2+j)-single.FloatAt(i)) > 1e-9 {
t.Fatalf("column %d, row %d: %.12g, want %.12g", j, i, x.FloatAt(i*2+j), single.FloatAt(i))
}
}
}
}
// TestSolveSVDErrors pins the validation contract of both solvers.
func TestSolveSVDErrors(t *testing.T) {
a := diagonalSpectrum(t, []float64{10, 1, 1e-8, 1e-12})
b := mustFloats(t, []float64{1, 1, 1, 1}, 4)
if _, err := SolveTruncated(a, b, 0); err == nil {
t.Fatal("expected an error for rank 0")
}
if _, err := SolveTruncated(a, b, 5); err == nil {
t.Fatal("expected an error for rank beyond the spectrum")
}
if _, err := SolveTikhonov(a, b, 0); err == nil {
t.Fatal("expected an error for a non-positive lambda")
}
shortB := mustFloats(t, []float64{1, 1}, 2)
if _, err := SolveTruncated(a, shortB, 2); err == nil {
t.Fatal("expected an error for mismatched b rows")
}
if _, err := SolveTikhonov(shortB, b, 1); err == nil {
t.Fatal("expected an error for a rank-1 matrix")
}
singular := mustFloats(t, []float64{1, 1, 1, 1, 1, 1, 1, 1, 1}, 3, 3)
if _, err := SolveTruncated(singular, mustFloats(t, []float64{1, 1, 1}, 3), 3); err == nil {
t.Fatal("expected an error for a system rank-deficient below the requested rank")
}
}