423 lines
14 KiB
Go
423 lines
14 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
|
|
// SPDX-License-Identifier: MIT
|
|||
|
|
|
|||
|
|
package linalg
|
|||
|
|
|
|||
|
|
import (
|
|||
|
|
"math"
|
|||
|
|
"slices"
|
|||
|
|
"sync"
|
|||
|
|
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/engine"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// Rank-revealing QR with column pivoting (Businger and Golub). The
|
|||
|
|
// plain Householder QR of `QR` answers a factorisation but hides the
|
|||
|
|
// rank: an exactly or nearly dependent column leaves a rounding-level
|
|||
|
|
// R diagonal and `LeastSquares` has to refuse the system outright.
|
|||
|
|
// Pivoting sweeps the column of largest remaining norm to the front at
|
|||
|
|
// every step, which forces the dependence into the trailing block: the
|
|||
|
|
// |R| diagonal comes out non-increasing and its decay is the rank
|
|||
|
|
// decision, the same tolerance convention `LeastSquares` applies to
|
|||
|
|
// its unpivoted diagonal.
|
|||
|
|
|
|||
|
|
// RRQR returns the rank-revealing QR factorisation with column
|
|||
|
|
// pivoting of an m×n matrix a with m ≥ n: A·P = Q·R, with Q m×m
|
|||
|
|
// orthogonal, R m×n upper triangular, and perm the column permutation
|
|||
|
|
// P, meaning the j-th column of A·P is column perm[j] of A. The rank
|
|||
|
|
// is the count of |R[i,i]| above n·eps·max|R|, the tolerance
|
|||
|
|
// `LeastSquares` applies; with column pivoting the diagonal decays
|
|||
|
|
// monotonically, so the count is where the decay crosses the floor.
|
|||
|
|
// Complex matrices, non-2-D inputs and underdetermined shapes are
|
|||
|
|
// refused, matching the dense surface.
|
|||
|
|
func RRQR(a *core.Array) (q, r *core.Array, perm []int, rank int, err error) {
|
|||
|
|
const name = "RRQR"
|
|||
|
|
if a.Dtype() == core.Complex {
|
|||
|
|
return nil, nil, nil, 0, base.Errf("%s: complex matrices are not supported", name)
|
|||
|
|
}
|
|||
|
|
if a.NDim() != 2 {
|
|||
|
|
return nil, nil, nil, 0, base.Errf("%s: needs a 2-D matrix, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
m, n := a.Shape()[0], a.Shape()[1]
|
|||
|
|
if m == 0 || n == 0 {
|
|||
|
|
return nil, nil, nil, 0, base.Errf("%s: zero-sized matrix, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
if m < n {
|
|||
|
|
return nil, nil, nil, 0, base.Errf("%s: needs m ≥ n, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
rMat, qMat, perm := rrqrFactor(denseFloats(a, m, n), m, n, true)
|
|||
|
|
return floatsToArray(qMat, []int{m, m}), floatsToArray(rMat, []int{m, n}), perm, rrqrRankOf(rMat, n), nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// RRQRRank returns the numerical rank a pivoted factorisation of a
|
|||
|
|
// reports under the package's tolerance convention: the count of
|
|||
|
|
// |R[i,i]| above n·eps·max|R|. It is the rank `RRQR` returns, computed
|
|||
|
|
// without the orthogonal factor for callers that only want the count.
|
|||
|
|
func RRQRRank(a *core.Array) (int, error) {
|
|||
|
|
const name = "RRQRRank"
|
|||
|
|
if a.Dtype() == core.Complex {
|
|||
|
|
return 0, base.Errf("%s: complex matrices are not supported", name)
|
|||
|
|
}
|
|||
|
|
if a.NDim() != 2 {
|
|||
|
|
return 0, base.Errf("%s: needs a 2-D matrix, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
m, n := a.Shape()[0], a.Shape()[1]
|
|||
|
|
if m == 0 || n == 0 {
|
|||
|
|
return 0, base.Errf("%s: zero-sized matrix, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
if m < n {
|
|||
|
|
return 0, base.Errf("%s: needs m ≥ n, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
rMat, _, _ := rrqrFactor(denseFloats(a, m, n), m, n, false)
|
|||
|
|
return rrqrRankOf(rMat, n), nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// rrqrRankOf counts the diagonal entries of the flat m×n R above
|
|||
|
|
// n·eps·max|R|, the rank convention `LeastSquares` states for its
|
|||
|
|
// unpivoted diagonal.
|
|||
|
|
func rrqrRankOf(rMat []float64, n int) int {
|
|||
|
|
rMax := 0.0
|
|||
|
|
for i := range n {
|
|||
|
|
if v := math.Abs(rMat[i*n+i]); v > rMax {
|
|||
|
|
rMax = v
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
rankTol := float64(n) * base.EpsF * rMax
|
|||
|
|
rank := 0
|
|||
|
|
for i := range n {
|
|||
|
|
if math.Abs(rMat[i*n+i]) > rankTol {
|
|||
|
|
rank++
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return rank
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// rrqrColumnBlock is the number of R columns one blocked reflector pass
|
|||
|
|
// covers, the blocking qrRaw applies to the same sweep and for the same
|
|||
|
|
// reason: R is row-major, so a lone column's update walks a stride-n
|
|||
|
|
// pattern that touches one cache line per element, where a block of
|
|||
|
|
// rrqrColumnBlock columns spans one line of float64 and walks each row
|
|||
|
|
// once.
|
|||
|
|
const rrqrColumnBlock = 8
|
|||
|
|
|
|||
|
|
// rrqrMinWork is the element-touch floor below which one worker's share
|
|||
|
|
// of a reflector application or of a column-norm scan falls back to the
|
|||
|
|
// calling goroutine: below it the dispatch costs more than the work.
|
|||
|
|
const rrqrMinWork = 2048
|
|||
|
|
|
|||
|
|
// rrqrFactor runs the pivoted Householder sweep on a flat m×n matrix.
|
|||
|
|
// The column of largest remaining norm moves to the front at every
|
|||
|
|
// step; norms are recomputed each step rather than downdated, so no
|
|||
|
|
// cancellation drift and no squared-overflow window: the cost is one
|
|||
|
|
// pass over the working block per step, the same order as the
|
|||
|
|
// reflector sweep itself. The Householder reflectors and their
|
|||
|
|
// application follow `qrRaw`: scale before squaring, one reflector per
|
|||
|
|
// step, strictly in sweep order, so the serial result is the
|
|||
|
|
// factorisation the parallel plain QR also answers.
|
|||
|
|
func rrqrFactor(aMat []float64, m, n int, wantQ bool) (rMat, qMat []float64, perm []int) {
|
|||
|
|
rMat = make([]float64, m*n)
|
|||
|
|
copy(rMat, aMat)
|
|||
|
|
perm = make([]int, n)
|
|||
|
|
for i := range perm {
|
|||
|
|
perm[i] = i
|
|||
|
|
}
|
|||
|
|
if wantQ {
|
|||
|
|
qMat = eye(m)
|
|||
|
|
}
|
|||
|
|
col := make([]float64, m)
|
|||
|
|
norms := make([]float64, n)
|
|||
|
|
for k := range n {
|
|||
|
|
best := rrqrPivot(rMat, m, n, k, perm, norms)
|
|||
|
|
if best == 0 {
|
|||
|
|
// Every remaining column is exactly zero: the trailing
|
|||
|
|
// block of R stays zero and the sweep is done.
|
|||
|
|
break
|
|||
|
|
}
|
|||
|
|
lv := col[:m-k]
|
|||
|
|
for i := range lv {
|
|||
|
|
lv[i] = rMat[(k+i)*n+k]
|
|||
|
|
}
|
|||
|
|
hh := householderVectorInto(lv, lv)
|
|||
|
|
if hh.beta == 0 {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
rrqrApply(rMat, m, n, k, hh, qMat, nil)
|
|||
|
|
}
|
|||
|
|
return rMat, qMat, perm
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// rrqrPivot selects the column of largest remaining norm among columns
|
|||
|
|
// k..n−1, swaps it into column k and records the swap in perm. The
|
|||
|
|
// norms themselves are computed one worker per contiguous column slab,
|
|||
|
|
// each column by the serial two-pass accumulation stridedNorm2 always
|
|||
|
|
// ran, and the largest is then picked by a serial scan whose strict
|
|||
|
|
// comparison keeps the first of equal maxima, so the tie rule is
|
|||
|
|
// unchanged. A zero return means every remaining column is exactly zero.
|
|||
|
|
//
|
|||
|
|
// A slab is a run of consecutive columns, which is what makes the scan
|
|||
|
|
// affordable at all: column j and column j+1 share seven of every eight
|
|||
|
|
// cache lines their strided walks touch, so a slab re-reads them from
|
|||
|
|
// the cache a lone column at a time cannot. The crew only forms when one
|
|||
|
|
// worker's slab holds a full dispatch floor of entries.
|
|||
|
|
func rrqrPivot(rMat []float64, m, n, k int, perm []int, norms []float64) float64 {
|
|||
|
|
cols, rows := n-k, m-k
|
|||
|
|
engine.ParallelMin(cols, max((rrqrMinWork+rows-1)/rows, 1), func(start, end int) {
|
|||
|
|
for j := start; j < end; j++ {
|
|||
|
|
norms[j] = stridedNorm2(rMat, k*n+k+j, n, rows)
|
|||
|
|
}
|
|||
|
|
})
|
|||
|
|
best, bestJ := 0.0, k
|
|||
|
|
for j := range cols {
|
|||
|
|
if nrm := norms[j]; nrm > best {
|
|||
|
|
best, bestJ = nrm, k+j
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
if best == 0 {
|
|||
|
|
return 0
|
|||
|
|
}
|
|||
|
|
if bestJ != k {
|
|||
|
|
rrqrSwapCols(rMat, m, n, bestJ, k)
|
|||
|
|
perm[bestJ], perm[k] = perm[k], perm[bestJ]
|
|||
|
|
}
|
|||
|
|
return best
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// rrqrApply applies the reflector to the R columns k..n−1 in column
|
|||
|
|
// blocks and, when qMat is non-nil, to the m rows of the orthonormal
|
|||
|
|
// accumulator, or when bv is non-nil to the tail of the right-hand side.
|
|||
|
|
// The targets are disjoint buffers, so they share one crew; the R blocks
|
|||
|
|
// and the accumulator rows take contiguous slices, and every dot keeps
|
|||
|
|
// its own i-ascending chain with the subtraction repeating the same
|
|||
|
|
// operand order, so no crew size and no block boundary moves a bit.
|
|||
|
|
func rrqrApply(rMat []float64, m, n, k int, hh householder, qMat, bv []float64) {
|
|||
|
|
ln := m - k
|
|||
|
|
blocks := (n - k + rrqrColumnBlock - 1) / rrqrColumnBlock
|
|||
|
|
items := blocks
|
|||
|
|
if qMat != nil {
|
|||
|
|
items += m
|
|||
|
|
}
|
|||
|
|
if bv != nil {
|
|||
|
|
items++
|
|||
|
|
}
|
|||
|
|
weight := (n - k) + m
|
|||
|
|
w := min(weight*ln/rrqrMinWork+1, engine.WorkersFor(items))
|
|||
|
|
if w < 2 {
|
|||
|
|
rrqrApplyWorker(rMat, m, n, k, hh, qMat, bv, blocks, items, 0, 1)
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
var wg sync.WaitGroup
|
|||
|
|
for gi := range w {
|
|||
|
|
wg.Go(func() {
|
|||
|
|
rrqrApplyWorker(rMat, m, n, k, hh, qMat, bv, blocks, items, gi, w)
|
|||
|
|
})
|
|||
|
|
}
|
|||
|
|
wg.Wait()
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// rrqrApplyWorker runs crew member gi of nw over its share of the
|
|||
|
|
// reflector's targets: a contiguous slab of the R column blocks, a
|
|||
|
|
// contiguous slab of the accumulator rows, and the right-hand side.
|
|||
|
|
func rrqrApplyWorker(rMat []float64, m, n, k int, hh householder, qMat, bv []float64, blocks, items, gi, nw int) {
|
|||
|
|
ln := m - k
|
|||
|
|
// Target layout in the item index space: the R blocks first, then the
|
|||
|
|
// accumulator rows when there is an accumulator, then the right-hand
|
|||
|
|
// side when there is one.
|
|||
|
|
qEnd := blocks
|
|||
|
|
if qMat != nil {
|
|||
|
|
qEnd += m
|
|||
|
|
}
|
|||
|
|
var d [rrqrColumnBlock]float64
|
|||
|
|
for it := gi * items / nw; it < (gi+1)*items/nw; it++ {
|
|||
|
|
switch {
|
|||
|
|
case it < blocks:
|
|||
|
|
j0 := k + it*rrqrColumnBlock
|
|||
|
|
j1 := min(j0+rrqrColumnBlock, n)
|
|||
|
|
wd := j1 - j0
|
|||
|
|
_ = d[wd-1] // bounds-check proof: the block never exceeds rrqrColumnBlock
|
|||
|
|
clear(d[:wd])
|
|||
|
|
for jj := range wd {
|
|||
|
|
jcol := j0 + jj
|
|||
|
|
for i := range ln {
|
|||
|
|
d[jj] += hh.v[i] * rMat[(k+i)*n+jcol]
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
for j := range wd {
|
|||
|
|
d[j] = hh.beta * d[j]
|
|||
|
|
}
|
|||
|
|
for i := range ln {
|
|||
|
|
vi := hh.v[i]
|
|||
|
|
base := (k+i)*n + j0
|
|||
|
|
for j := range wd {
|
|||
|
|
rMat[base+j] -= vi * d[j]
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
case it < qEnd:
|
|||
|
|
row := it - blocks
|
|||
|
|
base := row*m + k
|
|||
|
|
dot := 0.0
|
|||
|
|
for i := range ln {
|
|||
|
|
dot += hh.v[i] * qMat[base+i]
|
|||
|
|
}
|
|||
|
|
wq := hh.beta * dot
|
|||
|
|
for i := range ln {
|
|||
|
|
qMat[base+i] -= hh.v[i] * wq
|
|||
|
|
}
|
|||
|
|
default:
|
|||
|
|
dot := 0.0
|
|||
|
|
for i := range ln {
|
|||
|
|
dot += hh.v[i] * bv[k+i]
|
|||
|
|
}
|
|||
|
|
wb := hh.beta * dot
|
|||
|
|
for i := range ln {
|
|||
|
|
bv[k+i] -= hh.v[i] * wb
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// rrqrSwapCols exchanges columns j1 and j2 of a flat m×n matrix.
|
|||
|
|
func rrqrSwapCols(rMat []float64, m, n, j1, j2 int) {
|
|||
|
|
for i := range m {
|
|||
|
|
rMat[i*n+j1], rMat[i*n+j2] = rMat[i*n+j2], rMat[i*n+j1]
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// stridedNorm2 returns the L2 norm of count entries walking stride
|
|||
|
|
// step from start, summed relative to the largest magnitude so a
|
|||
|
|
// column near 1e154 does not overflow on the way to its norm.
|
|||
|
|
func stridedNorm2(vals []float64, start, step, count int) float64 {
|
|||
|
|
maxAbs := 0.0
|
|||
|
|
for i := range count {
|
|||
|
|
if a := math.Abs(vals[start+i*step]); a > maxAbs {
|
|||
|
|
maxAbs = a
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
if maxAbs == 0 {
|
|||
|
|
return 0
|
|||
|
|
}
|
|||
|
|
s := 0.0
|
|||
|
|
for i := range count {
|
|||
|
|
t := vals[start+i*step] / maxAbs
|
|||
|
|
s += t * t
|
|||
|
|
}
|
|||
|
|
return maxAbs * math.Sqrt(s)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// rrqrValidPerm reports whether p is a permutation of 0..n−1, the
|
|||
|
|
// property the returned permutation's contract rests on.
|
|||
|
|
func rrqrValidPerm(p []int) bool {
|
|||
|
|
seen := make([]bool, len(p))
|
|||
|
|
for _, v := range p {
|
|||
|
|
if v < 0 || v >= len(p) || seen[v] {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
seen[v] = true
|
|||
|
|
}
|
|||
|
|
return !slices.Contains(seen, false)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// SolveRRQR solves min ‖A·x − b‖₂ through the pivoted factorisation.
|
|||
|
|
// On full rank the answer is the ordinary back-substitution through R
|
|||
|
|
// against c = Qᵀb. On rank-deficient input it is the minimum-norm
|
|||
|
|
// least-squares solution: the rank r is read off the R diagonal decay,
|
|||
|
|
// and the leading r×n trapezoidal block R₁ answers its underdetermined
|
|||
|
|
// system in minimum norm through x = R₁ᵀ(R₁R₁ᵀ)⁻¹c₁, which squares the
|
|||
|
|
// condition of the leading block exactly the way the package's
|
|||
|
|
// BᵀB-based SVD squares the spectrum: honest for a rank decision taken
|
|||
|
|
// from a decay that has already been observed.
|
|||
|
|
func SolveRRQR(a, b *core.Array) (*core.Array, error) {
|
|||
|
|
const name = "SolveRRQR"
|
|||
|
|
if a.Dtype() == core.Complex {
|
|||
|
|
return nil, base.Errf("%s: complex matrices are not supported", name)
|
|||
|
|
}
|
|||
|
|
if b.Dtype() == core.Complex {
|
|||
|
|
return nil, base.Errf("%s: complex right-hand sides are not supported", name)
|
|||
|
|
}
|
|||
|
|
if a.NDim() != 2 {
|
|||
|
|
return nil, base.Errf("%s: needs a 2-D matrix, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
m, n := a.Shape()[0], a.Shape()[1]
|
|||
|
|
if m == 0 || n == 0 {
|
|||
|
|
return nil, base.Errf("%s: zero-sized matrix, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
if m < n {
|
|||
|
|
return nil, base.Errf("%s: needs m ≥ n, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
if b.NDim() != 1 || b.Len() != m {
|
|||
|
|
return nil, base.Errf("%s: right-hand side must be a vector of length %d, got shape %s",
|
|||
|
|
name, m, base.ShapeText(b.Shape()))
|
|||
|
|
}
|
|||
|
|
// The sweep below carries b through the reflectors as they are
|
|||
|
|
// built, so c = Qᵀb falls out of the factorisation without the
|
|||
|
|
// orthogonal factor ever being materialised.
|
|||
|
|
rMat := make([]float64, m*n)
|
|||
|
|
copy(rMat, denseFloats(a, m, n))
|
|||
|
|
bv := vectorF64(b, m)
|
|||
|
|
perm := make([]int, n)
|
|||
|
|
for i := range perm {
|
|||
|
|
perm[i] = i
|
|||
|
|
}
|
|||
|
|
col := make([]float64, m)
|
|||
|
|
norms := make([]float64, n)
|
|||
|
|
for k := range n {
|
|||
|
|
if rrqrPivot(rMat, m, n, k, perm, norms) == 0 {
|
|||
|
|
break
|
|||
|
|
}
|
|||
|
|
lv := col[:m-k]
|
|||
|
|
for i := range lv {
|
|||
|
|
lv[i] = rMat[(k+i)*n+k]
|
|||
|
|
}
|
|||
|
|
hh := householderVectorInto(lv, lv)
|
|||
|
|
if hh.beta == 0 {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
rrqrApply(rMat, m, n, k, hh, nil, bv)
|
|||
|
|
}
|
|||
|
|
x := make([]float64, n)
|
|||
|
|
switch rank := rrqrRankOf(rMat, n); {
|
|||
|
|
case rank == n:
|
|||
|
|
// Full rank: back-substitution through R.
|
|||
|
|
for i := n - 1; i >= 0; i-- {
|
|||
|
|
s := bv[i]
|
|||
|
|
for j := i + 1; j < n; j++ {
|
|||
|
|
s -= rMat[i*n+j] * x[j]
|
|||
|
|
}
|
|||
|
|
x[i] = s / rMat[i*n+i]
|
|||
|
|
}
|
|||
|
|
default:
|
|||
|
|
// Minimum norm through the leading trapezoidal block R₁:
|
|||
|
|
// R₁R₁ᵀw = c₁, x = R₁ᵀw. R₁R₁ᵀ is symmetric positive definite
|
|||
|
|
// at the rank the decay showed, and the dense `Solve` carries
|
|||
|
|
// the small r×r system.
|
|||
|
|
if rank > 0 {
|
|||
|
|
mm := make([]float64, rank*rank)
|
|||
|
|
for i := range rank {
|
|||
|
|
for j := range rank {
|
|||
|
|
s := 0.0
|
|||
|
|
for l := range n {
|
|||
|
|
s += rMat[i*n+l] * rMat[j*n+l]
|
|||
|
|
}
|
|||
|
|
mm[i*rank+j] = s
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
w, err := Solve(floatsToArray(mm, []int{rank, rank}), floatsToArray(append([]float64(nil), bv[:rank]...), []int{rank}))
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, base.Errf("%s: %w", name, err)
|
|||
|
|
}
|
|||
|
|
for i := range rank {
|
|||
|
|
for j := range n {
|
|||
|
|
x[j] += rMat[i*n+j] * w.FloatAt(i)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
// x lives in permuted coordinates; undo the permutation.
|
|||
|
|
out := make([]float64, n)
|
|||
|
|
for j := range n {
|
|||
|
|
out[perm[j]] = x[j]
|
|||
|
|
}
|
|||
|
|
return floatsToArray(out, []int{n}), nil
|
|||
|
|
}
|