311 lines
9.3 KiB
Go
311 lines
9.3 KiB
Go
// 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
|
||
}
|