feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,352 @@
|
||||
// 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 }
|
||||
Reference in New Issue
Block a user