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