Files

353 lines
12 KiB
Go
Raw Permalink 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 linalg
import (
"math"
"slices"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The sparse LU factorisation with partial pivoting: the left-looking
// column elimination of Gilbert and Peierls. Column k is gathered
// into a dense working vector, the columns that update it are found
// by a depth-first search over the factor built so far, the pivot row
// is the largest working entry at or below the diagonal, and a pivot
// swap relabels the stored factor's rows in place. L is unit lower
// triangular stored by columns, U is upper triangular stored by rows;
// the asymmetry is what lets a row swap move whole U rows while L's
// row labels travel with their values.
// SparseLU carries the factorisation of a square matrix: P·A = L·U
// with a row permutation only, the columns eliminated in stored
// order. One factorisation solves any number of right hand sides.
type SparseLU struct {
// piv[k] is the original index of the row eliminated k-th.
piv []int
n int
// L by columns, unit diagonal: colRows[j] holds the rows i > j of
// column j with colVals[j] beside it, sorted ascending after the
// elimination.
colRows [][]int
colVals [][]float64
// U by rows: rowCols[k] holds the columns j > k of row k with
// rowVals[k] beside it, diag[k] the pivot itself.
rowCols [][]int
rowVals [][]float64
diag []float64
nnz int
}
// NewSparseLU factors a square matrix with partial pivoting: the
// pivot row of column k is the largest working entry at or below the
// diagonal, ties resolved by the order the working positions were
// gathered in, which the input alone fixes, so the factorisation is a
// pure function of the input. Complex, rectangular
// and non-finite inputs are refused; a zero pivot column means the
// matrix is singular and is reported.
func NewSparseLU(a *core.SparseCOO) (*SparseLU, error) {
const name = "NewSparseLU"
if a.Values.Dtype() == core.Complex {
return nil, base.Errf("%s: complex sparse matrices are not supported", name)
}
if len(a.Shape) != 2 || a.Shape[0] != a.Shape[1] {
return nil, base.Errf("%s: needs a square 2-D sparse matrix, got shape %v", name, a.Shape)
}
c, err := CSCFromCOO(a)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
for j := range c.Cols {
for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ {
v := c.Values[p]
if math.IsNaN(v) || math.IsInf(v, 0) {
return nil, base.Errf("%s: entry [%d,%d] is not finite", name, c.RowIdx[p], j)
}
}
}
f := &SparseLU{n: c.Cols}
f.piv = make([]int, c.Cols)
for i := range f.piv {
f.piv[i] = i
}
if err := f.factor(c); err != nil {
return nil, base.Errf("%s: %w", name, err)
}
return f, nil
}
// factor runs the left-looking column elimination.
func (f *SparseLU) factor(c *SparseCSC) error {
const name = "NewSparseLU"
n := c.Cols
f.colRows = make([][]int, n)
f.colVals = make([][]float64, n)
f.rowCols = make([][]int, n)
f.rowVals = make([][]float64, n)
f.diag = make([]float64, n)
x := make([]float64, n)
mark := make([]bool, n)
visited := make([]bool, n)
pattern := make([]int, 0, n)
reach := make([]int, 0, n)
// inversePiv maps an original row to the position it sits at now, and
// piv is its inverse. A pivot swap exchanges two entries of each.
//
// L's stored rows carry the ORIGINAL index of the row, not its
// position, and every read translates through inversePiv. A swap
// then costs two entries of inversePiv instead of a walk over every
// stored multiplier, which for a system that pivots often is the
// difference between a pass over the factor per swap and a pass over
// the factor per factorisation. The final labels are exactly the
// ones the position-space storage would hold: the translation is
// applied once, after the elimination, before the columns are
// sorted.
inversePiv := make([]int, n)
for i := range inversePiv {
inversePiv[i] = i
}
// The depth-first search stack: a column and how far into its
// stored entries the walk has come.
type frame struct {
col int
deg int
}
stack := make([]frame, 0, n)
for k := range n {
// x holds column k of the working matrix: the stored entries
// minus the updates the search produces below. mark flags the
// working positions; visited flags the columns the search has
// walked, which the scatter must not suppress.
pattern = pattern[:0]
for p := c.ColStart[k]; p < c.ColStart[k+1]; p++ {
i := inversePiv[c.RowIdx[p]]
if !mark[i] {
mark[i] = true
pattern = append(pattern, i)
}
x[i] = c.Values[p]
}
reach = reach[:0]
// Depth-first search over the factor's columns from the
// stored entries above the diagonal: it visits every column
// the updates flow through, marks every position they can
// reach, and records the visit post-order, whose reverse is
// the order the updates must run in.
for p := c.ColStart[k]; p < c.ColStart[k+1]; p++ {
// The seed enters in the factor's current row space: the
// stored row may have been swapped since it was written,
// and the search walks the labels as they stand now.
seed := inversePiv[c.RowIdx[p]]
if seed >= k || visited[seed] {
continue
}
visited[seed] = true
stack = stack[:0]
stack = append(stack, frame{seed, 0})
for len(stack) > 0 {
// The frame is addressed by index, never by pointer:
// an append can reallocate the stack and a held
// pointer would keep writing into the stale array.
top := len(stack) - 1
entries := f.colRows[stack[top].col]
advanced := false
for stack[top].deg < len(entries) {
// The stored row is an original index; the search
// walks positions, so it translates as it reads.
i := inversePiv[entries[stack[top].deg]]
stack[top].deg++
if !visited[i] {
visited[i] = true
if !mark[i] {
mark[i] = true
pattern = append(pattern, i)
}
if i < k {
stack = append(stack, frame{i, 0})
advanced = true
break
}
}
}
if !advanced {
reach = append(reach, stack[top].col)
stack = stack[:top]
}
}
}
for _, j := range reach {
visited[j] = false
}
// The updates run over the post-order list backwards: a
// column's own inputs are all applied before the column
// pushes its U entry down its rows.
for _, j := range slices.Backward(reach) {
t := x[j]
if t == 0 {
continue
}
// First touch of a U row capacity-sizes it once, so the
// entries that arrive across eliminations append without a
// doubling chain per row.
if cap(f.rowCols[j]) == 0 {
f.rowCols[j] = make([]int, 0, 4)
f.rowVals[j] = make([]float64, 0, 4)
}
f.rowCols[j] = append(f.rowCols[j], k)
f.rowVals[j] = append(f.rowVals[j], t)
for p, i := range f.colRows[j] {
x[inversePiv[i]] -= f.colVals[j][p] * t
}
}
// Partial pivoting: the largest working entry at or below the
// diagonal; ties stay with the position the deterministic
// traversal gathered first.
pivot := -1
best := 0.0
for _, i := range pattern {
if i < k {
continue
}
if pivot == -1 || math.Abs(x[i]) > best {
pivot = i
best = math.Abs(x[i])
}
}
if pivot == -1 || best == 0 {
return base.Errf("the column %d has no pivot; the matrix is singular", k)
}
if pivot != k {
a, b := f.piv[k], f.piv[pivot]
f.piv[k], f.piv[pivot] = b, a
inversePiv[a], inversePiv[b] = inversePiv[b], inversePiv[a]
f.rowCols[k], f.rowCols[pivot] = f.rowCols[pivot], f.rowCols[k]
f.rowVals[k], f.rowVals[pivot] = f.rowVals[pivot], f.rowVals[k]
x[k], x[pivot] = x[pivot], x[k]
}
if math.IsInf(x[k], 0) {
return base.Errf("the pivot at [%d,%d] overflowed to %g; the elimination has left the float64 range", k, k, x[k])
}
if math.IsNaN(x[k]) || x[k] == 0 {
return base.Errf("the pivot at [%d,%d] is %g; the matrix is singular", k, k, x[k])
}
f.diag[k] = x[k]
// Column k of L: the working entries below the pivot row, stored
// under the original index of the row they belong to. The column
// is capacity-sized from the working pattern once, so the fill
// appends without a growth chain.
f.colRows[k] = make([]int, 0, max(len(pattern)-k-1, 0))
f.colVals[k] = make([]float64, 0, max(len(pattern)-k-1, 0))
for _, i := range pattern {
if i > k && x[i] != 0 {
f.colRows[k] = append(f.colRows[k], f.piv[i])
f.colVals[k] = append(f.colVals[k], x[i]/f.diag[k])
}
}
// Restore the working vector and both sets of flags; every
// visited node is in the pattern, so the reset is one pass.
for _, i := range pattern {
x[i] = 0
mark[i] = false
visited[i] = false
}
x[k] = 0
}
// Translate the stored original rows into the elimination positions
// they ended at, then sort the columns into canonical ascending row
// order. Both passes are per factorisation, and their result is what
// position-space storage would have held entry for entry. One sort
// scratch serves every column.
var sortScratch []rowValuePair
for j := range n {
rows := f.colRows[j]
for q, i := range rows {
rows[q] = inversePiv[i]
}
sortScratch = f.sortColumn(j, sortScratch)
}
// U's strict triangle is stored by rows, so it escapes the column
// walk's count: add it here, keeping NNZ at L strict + U strict +
// the diagonal.
for j := range n {
f.nnz += len(f.rowCols[j])
}
return nil
}
// sortColumn orders one L column by ascending row, values travelling
// with their rows, and records the factor's non-zero count. The caller
// supplies the scratch across columns; the grown buffer comes back for
// the next column. Stored rows are unique within a column, so the
// sorted placement is unique and the scratch cannot move a value.
func (f *SparseLU) sortColumn(j int, scratch []rowValuePair) []rowValuePair {
rows := f.colRows[j]
if !slices.IsSorted(rows) {
pairs := scratch[:0]
for p, r := range rows {
pairs = append(pairs, rowValuePair{r, f.colVals[j][p]})
}
slices.SortFunc(pairs, func(x, y rowValuePair) int { return x.row - y.row })
for p, e := range pairs {
f.colRows[j][p] = e.row
f.colVals[j][p] = e.val
}
scratch = pairs
}
f.nnz += len(rows)
return scratch
}
// Solve computes x = A⁻¹·b for a dense vector b: permute, forward
// substitution with the unit lower L, backward substitution with U,
// undo the permutation.
func (f *SparseLU) Solve(b *core.Array) (*core.Array, error) {
const name = "Solve"
if b.NDim() != 1 {
return nil, base.Errf("%s: the right hand side must be rank 1", name)
}
if b.Dtype() == core.Complex {
return nil, base.Errf("%s: complex right hand sides are not supported", name)
}
if b.Len() != f.n {
return nil, base.Errf("%s: right hand side length %d does not match %d rows", name, b.Len(), f.n)
}
x := core.New(core.Float, []int{f.n}...)
xf := x.RawFloats()
for k := range f.n {
xf[k] = b.FloatAt(f.piv[k])
}
// L is unit lower triangular: the divide at the diagonal is the
// identity and each column push reads its own x[j] first.
for j := range f.n {
xj := xf[j]
for p, i := range f.colRows[j] {
xf[i] -= f.colVals[j][p] * xj
}
}
for j := f.n - 1; j >= 0; j-- {
sum := xf[j]
for p, c := range f.rowCols[j] {
sum -= f.rowVals[j][p] * xf[c]
}
xf[j] = sum / f.diag[j]
}
// A = Pᵀ·L·U, so U⁻¹·L⁻¹·P·b is the answer itself: with the
// single-sided permutation nothing is undone at the end.
return x, nil
}
// Permutation returns the row elimination order: position k holds the
// original index of the row factored k-th, with P·A = L·U. The slice
// is a copy, so the caller cannot move the factor's state.
func (f *SparseLU) Permutation() []int {
return slices.Clone(f.piv)
}
// NNZ returns the count of stored non-zeros in the factor: L's strict
// columns, U's strict rows and the diagonal of U, L's unit diagonal
// excluded.
func (f *SparseLU) NNZ() int { return f.nnz + f.n }