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