Files
tensor/linalg/gmres_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

168 lines
4.3 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/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
func TestGMRES(t *testing.T) {
vals := []float64{
10, 1, 0, 2,
1, 12, 3, 0,
0, 3, 15, 1,
2, 0, 1, 8,
}
op := func(v *core.Array) (*core.Array, error) {
out := core.New(core.Float, 4)
for i := range 4 {
s := 0.0
for j := range 4 {
s += vals[i*4+j] * v.FloatAt(j)
}
out.RawFloats()[i] = s
}
return out, nil
}
b := mustFloats(t, []float64{1, 2, 3, 4}, 4)
x, err := GMRES(op, b, 0, 0, 1e-12)
if err != nil {
t.Fatalf("GMRES: %v", err)
}
a := mustFloats(t, vals, 4, 4)
ref, _ := Solve(a, b)
for i := range 4 {
if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-10 {
t.Fatalf("x[%d] = %.16g, want %.16g", i, x.FloatAt(i), ref.FloatAt(i))
}
}
}
func TestGMRESLarger(t *testing.T) {
const n = 50
vals := make([]float64, n*n)
for i := range n {
vals[i*n+i] = 4
if i+1 < n {
vals[i*n+i+1] = -1
vals[(i+1)*n+i] = -1
}
}
op := func(v *core.Array) (*core.Array, error) {
out := core.New(core.Float, n)
for i := range n {
s := 0.0
for j := range n {
s += vals[i*n+j] * v.FloatAt(j)
}
out.RawFloats()[i] = s
}
return out, nil
}
bv := make([]float64, n)
for i := range n {
bv[i] = float64(i + 1)
}
b := mustFloats(t, bv, n)
x, err := GMRES(op, b, 0, 0, 1e-12)
if err != nil {
t.Fatalf("GMRES: %v", err)
}
for i := range n {
s := 0.0
for j := range n {
s += vals[i*n+j] * x.FloatAt(j)
}
if math.Abs(s-bv[i]) > 1e-8 {
t.Fatalf("res[%d] = %e, want < 1e-8", i, math.Abs(s-bv[i]))
}
}
}
// TestGMRESRestarted forces several outer cycles by capping the
// Krylov dimension below the system size; the accumulated solution
// must still match the dense solve.
func TestGMRESRestarted(t *testing.T) {
vals := []float64{
10, 1, 0, 2,
1, 12, 3, 0,
0, 3, 15, 1,
2, 0, 1, 8,
}
op := func(v *core.Array) (*core.Array, error) {
out := core.New(core.Float, 4)
for i := range 4 {
s := 0.0
for j := range 4 {
s += vals[i*4+j] * v.FloatAt(j)
}
out.RawFloats()[i] = s
}
return out, nil
}
b := mustFloats(t, []float64{1, 2, 3, 4}, 4)
a := mustFloats(t, vals, 4, 4)
ref, err := Solve(a, b)
if err != nil {
t.Fatalf("Solve: %v", err)
}
x, err := GMRES(op, b, 2, 50, 1e-12)
if err != nil {
t.Fatalf("GMRES: %v", err)
}
for i := range 4 {
if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-10 {
t.Fatalf("x[%d] = %.16g, want %.16g", i, x.FloatAt(i), ref.FloatAt(i))
}
}
}
func TestGMRESErrors(t *testing.T) {
b := mustFloats(t, []float64{1, 2}, 2)
if _, err := GMRES(func(*core.Array) (*core.Array, error) {
return nil, base.Errf("operator failed")
}, b, 0, 0, 0); err == nil {
t.Fatal("operator error: want an error")
}
if x, err := GMRES(func(*core.Array) (*core.Array, error) {
return nil, base.Errf("operator failed")
}, mustFloats(t, []float64{}, 0), 0, 0, 0); err == nil {
t.Fatalf("empty system: want an error, got %v", x)
}
}
// TestGMRESUnconvergedErrors pins the contract shared with SpSolve and
// FindRootSystem: running out of cycles with the tolerance unmet is an
// error naming the residual, never a silent approximation.
func TestGMRESUnconvergedErrors(t *testing.T) {
// A = diag(1, 0.5): a legitimate but slow system for restart-1
// GMRES, which crawls towards the solution and cannot reach a
// 1e-12 relative residual in 2 cycles.
diag := []float64{1, 0.5}
op := func(v *core.Array) (*core.Array, error) {
out := core.New(core.Float, 2)
for i := range 2 {
out.RawFloats()[i] = diag[i] * v.FloatAt(i)
}
return out, nil
}
if _, err := GMRES(op, mustFloats(t, []float64{1, 1}, 2), 1, 2, 1e-12); err == nil {
t.Fatal("unconverged GMRES: want an error")
}
}
// TestNorm2F64LargeScale pins the overflow-safe norm: entries near
// 1e154 must not square their way to +Inf.
func TestNorm2F64LargeScale(t *testing.T) {
if got := norm2F64([]float64{1e154, 1e154}); math.IsInf(got, 1) || math.Abs(got-math.Sqrt2*1e154) > 1e-12*math.Sqrt2*1e154 {
t.Fatalf("norm2F64(1e154, 1e154) = %v, want ≈ %g", got, math.Sqrt2*1e154)
}
if got := norm2F64([]float64{0, 0}); got != 0 {
t.Fatalf("norm2F64(zeros) = %v, want 0", got)
}
}