Files
tensor/linalg/sparsecholesky.go
T

525 lines
17 KiB
Go
Raw 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 (
"cmp"
"math"
"slices"
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The sparse Cholesky factorisation: the elimination tree (Liu's
// construction with path compression) discovers the fill of every row
// cheaply, and the row-oriented elimination that follows is the
// up-looking form of the direct-methods standard (George and Liu;
// Davis). The factor exists only for symmetric positive definite
// matrices, where it is the direct solver of choice: no pivoting, no
// fill beyond the pattern the ordering decides.
// SparseOrdering selects the fill-reducing permutation a sparse
// factorisation applies before it eliminates.
type SparseOrdering int
const (
// SparseOrderingNatural eliminates rows and columns in stored
// order. Only right for matrices whose pattern is already dense
// near the diagonal, and the reference point every ordering is
// measured against.
SparseOrderingNatural SparseOrdering = iota
// SparseOrderingReverseCuthillMcKee orders every component in
// reverse breadth first order from a pseudo-peripheral start: the
// classic choice for mesh-shaped patterns, where it keeps the
// factor banded and the fill small.
SparseOrderingReverseCuthillMcKee
// SparseOrderingMinimumDegree eliminates the vertex with the
// fewest remaining neighbours at every step, absorbing each
// eliminated pattern into its neighbours: the stronger choice on
// irregular patterns, where a bandwidth order leaves the factor
// dense in the middle.
SparseOrderingMinimumDegree
)
// rowValuePair carries a factor entry through a column sort: the row
// index travels with its value.
type rowValuePair struct {
row int
val float64
}
// SparseCholesky carries the factorisation of a symmetric positive
// definite matrix: L with A permuted to the elimination order, so
// that P·A·Pᵀ = L·Lᵀ. One factorisation solves any number of right
// hand sides.
type SparseCholesky struct {
// perm[k] is the original index of the row eliminated k-th;
// inversePerm maps an original index back to its position.
perm []int
inversePerm []int
n int
// L in CSC form: column j holds the off-diagonal rows i > j in
// ascending order, and diag carries L[j][j].
colStart []int
rowIdx []int
values []float64
diag []float64
}
// NewSparseCholesky factors a symmetric positive definite matrix. The
// lower triangle defines the matrix: a stored upper triangle entry is
// refused unless its lower counterpart is stored with exactly the
// same value, so a silently asymmetric input cannot be factored. The
// ordering argument picks the fill-reducing permutation; see
// SparseOrdering.
func NewSparseCholesky(a *core.SparseCOO, ordering SparseOrdering) (*SparseCholesky, error) {
const name = "NewSparseCholesky"
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++ {
i := c.RowIdx[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, i, j)
}
// The lower triangle is the authority: an upper entry
// must mirror a stored lower entry bit for bit.
if i < j {
q, ok := findEntry(c, j, i)
if !ok || c.Values[q] != v {
return nil, base.Errf("%s: entry [%d,%d] = %g has no matching lower counterpart [%d,%d]",
name, i, j, v, j, i)
}
}
}
}
var perm []int
switch ordering {
case SparseOrderingNatural:
perm = make([]int, c.Cols)
for i := range perm {
perm[i] = i
}
case SparseOrderingReverseCuthillMcKee:
perm, err = reverseCuthillMcKee(c)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
case SparseOrderingMinimumDegree:
perm, err = minimumDegree(c)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
default:
return nil, base.Errf("%s: unknown ordering %d", name, int(ordering))
}
f := &SparseCholesky{perm: perm, n: c.Cols}
f.inversePerm = make([]int, c.Cols)
for k, i := range perm {
f.inversePerm[i] = k
}
lower := permutedLower(c, perm)
rows := lowerRows(lower)
if err := f.factor(lower, rows); err != nil {
return nil, base.Errf("%s: %w", name, err)
}
return f, nil
}
// lowerRows returns the row view of a lower triangular matrix: row k
// holds its columns in ascending order, diagonal included. The full
// transpose would carry the upper entries into the rows, which the
// elimination must not see.
func lowerRows(c *SparseCSC) *SparseCSR {
rows := &SparseCSR{Rows: c.Cols, Cols: c.Cols, RowStart: make([]int, c.Cols+1)}
for j := range c.Cols {
for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ {
if c.RowIdx[p] >= j {
rows.RowStart[c.RowIdx[p]+1]++
}
}
}
for i := range c.Cols {
rows.RowStart[i+1] += rows.RowStart[i]
}
rows.ColIdx = make([]int, rows.RowStart[c.Cols])
rows.Values = make([]float64, rows.RowStart[c.Cols])
next := make([]int, c.Cols)
copy(next, rows.RowStart[:c.Cols])
for j := range c.Cols {
for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ {
i := c.RowIdx[p]
if i < j {
continue
}
q := next[i]
rows.ColIdx[q] = j
rows.Values[q] = c.Values[p]
next[i] = q + 1
}
}
return rows
}
// findEntry looks a stored entry (i, j) up by binary search; the
// canonical form keeps every column's row indices sorted.
func findEntry(c *SparseCSC, i, j int) (int, bool) {
at, found := slices.BinarySearch(c.RowIdx[c.ColStart[j]:c.ColStart[j+1]], i)
if !found {
return 0, false
}
return c.ColStart[j] + at, true
}
// permutedLower builds the lower triangle of P·A·Pᵀ in CSC form: new
// position k carries the original row perm[k]. Permuting does not
// keep the triangle entry for entry, so every stored lower entry (i,
// j) is placed at the unordered pair of new positions: row max of
// inversePerm[i] and inversePerm[j], column min. The symmetry check
// in NewSparseCholesky has already pinned each upper entry to its
// lower counterpart, so walking the lower triangle alone visits
// every pair once.
func permutedLower(c *SparseCSC, perm []int) *SparseCSC {
inversePerm := make([]int, c.Cols)
for k, i := range perm {
inversePerm[i] = k
}
counts := make([]int, c.Cols+1)
for j := range c.Cols {
for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ {
i := c.RowIdx[p]
if i < j {
continue
}
// The entry lands below the new diagonal: row max of the
// two new positions, column min.
b := min(inversePerm[j], inversePerm[i])
counts[b+1]++
}
}
for i := range c.Cols {
counts[i+1] += counts[i]
}
out := &SparseCSC{Rows: c.Cols, Cols: c.Cols, ColStart: counts}
out.RowIdx = make([]int, counts[c.Cols])
out.Values = make([]float64, counts[c.Cols])
next := make([]int, c.Cols)
copy(next, counts[:c.Cols])
for j := range c.Cols {
for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ {
i := c.RowIdx[p]
if i < j {
continue
}
ki, kj := inversePerm[i], inversePerm[j]
b := ki
r := kj
if kj < ki {
b = kj
r = ki
}
q := next[b]
out.RowIdx[q] = r
out.Values[q] = c.Values[p]
next[b] = q + 1
}
}
// Each new column receives entries from several old columns, so
// the rows must be sorted into the canonical ascending order,
// values travelling with their rows. One scratch buffer serves
// every column: it grows to the longest unsorted column once
// instead of allocating per column. The row indices within a
// column are unique, so the sorted placement is unique and the
// scratch cannot move a value.
pairsScratch := make([]rowValuePair, 0, c.Cols)
for b := range c.Cols {
segment := out.RowIdx[counts[b]:counts[b+1]]
if slices.IsSorted(segment) {
continue
}
pairs := pairsScratch[:0]
for p := counts[b]; p < counts[b+1]; p++ {
pairs = append(pairs, rowValuePair{out.RowIdx[p], out.Values[p]})
}
slices.SortFunc(pairs, func(x, y rowValuePair) int { return cmp.Compare(x.row, y.row) })
for p, e := range pairs {
out.RowIdx[counts[b]+p] = e.row
out.Values[counts[b]+p] = e.val
}
pairsScratch = pairs
}
return out
}
// eliminationTree returns the parent of every column in the
// elimination tree of the lower triangular pattern, Liu's
// construction with path compression: for every stored entry (k, j),
// j < k, the root of j's tree is hung under k, so parent[i] > i and
// the chain from any j whose row reaches k ends exactly at k.
// Unparented columns carry −1.
func eliminationTree(n int, rows *SparseCSR) []int {
parent := make([]int, n)
ancestor := make([]int, n)
for i := range ancestor {
ancestor[i] = -1
}
for k := range n {
parent[k] = -1
for p := rows.RowStart[k]; p < rows.RowStart[k+1]; p++ {
j := rows.ColIdx[p]
if j >= k {
continue
}
r := j
for r != -1 && r != k {
next := ancestor[r]
ancestor[r] = k
if next == -1 {
parent[r] = k
break
}
r = next
}
}
}
return parent
}
// rowReach collects the update columns of row k: every stored entry
// (k, j) walked up its parent chain until the walk reaches k. The
// chain can also run through positions whose working entry stays
// zero, which subtracts nothing later, so the superset costs
// arithmetic, never correctness.
//
// The result is sorted ascending, and the walk keeps it that way as
// it grows: a chain climbs, so its next vertex usually lands past the
// tail and appends, and one that lands below inserts into its sorted
// place. The set's order is part of the elimination's arithmetic: the
// columns below are processed in it, and two of them can push into
// one working entry.
func rowReach(rows *SparseCSR, parent []int, k int, mark []bool, set []int) []int {
set = rowReachUnsorted(rows, parent, k, mark, set)
if !slices.IsSorted(set) {
slices.Sort(set)
}
return set
}
// rowReachUnsorted collects the same set without ordering it. The
// counting pass reads the set's length alone, where the walk order
// costs nothing, so half the reach walks pay no sort.
func rowReachUnsorted(rows *SparseCSR, parent []int, k int, mark []bool, set []int) []int {
set = set[:0]
for p := rows.RowStart[k]; p < rows.RowStart[k+1]; p++ {
j := rows.ColIdx[p]
for j != k && j != -1 && !mark[j] {
mark[j] = true
set = append(set, j)
j = parent[j]
}
}
for _, j := range set {
mark[j] = false
}
return set
}
// cholColumnCounts returns the exact number of entries every column of
// the factor receives, by running the same reach walk the numeric
// elimination runs and counting the columns it visits. The numeric pass
// appends an entry to column j for every j in the reach of a row k
// except where a working entry has cancelled to zero, which only makes
// the real count smaller, and both passes call the same pure walk on
// the same elimination tree. The counts are therefore an exact upper
// bound: the arena below needs no growth and leaves no gap between
// columns. The walk runs unsorted here: the counts read the set's
// size, which no order changes.
//
// The result is returned in prefix-sum form: column j's entries start
// at counts[j], and counts[n] is the factor's off-diagonal entry count.
func cholColumnCounts(n int, rows *SparseCSR, parent []int, mark []bool, set []int) []int {
counts := make([]int, n+1)
for k := range n {
set = rowReachUnsorted(rows, parent, k, mark, set)
for _, j := range set {
counts[j+1]++
}
}
for j := range n {
counts[j+1] += counts[j]
}
return counts
}
// factor runs the row-oriented elimination. Row k of L is its
// working row divided through the square-rooted diagonal, the update
// columns come from row k's stored entries walked through the
// elimination tree, and each settled entry L[k,j] pushes its
// contribution down column j of the factor, the update the left-looking
// recurrence demands: L[k,i] loses L[k,j]·L[i,j] for every earlier
// column j, so the columns are walked ascending and every working
// entry has settled by the time its turn comes. The elimination
// tree's guarantee makes the walk safe: every stored entry (k, j)
// sits on j's parent chain, and a chain position whose working entry
// stays zero simply subtracts nothing.
//
// The factor is written straight into one flat arena, a contiguous
// block per column with a cursor, so the whole numeric factor costs two
// allocations and the scatter into the working row reads one base array
// instead of a slice header per column. A column's entries land in
// ascending k, the order the solver reads.
func (f *SparseCholesky) factor(lower *SparseCSC, rows *SparseCSR) error {
const name = "NewSparseCholesky"
n := lower.Cols
parent := eliminationTree(n, rows)
f.diag = make([]float64, n)
x := make([]float64, n)
set := make([]int, 0, n)
mark := make([]bool, n)
f.colStart = cholColumnCounts(n, rows, parent, mark, set)
arenaRows := make([]int, f.colStart[n])
arenaVals := make([]float64, f.colStart[n])
// next[j] is column j's cursor: where its next entry lands. It
// starts at the column's block and ends one past its last written
// entry, so next[j]−colStart[j] is the column's actual length.
next := make([]int, n)
copy(next, f.colStart[:n])
for k := range n {
// x holds row k of the working matrix: the stored entries
// minus the updates below.
rs, re := rows.RowStart[k], rows.RowStart[k+1]
ci := rows.ColIdx[rs:re:re]
cv := rows.Values[rs:re:re]
for i, c := range ci {
x[c] = cv[i]
}
set = rowReach(rows, parent, k, mark, set)
for _, j := range set {
dj := f.diag[j]
lkj := x[j] / dj
if lkj != 0 {
at := next[j]
arenaRows[at] = k
arenaVals[at] = lkj
next[j] = at + 1
}
// Push the entry down column j: the diagonal term zeroes
// x[j], the rows below carry the fill, and the k term is
// the diagonal sum of this very row. The column is walked
// through slices cut to its written range, so the loop
// carries no bound checks and no reload of the cursor
// next[j] the writes above could be thought to alias.
x[j] -= lkj * dj
cs, ce := f.colStart[j], next[j]
rw := arenaRows[cs:ce:ce]
rv := arenaVals[cs:ce:ce]
for i, r := range rw {
x[r] -= lkj * rv[i]
}
}
d := x[k]
if math.IsInf(d, 0) {
return base.Errf("%s: the pivot at [%d,%d] overflowed to %g; the elimination has left the float64 range",
name, f.perm[k], f.perm[k], d)
}
if math.IsNaN(d) || d <= 0 {
return base.Errf("%s: the pivot at [%d,%d] is %g; the matrix is not positive definite",
name, f.perm[k], f.perm[k], d)
}
f.diag[k] = math.Sqrt(d)
// Restore every x position the row touched: the update
// columns, their factor entries, and the diagonal.
for _, j := range set {
x[j] = 0
cs, ce := f.colStart[j], next[j]
for _, r := range arenaRows[cs:ce:ce] {
x[r] = 0
}
}
x[k] = 0
}
// Pack the columns: a column whose working entry cancelled along
// the way sits shorter than its block, so only the written range
// moves, and a column's entries keep their ascending-k order.
total := 0
for j := range n {
total += next[j] - f.colStart[j]
}
colStart := make([]int, n+1)
f.rowIdx = make([]int, total)
f.values = make([]float64, total)
for j := range n {
b, w := f.colStart[j], next[j]
colStart[j+1] = colStart[j] + (w - b)
copy(f.rowIdx[colStart[j]:], arenaRows[b:w])
copy(f.values[colStart[j]:], arenaVals[b:w])
}
f.colStart = colStart
return nil
}
// Solve computes x = A⁻¹·b for a dense vector b: permute, forward
// substitution with L, backward substitution with Lᵀ, undo the
// permutation.
func (f *SparseCholesky) 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.perm[k])
}
for j := range f.n {
xj := xf[j] / f.diag[j]
xf[j] = xj
for p := f.colStart[j]; p < f.colStart[j+1]; p++ {
xf[f.rowIdx[p]] -= f.values[p] * xj
}
}
for j := f.n - 1; j >= 0; j-- {
sum := xf[j]
for p := f.colStart[j]; p < f.colStart[j+1]; p++ {
sum -= f.values[p] * xf[f.rowIdx[p]]
}
xf[j] = sum / f.diag[j]
}
// z solves P·A·Pᵀ·z = P·b, so z[k] is the answer at original
// position perm[k]: the scatter walks the same side as the
// gather.
out := core.New(core.Float, []int{f.n}...)
of := out.RawFloats()
for k := range f.n {
of[f.perm[k]] = xf[k]
}
return out, nil
}
// Permutation returns the elimination order: position k holds the
// original index factored k-th, with P·A·Pᵀ = L·Lᵀ. The slice is a
// copy, so the caller cannot move the factor's state.
func (f *SparseCholesky) Permutation() []int {
return slices.Clone(f.perm)
}
// NNZ returns the count of stored non-zeros in the factor, diagonal
// included.
func (f *SparseCholesky) NNZ() int { return len(f.values) + f.n }