353 lines
12 KiB
Go
353 lines
12 KiB
Go
// 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 }
|