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