Files
tensor/linalg/gmres.go
T

311 lines
9.3 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package linalg
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import "math"
// GMRES returns the solution of A·x = b by restarted GMRES. op maps a
// vector of length n to A·v. restart bounds the Krylov dimension per
// cycle (≤ 0 means n); maxIter bounds the outer cycles (≤ 0 means 20);
// tol is the relative residual target (≤ 0 means 1e-10). Running out
// of cycles with the target unmet is an error naming the best residual
// achieved, never a silent approximation.
//
// The vector op is handed is a read-only buffer the solver owns and
// refills for every call: op must read it within the call and must not
// retain or modify it. An operator that needs to keep its argument has
// to copy it. The array op returns is op's own and the solver only
// reads it, so an operator may hand back a buffer it reuses.
//
// Every new Hessenberg column the Arnoldi process produces is folded
// through the Givens rotations of the previous columns the moment it
// appears, which leaves a triangular projected system behind and
// tracks the least-squares residual incrementally: s starts at β and
// each rotation j updates s[j+1] = −sn[j]·s[j], so |s[j+1]| is the
// residual of the growing fit without ever solving it. That keeps the
// projected problem at the condition of the Hessenberg matrix instead
// of its square, which is what the normal-equations route would pay.
func GMRES(op func(*core.Array) (*core.Array, error), b *core.Array, restart, maxIter int, tol float64) (*core.Array, error) {
if b.NDim() != 1 {
return nil, base.Errf("GMRES: b must be a rank-1 vector, got shape %s", base.ShapeText(b.Shape()))
}
n := b.Len()
if n == 0 {
return nil, base.Errf("GMRES: empty system")
}
if restart <= 0 || restart > n {
restart = n
}
if maxIter <= 0 {
maxIter = 20
}
if tol <= 0 {
tol = 1e-10
}
if !isRealDtype(b) {
return nil, base.Errf("GMRES: b must be a real dtype, got %s", b.Dtype())
}
bv := vectorF64(b, n)
bNorm := norm2F64(bv)
if bNorm == 0 {
return core.Zeros(core.Float, n)
}
x := make([]float64, n)
m := min(restart, n)
// The whole (m+1)×n basis and the (m+1)×m Hessenberg matrix live in
// one flat buffer each, allocated once per solve: a row is a
// subslice, so the Arnoldi process adds no allocation per step.
basisFlat := make([]float64, (m+1)*n)
basis := make([][]float64, m+1)
for i := range m + 1 {
basis[i] = basisFlat[i*n : (i+1)*n]
}
hessFlat := make([]float64, (m+1)*m)
hess := make([][]float64, m+1)
for i := range m + 1 {
hess[i] = hessFlat[i*m : (i+1)*m]
}
cs := make([]float64, m)
sn := make([]float64, m)
s := make([]float64, m+1)
// The working vectors are reused across columns and cycles: w holds
// the current Arnoldi remainder, r the residual, ax A·x, and y the
// triangular solve (only its first nsteps entries are read per
// cycle).
w := make([]float64, n)
r := make([]float64, n)
axv := make([]float64, n)
y := make([]float64, m)
// The operand op receives is one buffer the solver owns and
// refills before every call, so a solve allocates nothing per
// right-hand side. The contract that makes it safe is stated in the
// doc comment: op must not retain or modify the vector.
opBuf := core.New(core.Float, n)
opVals := opBuf.RawFloats()
applyOp := func(src []float64) (*core.Array, error) {
copy(opVals, src)
return op(opBuf)
}
best := math.Inf(1)
converged := false
// The op call that closes a cycle computes A·x for the x the cycle
// produced, which is exactly what the next cycle's residual needs,
// so the result is reused rather than recomputed for the same x.
// Every call to op still gets a fresh array: the callback may keep
// the vector it is given, so the operand is the one buffer here
// that cannot be recycled between calls.
freshAx := false
// The Arnoldi breakdown floor is relative to the operator's own
// Hessenberg scale, not to ‖b‖: a legitimate system with ‖A‖ ≪ ‖b‖
// would otherwise collapse every cycle at its first column. This is
// the same purely relative floor the Lanczos code documents.
hScale := 0.0
for range maxIter {
if !freshAx {
ax, aerr := applyOp(x)
if aerr != nil {
return nil, base.Errf("GMRES: %w", aerr)
}
if !isRealDtype(ax) {
return nil, base.Errf("GMRES: op returned %s, want a real dtype", ax.Dtype())
}
if ax.Len() != n {
return nil, base.Errf("GMRES: op returned %d elements for a problem of %d", ax.Len(), n)
}
vecF64Into(axv, ax, n)
}
// Residual r = b − A·x.
for i := range n {
r[i] = bv[i] - axv[i]
}
// norm2F64 skips NaN entries, so an all-NaN residual would read
// as a zero norm and as convergence; a non-finite op output is a
// breakdown before that.
if !vecFinite(r) {
return nil, base.Errf("GMRES: non-finite residual from op")
}
beta := norm2F64(r)
if beta < best {
best = beta
}
if beta <= tol*bNorm {
converged = true
break
}
for i := range n {
basis[0][i] = r[i] / beta
}
s[0] = beta
// Arnoldi, one column at a time, with the rotations folded in
// before the next column is built.
nsteps := 0
for j := range m {
opv, oerr := applyOp(basis[j])
if oerr != nil {
return nil, base.Errf("GMRES: %w", oerr)
}
if !isRealDtype(opv) {
return nil, base.Errf("GMRES: op returned %s, want a real dtype", opv.Dtype())
}
if opv.Len() != n {
return nil, base.Errf("GMRES: op returned %d elements for a problem of %d", opv.Len(), n)
}
vecF64Into(w, opv, n)
if !vecFinite(w) {
return nil, base.Errf("GMRES: non-finite op output at Arnoldi column %d", j)
}
for i := 0; i <= j; i++ {
dot := dotF64(w, basis[i])
hess[i][j] = dot
hScale = max(hScale, math.Abs(dot))
for k := range n {
w[k] -= dot * basis[i][k]
}
}
hSub := norm2F64(w)
hess[j+1][j] = hSub
hScale = max(hScale, hSub)
nsteps = j + 1
// Fold rotations 0..j−1 into the fresh column.
for k := range j {
h1, h2 := hess[k][j], hess[k+1][j]
hess[k][j] = cs[k]*h1 + sn[k]*h2
hess[k+1][j] = -sn[k]*h1 + cs[k]*h2
}
// Rotation j zeroes the subdiagonal entry: the pair
// (hess[j][j], hSub) turns into (denom, 0).
denom := math.Hypot(hess[j][j], hSub)
hScale = max(hScale, denom)
if denom == 0 {
cs[j], sn[j] = 1, 0
} else {
cs[j], sn[j] = hess[j][j]/denom, hSub/denom
}
hess[j][j] = denom
hess[j+1][j] = 0
s[j+1] = -sn[j] * s[j]
s[j] *= cs[j]
// A collapsed direction means the Krylov space is
// exhausted; a converged estimate means the basis already
// spans the answer. Either way the cycle ends here.
if hSub <= float64(n)*base.EpsF*hScale || math.Abs(s[j+1]) <= tol*bNorm {
break
}
next := basis[j+1]
for k := range n {
next[k] = w[k] / hSub
}
}
// Triangular solve R·y = s[0:nsteps] over the rotated leading
// block, then x = x + V·y.
for i := nsteps - 1; i >= 0; i-- {
t := s[i]
for k := i + 1; k < nsteps; k++ {
t -= hess[i][k] * y[k]
}
if hess[i][i] == 0 {
return nil, base.Errf("GMRES: singular projected system at step %d", i)
}
y[i] = t / hess[i][i]
}
for i := range nsteps {
for k := range n {
x[k] += y[i] * basis[i][k]
}
}
// Trust only the true residual between cycles.
ax, aerr := applyOp(x)
if aerr != nil {
return nil, base.Errf("GMRES: %w", aerr)
}
if !isRealDtype(ax) {
return nil, base.Errf("GMRES: op returned %s, want a real dtype", ax.Dtype())
}
vecF64Into(axv, ax, n)
if !vecFinite(axv) {
return nil, base.Errf("GMRES: non-finite op output between cycles")
}
freshAx = true
res := 0.0
for k := range n {
d := bv[k] - axv[k]
res += d * d
}
residual := math.Sqrt(res)
if residual < best {
best = residual
}
if residual <= tol*bNorm {
converged = true
break
}
}
// Contract consistency with SpSolve and FindRootSystem: running out
// of cycles is an error naming the best residual achieved, never a
// silent approximation.
if !converged {
return nil, base.Errf("GMRES: no convergence in %d cycles, residual %.3g (tolerance %.3g)",
maxIter, best, tol*bNorm)
}
return floatsToArray(x, []int{n}), nil
}
// norm2F64 returns the L2 norm of a float64 slice. The squares are
// summed relative to the largest magnitude so entries near 1e154 do
// not overflow on the way to the norm.
func norm2F64(a []float64) float64 {
maxAbs := 0.0
for _, v := range a {
if v := math.Abs(v); v > maxAbs {
maxAbs = v
}
}
if maxAbs == 0 {
return 0
}
s := 0.0
for _, v := range a {
v /= maxAbs
s += v * v
}
return maxAbs * math.Sqrt(s)
}
// ArrayFromFloatsSafe builds an array, copying the values, so callers
// keep ownership of the slice they passed in.
func ArrayFromFloatsSafe(v []float64, n int) *core.Array {
vals := make([]float64, n)
copy(vals, v)
a, _ := core.FloatsFromArray(vals, n)
return a
}
// vecF64Into fills dst with an array's leading n elements, the
// writing twin of vectorF64 for a destination that is reused.
func vecF64Into(dst []float64, a *core.Array, n int) {
for i := range n {
dst[i] = a.FloatAt(i)
}
}
// dotF64 returns the dot product of two float64 slices.
func dotF64(a, b []float64) float64 {
s := 0.0
for i := range a {
s += a[i] * b[i]
}
return s
}