// Copyright (c) 2026 Petr Balvín (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) } }