181 lines
5.6 KiB
Go
181 lines
5.6 KiB
Go
// 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")
|
||
}
|
||
}
|