Files

609 lines
20 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package linalg
import (
"math"
"sync"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"sourcedock.dev/petrbalvin/tensor/internal/engine"
)
// Matrix decompositions. All routines operate on float64 in
// and out; complex matrices are rejected. The QR routine uses
// Householder reflections, numerically stable and standard. Cholesky
// uses the Banachiewicz algorithm. LeastSquares builds on QR.
// QR returns the QR decomposition a = Q * R of an m×n matrix a, with
// m ≥ n. Q is m×m orthogonal and R is m×n upper triangular.
func QR(a *core.Array) (q, r *core.Array, err error) {
if a.Dtype() == core.Complex {
return nil, nil, base.Errf("QR: complex matrices are not supported")
}
if a.NDim() != 2 {
return nil, nil, base.Errf("QR: needs a 2-D matrix, got shape %s", base.ShapeText(a.Shape()))
}
m, n := a.Shape()[0], a.Shape()[1]
if m < n {
return nil, nil, base.Errf("QR: needs m ≥ n, got shape %s", base.ShapeText(a.Shape()))
}
aMat := denseFloats(a, m, n)
qMat, rMat, err := qrRaw(aMat, m, n)
if err != nil {
return nil, nil, err
}
return floatsToArray(qMat, []int{m, m}), floatsToArray(rMat, []int{m, n}), nil
}
// Cholesky returns the lower-triangular Cholesky factor L such that
// a = L * Lᵀ for a symmetric positive definite matrix a.
func Cholesky(a *core.Array) (*core.Array, error) {
if a.Dtype() == core.Complex {
return nil, base.Errf("Cholesky: complex matrices are not supported")
}
if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] {
return nil, base.Errf("Cholesky: needs a square 2-D matrix, got shape %s", base.ShapeText(a.Shape()))
}
n := a.Shape()[0]
aMat := denseFloats(a, n, n)
// A non-finite entry is neither positive nor definite, but the pivot
// test below cannot see it: NaN fails every comparison, so the sweep
// would return an all-NaN factor with a nil error. The sparse twin
// refuses the same input, and so does this one.
for i := range n {
for j := range n {
if math.IsNaN(aMat[i*n+j]) || math.IsInf(aMat[i*n+j], 0) {
return nil, base.Errf("Cholesky: entry [%d,%d] is not finite", i, j)
}
}
}
lMat, err := denseCholFactor(aMat, n)
if err != nil {
return nil, err
}
return floatsToArray(lMat, []int{n, n}), nil
}
// denseCholBlock is the column-block width of the Cholesky sweep. A wider
// block shortens the serial diagonal factorisation and lengthens the
// parallel panel update's k run; 64 keeps the serial share near a
// twentieth of the arithmetic at the sizes this package serves while
// both panel passes stay inside a cache line's neighbourhood.
const denseCholBlock = 64
// denseCholMinWork is the element-touch floor one Cholesky panel worker must
// receive before the dispatch is worth its spawn cost, counted in
// update terms (rows times columns times the k run). Below it the panel
// runs on the calling goroutine.
const denseCholMinWork = 8192
// denseCholFactor factors a flat n×n row-major matrix into its lower
// triangular Cholesky factor, consuming the input.
//
// The sweep runs in column blocks of denseCholBlock: each block is first
// updated from the finished columns, then factored serially, then
// divided into the trapezoid below it. For an element (i, j) the three
// stages subtract the columns k < j in the one ascending order the
// unblocked dot always used, since the panel update covers exactly the
// k < j0 prefix and the factorisation or the divide covers the rest,
// so the running value the last stage divides is bit for bit the value
// the unblocked loop held. Only the order in which independent elements
// are visited changes, which is why the panel update and the trapezoid
// divide, both disjoint per row, run over a crew.
func denseCholFactor(mat []float64, n int) ([]float64, error) {
lMat := make([]float64, n*n)
for j0 := 0; j0 < n; j0 += denseCholBlock {
j1 := min(j0+denseCholBlock, n)
// Panel update: every element of rows j0..n−1 in columns
// j0..j1−1 loses the contribution of the finished columns k < j0.
// A row's elements depend on those columns alone, so rows split
// over the crew.
if j0 > 0 {
cols := j1 - j0
panel := func(start, end int) {
for i := start; i < end; i++ {
for j := j0; j < j1 && j <= i; j++ {
s := mat[i*n+j]
for k := range j0 {
s -= lMat[i*n+k] * lMat[j*n+k]
}
mat[i*n+j] = s
}
}
}
if rows := n - j0; rows*cols*j0 >= denseCholMinWork {
engine.ParallelMin(rows, denseCholRows(denseCholMinWork, cols*j0), func(start, end int) {
panel(j0+start, j0+end)
})
} else {
panel(j0, n)
}
}
// Diagonal block: columns in order, each finished before the
// next reads it. This is the only serial arithmetic left, and it
// shrinks with the block width: the block's own triangle is
// denseCholBlock³/6 of the matrix's n³/6.
for j := j0; j < j1; j++ {
for i := j; i < j1; i++ {
s := mat[i*n+j]
for k := j0; k < j; k++ {
s -= lMat[i*n+k] * lMat[j*n+k]
}
if i == j {
if s <= 0 {
return nil, base.Errf("Cholesky: matrix is not positive definite (pivot %d = %v)", i, s)
}
lMat[i*n+j] = math.Sqrt(s)
} else {
if lMat[j*n+j] == 0 {
return nil, base.Errf("Cholesky: zero diagonal at %d", j)
}
lMat[i*n+j] = s / lMat[j*n+j]
}
}
}
// Trapezoid below the block: row i of the panel needs row j of the
// diagonal block, which the sweep above has finished, and then its
// own earlier columns of the same block; both are in hand, so the
// rows divide independently. The blocks are disjoint per row, so
// one crew covers the divide.
if j1 < n {
divide := func(start, end int) {
for j := j0; j < j1; j++ {
d := lMat[j*n+j]
for i := start; i < end; i++ {
s := mat[i*n+j]
for k := j0; k < j; k++ {
s -= lMat[i*n+k] * lMat[j*n+k]
}
lMat[i*n+j] = s / d
}
}
}
if rows := n - j1; rows*(j1-j0)*(j1-j0)/2 >= denseCholMinWork {
engine.ParallelMin(rows, denseCholRows(denseCholMinWork, (j1-j0)*(j1-j0)/2), func(start, end int) {
divide(j1+start, j1+end)
})
} else {
divide(j1, n)
}
}
}
return lMat, nil
}
// denseCholRows returns the rows one worker needs for the given per-row
// cost to reach the dispatch floor, never below one.
func denseCholRows(floor, cost int) int {
if cost <= 0 {
return 1
}
return max((floor+cost-1)/cost, 1)
}
// LeastSquares solves Ax = b in the least-squares sense. a is m×n
// with m ≥ n. b is m or m×k. Implemented via QR.
func LeastSquares(a, b *core.Array) (*core.Array, error) {
if a.Dtype() == core.Complex {
return nil, base.Errf("LeastSquares: complex matrices are not supported")
}
// The QR route reads b through FloatAt as well: without the gate a
// complex right-hand side reaches the empty int payload and panics.
if b.Dtype() == core.Complex {
return nil, base.Errf("LeastSquares: complex right-hand sides are not supported")
}
if a.NDim() != 2 {
return nil, base.Errf("LeastSquares: 'a' must be 2-D, got shape %s", base.ShapeText(a.Shape()))
}
m, n := a.Shape()[0], a.Shape()[1]
if m < n {
return nil, base.Errf("LeastSquares: needs m ≥ n, got shape %s", base.ShapeText(a.Shape()))
}
if b.NDim() != 1 && b.NDim() != 2 {
return nil, base.Errf("LeastSquares: 'b' must be 1-D or 2-D, got shape %s", base.ShapeText(b.Shape()))
}
if b.Shape()[0] != m {
return nil, base.Errf("LeastSquares: b rows (%d) must match a rows (%d)", b.Shape()[0], m)
}
k := 1
if b.NDim() == 2 {
k = b.Shape()[1]
}
aMat := denseFloats(a, m, n)
bMat := make([]float64, m*k)
if b.Dtype() == core.Float && !b.Strided() && len(b.RawFloats()) == m*k {
copy(bMat, b.RawFloats())
} else if b.NDim() == 1 {
for i := range m {
bMat[i] = b.FloatAt(i)
}
} else {
for i := range m {
for j := range k {
bMat[i*k+j] = b.FloatAt(i*k + j)
}
}
}
// The reflectors go straight onto b; no Q is formed, because the
// only thing the solve wants from it is Qᵀ·b.
qtB, rMat := qrSolveRows(aMat, m, n, k, bMat)
x := make([]float64, n*k)
// Rank guard: with no column pivoting an exactly or nearly
// dependent column leaves a rounding-level R diagonal instead of an
// exact zero, which back-substitution would amplify into a huge
// silent answer. Pivots below the Householder-QR analogue of the
// SVD route's threshold, n·eps·max|R|, refuse the system.
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
for col := 0; col < k; col++ {
for i := n - 1; i >= 0; i-- {
s := qtB[i*k+col]
for j := i + 1; j < n; j++ {
s -= rMat[i*n+j] * x[j*k+col]
}
if math.Abs(rMat[i*n+i]) <= rankTol {
return nil, base.Errf("LeastSquares: rank-deficient system (R pivot %d = %g is at the rounding floor %g)", i, rMat[i*n+i], rankTol)
}
x[i*k+col] = s / rMat[i*n+i]
}
}
if b.NDim() == 1 {
return floatsToArray(x, []int{n}), nil
}
return floatsToArray(x, []int{n, k}), nil
}
// householderReflector builds the reflector that zeroes the entries
// below the diagonal of column k of the row-major m×n matrix r and
// writes its scaled vector into v[:m−k]. It reports the vector length,
// beta = 2/(vᵀv) and whether a reflector was needed at all: a zero
// column, a zero norm or a zero vector leaves the column as the sweep
// wants it and the step is skipped.
//
// The scaling is the one the sweep has always used: entries near 1e154
// would overflow the raw squares to +Inf, which zeroes beta and forms
// 0*Inf = NaN in the update, so the column is divided by its largest
// magnitude first and beta absorbs the square of that scale, leaving
// the update beta·(vᵀy)·v at the value the unscaled arithmetic would
// have had, with every intermediate O(1).
func householderReflector(r []float64, m, n, k int, v []float64) (ln int, beta float64, ok bool) {
scale := 0.0
for i := k; i < m; i++ {
if a := math.Abs(r[i*n+k]); a > scale {
scale = a
}
}
if scale == 0 {
return 0, 0, false
}
sum := 0.0
for i := k; i < m; i++ {
t := r[i*n+k] / scale
sum += t * t
}
norm := math.Sqrt(sum)
if norm == 0 {
return 0, 0, false
}
xk := r[k*n+k] / scale
s := -sign(xk)
if s == 0 {
s = -1
}
u0 := xk - s*norm
vv := v[:m-k]
vv[0] = u0
for i := 1; i < m-k; i++ {
vv[i] = r[(k+i)*n+k] / scale
}
vtv := u0 * u0
for i := 1; i < m-k; i++ {
vtv += vv[i] * vv[i]
}
if vtv == 0 {
return 0, 0, false
}
return m - k, 2.0 / vtv, true
}
// qrSolveRows runs the Householder sweep on a row-major m×n system and
// applies every reflector to the k right-hand sides as it goes, so the
// projection Qᵀ·b comes out without Q ever being formed: the sweep costs
// O(m·n²) for R and O(m·n·k) for the right-hand sides, where forming Q
// costs O(m²·n) and the dense product after it another O(m²·k). For the
// tall systems this serves, m ≫ n, that product is nearly the whole
// solve, and applying the reflectors is also the more accurate route:
// the projection stays orthogonal by construction instead of carrying
// the rounding of an explicitly accumulated Q. It returns the leading n
// rows of the transformed right-hand sides, concatenated by column, and
// R.
func qrSolveRows(aMat []float64, m, n, k int, rhs []float64) (qtB, rMat []float64) {
rMat = make([]float64, m*n)
copy(rMat, aMat)
y := make([]float64, m*k)
copy(y, rhs)
v := make([]float64, m)
for kk := range n {
ln, beta, ok := householderReflector(rMat, m, n, kk, v)
if !ok {
continue
}
vv := v[:ln]
// The right-hand sides first: each column's dot runs over i
// ascending and its update repeats the same operand order, and
// the columns are independent outputs, so a split over them
// cannot move a value.
if k*ln >= qrRHSMinWork {
engine.ParallelMin(k, 1, func(start, end int) {
for j := start; j < end; j++ {
applyReflectorToColumn(y, j, k, kk, ln, beta, vv)
}
})
} else {
for j := range k {
applyReflectorToColumn(y, j, k, kk, ln, beta, vv)
}
}
// Then the trailing columns of R, blocked and dispatched exactly
// as the explicit-Q sweep does it, with the Q columns it also
// carries left out.
rBlocks := (n - kk + qrColumnBlock - 1) / qrColumnBlock
worker := func(gi, nw int) {
var d [qrColumnBlock]float64
for b := gi * rBlocks / nw; b < (gi+1)*rBlocks/nw; b++ {
j0 := kk + b*qrColumnBlock
j1 := min(j0+qrColumnBlock, n)
wd := j1 - j0
_ = d[wd-1] // bounds-check proof: the block never exceeds qrColumnBlock
clear(d[:wd])
for jj := range wd {
jcol := j0 + jj
for i := range ln {
d[jj] += vv[i] * rMat[(kk+i)*n+jcol]
}
}
for j := range wd {
d[j] = beta * d[j]
}
for i := range ln {
vi := vv[i]
base := (kk+i)*n + j0
for j := range wd {
rMat[base+j] -= vi * d[j]
}
}
}
}
weight := (n - kk) + k
w := min(weight*ln/qrWorkQuantum+1, engine.WorkersFor(rBlocks))
if w < 2 {
worker(0, 1)
continue
}
var wg sync.WaitGroup
for gi := range w {
wg.Go(func() {
worker(gi, w)
})
}
wg.Wait()
}
qtB = make([]float64, n*k)
for j := range k {
for i := range n {
qtB[i*k+j] = y[i*k+j]
}
}
return qtB, rMat
}
// applyReflectorToColumn applies one reflector to column j of an m×k
// right-hand-side block: the dot runs over i ascending and the update
// repeats the same operand order.
func applyReflectorToColumn(y []float64, j, k, kk, ln int, beta float64, vv []float64) {
dot := 0.0
base := kk * k
for i := range ln {
dot += vv[i] * y[base+i*k+j]
}
w := beta * dot
for i := range ln {
y[base+i*k+j] -= vv[i] * w
}
}
// qrRHSMinWork is the element work one worker must carry before the
// right-hand-side columns of a reflector are split across goroutines.
const qrRHSMinWork = 1 << 15
// qrWorkQuantum is the element-touch budget one worker receives in a
// qrRaw reflector dispatch, counted as the reflector's reach (columns
// of R plus columns of Q) times the reflector length. A reflector whose
// application does not fill one quantum runs on the calling goroutine:
// the dispatch would cost more than the work. A fixed 32-way split
// measured 15 000 allocations and several milliseconds of scheduler
// time per 256² QR, because 32 goroutines spawn and synchronise per
// reflector however small the reflector's reach has grown; sizing the
// crew by the work keeps each dispatch's spawn cost proportional to the
// work it spreads.
const qrWorkQuantum = 8192
// qrColumnBlock is the number of R columns one blocked update pass
// covers. R is row-major, so a single column's update walks a stride-n
// pattern that touches one cache line per element; a block of columns
// walks the same rows linearly and reuses each line across the block.
// Eight columns span one 64-byte line of float64.
const qrColumnBlock = 8
// qrRaw performs Householder QR on a flat m×n matrix and returns the
// flat Q (m, m) and R (m, n).
//
// The reflectors apply one after another: their ORDER defines Q and R
// and is never changed. Within a single reflector every column of R
// and every column of Q is updated exactly once, reading only the
// frozen reflector v and beta, so the per-column work (dot along i
// ascending, then the same-operand rank-1 subtraction) is dispatched
// over the columns. Identical inputs, identical per-element sequence:
// the result is bit-identical to the serial sweep, whichever crew size
// or column blocking the dispatcher picks.
func qrRaw(aMat []float64, m, n int) (qMat, rMat []float64, err error) {
rMat = make([]float64, m*n)
copy(rMat, aMat)
qMat = eye(m)
// One reflector scratch reused across the sweep; each step uses its
// first m−k entries.
v := make([]float64, m)
// Per-reflector state the crew reads: the sweep writes it on the
// calling goroutine and the crew only reads it, so one body closure
// and its captures live for the whole sweep instead of a fresh pair
// of closures escaping to the heap per reflector.
var (
k int // reflector's leading row and column
ln int // reflector length, m−k
beta float64 // 2/(vᵀv) of the current reflector
vv []float64 // reflector prefix view, re-sliced per step
)
// worker runs the reflector's R columns and Q columns that fall to
// crew member gi of nw: a contiguous slab of the R blocks and a
// contiguous slab of the Q columns. Every member owns its own block
// scratch d, and slabs are disjoint, so no two members ever touch
// the same line; contiguous slabs keep each member's working set at
// its own rows of qMat and its own column stripe of rMat instead of
// walking the whole window. Each column's dot runs over i ascending
// and the subtraction repeats the same operand order, so the slab
// boundaries never move a bit.
worker := func(gi, nw int) {
rBlocks := (n - k + qrColumnBlock - 1) / qrColumnBlock
// R columns in blocks of qrColumnBlock. R is row-major, so a
// lone column's update walks a stride-n pattern that touches one
// cache line per element; a block keeps the hot window at one
// line per row. The per-column dot below keeps the exact shape
// the per-column sweep always had (same accumulator shape, same
// operand order), because reshaping the dot loop changes the
// compiler's fusing of the multiply into the add and that moves
// bits. Each column's w is beta·dot as before, and the blocked
// subtraction walks the same i order over the block linearly.
var d [qrColumnBlock]float64
for b := gi * rBlocks / nw; b < (gi+1)*rBlocks/nw; b++ {
j0 := k + b*qrColumnBlock
j1 := min(j0+qrColumnBlock, n)
wd := j1 - j0
_ = d[wd-1] // bounds-check proof: the block never exceeds qrColumnBlock
clear(d[:wd])
for jj := range wd {
jcol := j0 + jj
for i := range ln {
d[jj] += vv[i] * rMat[(k+i)*n+jcol]
}
}
for j := range wd {
d[j] = beta * d[j]
}
for i := range ln {
vi := vv[i]
base := (k+i)*n + j0
for j := range wd {
rMat[base+j] -= vi * d[j]
}
}
}
// Q columns, one at a time: per column the access is already
// linear in i (row-major, row col, columns k..m−1), and a slab
// of consecutive columns is a slab of consecutive rows.
for col := gi * m / nw; col < (gi+1)*m/nw; col++ {
dot := 0.0
base := col*m + k
for i := range ln {
dot += vv[i] * qMat[base+i]
}
w := beta * dot
for i := range ln {
qMat[base+i] -= vv[i] * w
}
}
}
for k = range n {
// The reflector's construction and its scaling live in
// householderReflector, shared with the explicit-Q-free solve
// below.
var ok bool
ln, beta, ok = householderReflector(rMat, m, n, k, v)
if !ok {
continue
}
vv = v[:ln]
// Columns k..n-1 of R (in blocks, one item per block) and all m
// columns of Q: disjoint writes, read-only v and beta, so they
// can run as one parallel crew. The work weight counts a block
// as the qrColumnBlock columns it covers, which keeps the crew
// size honest about R's share.
items := (n-k+qrColumnBlock-1)/qrColumnBlock + m
weight := (n - k) + m
w := min(weight*ln/qrWorkQuantum+1, engine.WorkersFor(items))
if w < 2 {
worker(0, 1)
continue
}
var wg sync.WaitGroup
for gi := range w {
wg.Go(func() {
worker(gi, w)
})
}
wg.Wait()
}
return qMat, rMat, nil
}
// denseFloats copies a 2-D array's payload into a flat m×n float64
// matrix, converting from any numeric dtype via floatAt. Contiguous
// float64 payloads take a direct-copy fast path; floatAt gives the
// identical values for every dtype and layout either way.
func denseFloats(a *core.Array, m, n int) []float64 {
out := make([]float64, m*n)
if a.Dtype() == core.Float && !a.Strided() && len(a.RawFloats()) == m*n {
copy(out, a.RawFloats())
return out
}
for i := range m {
for j := range n {
out[i*n+j] = a.FloatAt(i*n + j)
}
}
return out
}
// floatsToArray wraps a flat m×n float64 matrix as a core.Array.
func floatsToArray(data []float64, shape []int) *core.Array {
out := core.New(core.Float, shape...)
copy(out.RawFloats(), data)
return out
}
func eye(n int) []float64 {
out := make([]float64, n*n)
for i := range n {
out[i*n+i] = 1
}
return out
}
func sign(x float64) float64 {
switch {
case x > 0:
return 1
case x < 0:
return -1
default:
return 0
}
}