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