168 lines
4.3 KiB
Go
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)
|
|
}
|
|
}
|