Files
tensor/internal/base/base.go
T

429 lines
14 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Package base holds the low-level primitives the library's domain
// packages share: error construction, shape formatting, the machine
// epsilon and the generic LU factorisation the solvers build on. It
// touches no arrays, so every package in the module can depend on it
// without a cycle.
package base
import (
"fmt"
"math"
"math/cmplx"
"strconv"
"strings"
"sync"
"sourcedock.dev/petrbalvin/tensor/internal/engine"
)
// EpsF is the float64 machine epsilon.
const EpsF = 2.220446049250313e-16
// Errf builds a package-prefixed error: every error the library
// returns starts with "tensor: ", whatever package raised it.
func Errf(format string, args ...any) error {
return fmt.Errorf("tensor: "+format, args...)
}
// WrapErr wraps an error the library already prefixed, naming the
// operation that surfaces it. The chain stays intact for errors.Is and
// errors.As, and the message carries the tensor tag and the operation
// once each instead of the doubled prefix a plain Errf wrap of a
// prefixed error prints.
func WrapErr(operation string, err error) error {
if err == nil {
return nil
}
return &wrappedErr{op: operation, err: err}
}
type wrappedErr struct {
op string
err error
}
func (e *wrappedErr) Error() string {
return "tensor: " + e.op + ": " + strings.TrimPrefix(e.err.Error(), "tensor: ")
}
func (e *wrappedErr) Unwrap() error { return e.err }
// ShapeText renders a shape as (n, m, ...).
func ShapeText(shape []int) string {
if len(shape) == 0 {
return "scalar"
}
parts := make([]string, len(shape))
for i, d := range shape {
parts[i] = strconv.Itoa(d)
}
return "(" + strings.Join(parts, ", ") + ")"
}
// AbsComplex returns |z|, the magnitude of a complex number. It is
// math.Hypot rather than sqrt(r² + i²), which overflows above about
// 1.34e154 and underflows below about 1.5e-162 even though the true
// magnitude is an ordinary number there.
func AbsComplex(z complex128) float64 {
return math.Hypot(real(z), imag(z))
}
// RangeN returns 0..n-1.
func RangeN(n int) []int {
out := make([]int, n)
for i := range out {
out[i] = i
}
return out
}
// CmplxPolar builds a complex number from magnitude and angle.
func CmplxPolar(r, theta float64) complex128 {
return complex(r*math.Cos(theta), r*math.Sin(theta))
}
// Real2 returns |z|², the squared magnitude of a complex number.
func Real2(z complex128) float64 {
return real(z)*real(z) + imag(z)*imag(z)
}
// transposeBlock is the tile side of TransposeFlat: the destination
// column segment a tile writes spans transposeBlock consecutive
// entries of the tile's own rows, which keeps the strided side of the
// copy inside the cache instead of paying a miss per element.
const transposeBlock = 16
// TransposeFlat transposes a row-major m×n matrix held flat. The copy
// is tiled, so the strided side walks transposeBlock rows at a time
// rather than the whole matrix. A pure permutation of values, so the
// tiling cannot move a bit.
func TransposeFlat(a []float64, m, n int) []float64 {
out := make([]float64, m*n)
for i0 := 0; i0 < m; i0 += transposeBlock {
i1 := min(i0+transposeBlock, m)
for j0 := 0; j0 < n; j0 += transposeBlock {
j1 := min(j0+transposeBlock, n)
for i := i0; i < i1; i++ {
row := a[i*n+j0 : i*n+j1]
dst := j0*m + i
_ = out[dst] // bounds-check proof: the first tile entry
for _, v := range row {
out[dst] = v
dst += m
}
}
}
}
return out
}
// Scalar is the element type the generic linear algebra serves.
type Scalar interface {
float64 | float32 | complex128
}
// factorWorkQuantum is the element-update budget one worker receives in
// a rank-1 update dispatch, counted as rows remaining × columns
// remaining. A pivot whose update does not fill one quantum runs on the
// calling goroutine: the dispatch would cost more than the update. The
// crew is sized by the work, never by the machine, because the updates
// shrink every pivot and a wide dispatch of a small update costs more
// than it saves. Measured over the rank-1 update alone at the sizes the
// solvers reach, the best crew shrinks with the quantum and the curve is
// flat between 8 192 and 49 152, one size preferring a smaller crew and
// the next a larger one; this value keeps both within a few percent of
// their own optimum. The pivot SEARCH and the row SWAP stay strictly
// serial: they pick the pivot every update reads, and their order
// defines the factorisation.
const factorWorkQuantum = 24576
// factorSpawnFloor is the update size from which any crew may spawn.
// Crew SIZE above it stays tied to factorWorkQuantum; below it the whole
// update runs on the calling goroutine. The paired crossover probe
// (BenchmarkFactorCrossoverAB) measures the lone walk ahead of every
// crew through n = 384, whose largest update is 147 072, and the shipped
// crew ahead of the lone walk from n = 448, whose largest update is
// 200 256; the floor sits between the two measured updates. The crew
// size never moves a bit: every row keeps the serial per-row sequence
// whatever dispatch carries it (TestFactorDispatchBitIdentical).
const factorSpawnFloor = 147456
// factorMaxCrew caps the crew a single dispatch may use. Past this the
// per-goroutine spawn and join cost outgrows the update it spreads: an
// empty fork-join of w goroutines costs about 0.2 µs per goroutine, so a
// 512² pivot on a 16-way crew already pays more in synchronisation than
// the last worker's chunk earns.
const factorMaxCrew = 16
// factorMinRows is the floor on rows per worker for the rank-1 update;
// fewer rows per goroutine than this only adds scheduling latency.
const factorMinRows = 4
// factorJob is one crew member's slice of a pivot's rank-1 update: the
// operands every member shares and the row range this member owns. The
// jobs are allocated once per Factor call and rewritten in place per
// pivot, so a dispatch allocates the goroutine's argument frame only,
// never a closure.
type factorJob[T Scalar] struct {
rows [][]T
pivotRow []T
k, n int
start int
end int
wg *sync.WaitGroup
}
// factorUpdate runs one job through the serial per-row sequence, the
// same call the inline path makes, so the crew size never moves a bit.
func factorUpdate[T Scalar](job *factorJob[T]) {
factorRows(job.rows[job.start:job.end], job.k, job.n, job.pivotRow)
job.wg.Done()
}
// factorRows subtracts f·pivotRow from every row of rows, the rank-1
// update of one pivot: f is the row's multiplier, then the products are
// subtracted along j ascending. The whole family of entry points shares
// this one function, which is what makes the parallel result
// bit-identical to the serial one.
func factorRows[T Scalar](rows [][]T, k, n int, pivotRow []T) {
pk := pivotRow[k]
_ = pivotRow[n-1] // bounds-check proof: every pivot row spans n entries
for _, row := range rows {
f := row[k] / pk
row[k] = f
_ = row[n-1] // bounds-check proof: every row spans n entries
for j := k + 1; j < n; j++ {
row[j] -= f * pivotRow[j]
}
}
}
// Factor factors a row-major square matrix in place with partial
// pivoting, returning the row permutation and the parity of the
// permutation. The stored subdiagonal holds the elimination factors,
// so a zero pivot column is skipped and reported through the zero
// diagonal the singular check reads.
//
// The elimination is parallelised per pivot over the rows below it.
// Every row's update reads only the pivot row, which is final once the
// search and swap have run, and writes only its own row: each element
// is updated exactly once per pivot with the same operands and in the
// same order as the serial loop (f, then f·pivot[j] subtracted along
// j ascending), so the result is bit-identical for either element
// type, complex128 included.
func Factor[T Scalar](m [][]T) ([]int, int) {
n := len(m)
perm := make([]int, n)
for i := range perm {
perm[i] = i
}
parity := 1
var wg sync.WaitGroup
var jobs []factorJob[T]
for k := range n {
pivot := k
for i := k + 1; i < n; i++ {
if absOf(m[i][k]) > absOf(m[pivot][k]) {
pivot = i
}
}
if m[pivot][k] == 0 {
continue // the whole column is zero: singular from here on
}
if pivot != k {
m[pivot], m[k] = m[k], m[pivot]
perm[pivot], perm[k] = perm[k], perm[pivot]
parity = -parity
}
pivotRow := m[k]
rows := m[k+1:]
// Crew sized by the update's size, not by the machine: the
// updates shrink every pivot, and a 32-goroutine dispatch of a
// small update costs more than it saves. Below one quantum the
// update cannot fund even a two-way split, and below the spawn
// floor no crew at all earns its fork-join, so the update runs
// inline.
work := len(rows) * (n - k)
w := min(work/factorWorkQuantum+1, engine.WorkersFor(len(rows)))
w = min(w, factorMaxCrew)
if work < factorSpawnFloor {
w = 1
}
if w < 2 {
factorRows(rows, k, n, pivotRow)
continue
}
chunk := (len(rows) + w - 1) / w
if chunk < factorMinRows {
w = max(len(rows)/factorMinRows, 1)
chunk = (len(rows) + w - 1) / w
if w < 2 {
factorRows(rows, k, n, pivotRow)
continue
}
}
// Each row's update reads only the pivot row, which is final
// once the search and swap have run, and writes only its own
// row: the per-row sequence (f, then f·pivotRow[j] subtracted
// along j ascending) is the serial one, so the crew size and
// the chunk boundaries never move a bit. The crew runs over
// preallocated slots, so a dispatch allocates the goroutine's
// argument frame and nothing else.
if jobs == nil {
jobs = make([]factorJob[T], factorMaxCrew)
}
spawned := 0
for start := 0; start < len(rows); start += chunk {
job := &jobs[spawned]
job.rows, job.pivotRow, job.k, job.n = rows, pivotRow, k, n
job.start, job.end, job.wg = start, min(start+chunk, len(rows)), &wg
spawned++
}
wg.Add(spawned)
for i := range spawned {
go factorUpdate(&jobs[i])
}
wg.Wait()
}
return perm, parity
}
// CheckSingular reports an error when a factored matrix carries a
// zero pivot.
func CheckSingular[T Scalar](name string, m [][]T) error {
for i := range m {
if m[i][i] == 0 {
return Errf("%s: matrix is singular", name)
}
}
return nil
}
// SolveColumn back-substitutes one right-hand column through the
// factored matrix. It is a recurrence along the column (each col[i]
// reads every earlier entry) and is intentionally serial; the hoisted
// row slices and the index proofs only strip bounds checks, the
// per-element operation sequence is untouched.
func SolveColumn[T Scalar](m [][]T, col []T) {
n := len(m)
if n == 0 {
return
}
_ = col[n-1] // bounds-check proof: every index below stays in range
for i := 1; i < n; i++ {
row := m[i]
_ = row[i-1]
ci := col[i]
for j := range i {
ci -= row[j] * col[j]
}
col[i] = ci
}
for i := n - 1; i >= 0; i-- {
row := m[i]
_ = row[n-1]
ci := col[i]
for j := i + 1; j < n; j++ {
ci -= row[j] * col[j]
}
col[i] = ci / m[i][i]
}
}
// PermuteColumn reorders col in place so that col[i] takes the value
// that sat at perm[i]. The permutation is applied by walking the cycles
// of perm: every value moves exactly once with plain assignment, bit
// for bit, and perm itself is never written (callers reuse it across
// columns). The only allocation is the visit bitmap, not a copy of the
// column.
func PermuteColumn[T Scalar](col []T, perm []int) {
permuteColumn(col, perm, make([]bool, len(col)))
}
// permuteColumn is PermuteColumn with the caller supplying the visit
// bitmap, so a caller solving many columns can reuse one bitmap instead
// of allocating one per column. The bitmap is cleared on entry, so a
// dirty buffer behaves exactly like a fresh one.
func permuteColumn[T Scalar](col []T, perm []int, visited []bool) {
clear(visited[:len(col)])
for i := range col {
if visited[i] || perm[i] == i {
visited[i] = true
continue
}
// Carry the displaced value around the cycle.
tmp := col[i]
j := i
for {
visited[j] = true
k := perm[j]
if k == i {
break
}
col[j] = col[k]
j = k
}
col[j] = tmp
}
}
// SolveSystem factors a fresh copy-safe matrix, rejects singularity
// and solves every right-hand column in one pass: each column is
// permuted the way factoring moved the rows, then substitutes through
// L and U. m is consumed in place; callers hand over freshly built
// matrices.
//
// With several right-hand columns the substitution runs in parallel
// over the columns: a column's permutation and back-substitution read
// only the factored matrix and write only that column, so each column
// runs the identical serial sequence on identical inputs, and the
// results are bit-identical to solving the columns one after another.
// SolveColumn itself is a recurrence along the column and stays
// serial.
func SolveSystem[T Scalar](name string, m [][]T, rhs [][]T) ([][]T, error) {
perm, _ := Factor(m)
if err := CheckSingular(name, m); err != nil {
return nil, err
}
if len(rhs) > 1 {
// Each crew member owns one visit bitmap reused across its
// columns: permuteColumn clears it on entry, so the walk sees
// a fresh bitmap without an allocation per column. Bitmaps are
// never shared between crew members, so there is no race.
engine.ParallelMin(len(rhs), 1, func(start, end int) {
visited := make([]bool, len(m))
for _, col := range rhs[start:end] {
permuteColumn(col, perm, visited)
SolveColumn(m, col)
}
})
return rhs, nil
}
// The serial path reuses one bitmap across the columns for the
// same reason the parallel path does: permuteColumn clears it on
// entry, so a dirty buffer behaves exactly like a fresh one.
visited := make([]bool, len(m))
for _, col := range rhs {
permuteColumn(col, perm, visited)
SolveColumn(m, col)
}
return rhs, nil
}
func absOf[T Scalar](v T) float64 {
switch x := any(v).(type) {
case float64:
return math.Abs(x)
case float32:
return math.Abs(float64(x))
case complex128:
return cmplx.Abs(x)
}
// Every type the Scalar constraint admits is handled above; the
// clause is what tells the compiler so.
return 0
}