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")
|
|||
|
|
}
|
|||
|
|
}
|