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
|
|||
|
|
}
|