Files
tensor/internal/base/base.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

429 lines
14 KiB
Go
Raw 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 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
}