310 lines
11 KiB
Go
310 lines
11 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
|
|
// SPDX-License-Identifier: MIT
|
|||
|
|
|
|||
|
|
package linalg
|
|||
|
|
|
|||
|
|
import (
|
|||
|
|
"slices"
|
|||
|
|
"sync"
|
|||
|
|
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/engine"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// The compressed sparse column view of a COO matrix: the format direct
|
|||
|
|
// factorisations run on, where a column at a time is eliminated and
|
|||
|
|
// the fill of one column extends the entries below the diagonal.
|
|||
|
|
type SparseCSC struct {
|
|||
|
|
ColStart []int
|
|||
|
|
RowIdx []int
|
|||
|
|
Values []float64
|
|||
|
|
Rows int
|
|||
|
|
Cols int
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// stableOrderEntriesByCol orders entries by (col, row) and returns the
|
|||
|
|
// slice holding the result, with the same stability the row-major order
|
|||
|
|
// keeps: one counting pass fixes the major key, the other the minor one,
|
|||
|
|
// and entries that share both coordinates keep their input order and
|
|||
|
|
// accumulate in it.
|
|||
|
|
func stableOrderEntriesByCol(entries []cooEntry, rows, cols int) []cooEntry {
|
|||
|
|
buf := make([]cooEntry, len(entries))
|
|||
|
|
count := make([]int, max(rows, cols)+1)
|
|||
|
|
countingPlace(buf, entries, count, func(e cooEntry) int { return e.row })
|
|||
|
|
countingPlace(entries, buf, count, func(e cooEntry) int { return e.col })
|
|||
|
|
return entries
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// CSCFromCOO converts a core.SparseCOO to CSC form with the same
|
|||
|
|
// canonicalisation CSRFromCOO applies: duplicate coordinates sum, the
|
|||
|
|
// way COO semantics require, explicit zeros drop, and every column's
|
|||
|
|
// row indices end up sorted and unique. core.Complex values are
|
|||
|
|
// refused.
|
|||
|
|
func CSCFromCOO(s *core.SparseCOO) (*SparseCSC, error) {
|
|||
|
|
if s.Values.Dtype() == core.Complex {
|
|||
|
|
return nil, base.Errf("CSCFromCOO: complex sparse matrices are not supported")
|
|||
|
|
}
|
|||
|
|
if len(s.Shape) != 2 {
|
|||
|
|
return nil, base.Errf("CSCFromCOO: needs a 2-D matrix, got rank %d", len(s.Shape))
|
|||
|
|
}
|
|||
|
|
rows, cols := s.Shape[0], s.Shape[1]
|
|||
|
|
nnz := s.Indices.Shape()[0]
|
|||
|
|
idx := s.Indices.RawInts()
|
|||
|
|
vals := s.Values.RawFloats()
|
|||
|
|
payload := s.Values.Dtype() == core.Float && !s.Values.Strided()
|
|||
|
|
entries := make([]cooEntry, nnz)
|
|||
|
|
for i := range nnz {
|
|||
|
|
r := int(idx[i*2])
|
|||
|
|
c := int(idx[i*2+1])
|
|||
|
|
if r < 0 || r >= rows || c < 0 || c >= cols {
|
|||
|
|
return nil, base.Errf("CSCFromCOO: index [%d,%d] out of range for %d×%d", r, c, rows, cols)
|
|||
|
|
}
|
|||
|
|
v := 0.0
|
|||
|
|
if payload {
|
|||
|
|
v = vals[i]
|
|||
|
|
} else {
|
|||
|
|
v = s.Values.FloatAt(i)
|
|||
|
|
}
|
|||
|
|
entries[i] = cooEntry{r, c, v}
|
|||
|
|
}
|
|||
|
|
// Sort by (col, row) so duplicates merge and columns are
|
|||
|
|
// contiguous, keeping equal coordinates in COO order.
|
|||
|
|
sorted := stableOrderEntriesByCol(entries, rows, cols)
|
|||
|
|
csc := &SparseCSC{Rows: rows, Cols: cols}
|
|||
|
|
csc.ColStart = make([]int, cols+1)
|
|||
|
|
csc.RowIdx = make([]int, 0, nnz)
|
|||
|
|
csc.Values = make([]float64, 0, nnz)
|
|||
|
|
// Merge duplicates before the zero drop, exactly as CSRFromCOO
|
|||
|
|
// argues: dropping first would let a later duplicate of a dropped
|
|||
|
|
// coordinate accumulate onto an unrelated slot.
|
|||
|
|
for p := 0; p < len(sorted); {
|
|||
|
|
e := sorted[p]
|
|||
|
|
v := e.val
|
|||
|
|
p++
|
|||
|
|
for p < len(sorted) && sorted[p].row == e.row && sorted[p].col == e.col {
|
|||
|
|
v += sorted[p].val
|
|||
|
|
p++
|
|||
|
|
}
|
|||
|
|
if v == 0 {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
csc.RowIdx = append(csc.RowIdx, e.row)
|
|||
|
|
csc.Values = append(csc.Values, v)
|
|||
|
|
csc.ColStart[e.col+1]++
|
|||
|
|
}
|
|||
|
|
for i := range cols {
|
|||
|
|
csc.ColStart[i+1] += csc.ColStart[i]
|
|||
|
|
}
|
|||
|
|
return csc, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// ToCSR returns the same matrix in CSR form, built by one counting
|
|||
|
|
// pass over the stored entries. (The Cholesky lower-triangle row view
|
|||
|
|
// is lowerRows, not this conversion.)
|
|||
|
|
func (c *SparseCSC) ToCSR() (*SparseCSR, error) {
|
|||
|
|
csr := &SparseCSR{Rows: c.Rows, Cols: c.Cols, RowStart: make([]int, c.Rows+1)}
|
|||
|
|
csr.ColIdx = make([]int, len(c.Values))
|
|||
|
|
csr.Values = make([]float64, len(c.Values))
|
|||
|
|
for _, r := range c.RowIdx {
|
|||
|
|
if r < 0 || r >= c.Rows {
|
|||
|
|
return nil, base.Errf("ToCSR: row index %d out of range for %d rows", r, c.Rows)
|
|||
|
|
}
|
|||
|
|
csr.RowStart[r+1]++
|
|||
|
|
}
|
|||
|
|
for i := range c.Rows {
|
|||
|
|
csr.RowStart[i+1] += csr.RowStart[i]
|
|||
|
|
}
|
|||
|
|
// next holds the insertion point of each row while the columns
|
|||
|
|
// are walked in order, which leaves every row's entries sorted by
|
|||
|
|
// column.
|
|||
|
|
next := make([]int, c.Rows)
|
|||
|
|
copy(next, csr.RowStart[:c.Rows])
|
|||
|
|
for j := range c.Cols {
|
|||
|
|
for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ {
|
|||
|
|
r := c.RowIdx[p]
|
|||
|
|
q := next[r]
|
|||
|
|
csr.ColIdx[q] = j
|
|||
|
|
csr.Values[q] = c.Values[p]
|
|||
|
|
next[r] = q + 1
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return csr, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// ToCSC returns the same matrix in CSC form, built by one counting
|
|||
|
|
// pass over the stored entries.
|
|||
|
|
func (c *SparseCSR) ToCSC() (*SparseCSC, error) {
|
|||
|
|
csc := &SparseCSC{Rows: c.Rows, Cols: c.Cols, ColStart: make([]int, c.Cols+1)}
|
|||
|
|
csc.RowIdx = make([]int, len(c.Values))
|
|||
|
|
csc.Values = make([]float64, len(c.Values))
|
|||
|
|
for _, j := range c.ColIdx {
|
|||
|
|
if j < 0 || j >= c.Cols {
|
|||
|
|
return nil, base.Errf("ToCSC: column index %d out of range for %d columns", j, c.Cols)
|
|||
|
|
}
|
|||
|
|
csc.ColStart[j+1]++
|
|||
|
|
}
|
|||
|
|
for i := range c.Cols {
|
|||
|
|
csc.ColStart[i+1] += csc.ColStart[i]
|
|||
|
|
}
|
|||
|
|
// next holds the insertion point of each column while the rows
|
|||
|
|
// are walked in order, which leaves every column's entries sorted
|
|||
|
|
// by row.
|
|||
|
|
next := make([]int, c.Cols)
|
|||
|
|
copy(next, csc.ColStart[:c.Cols])
|
|||
|
|
for i := range c.Rows {
|
|||
|
|
for p := c.RowStart[i]; p < c.RowStart[i+1]; p++ {
|
|||
|
|
j := c.ColIdx[p]
|
|||
|
|
q := next[j]
|
|||
|
|
csc.RowIdx[q] = i
|
|||
|
|
csc.Values[q] = c.Values[p]
|
|||
|
|
next[j] = q + 1
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return csc, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// cscMatVecColumns accumulates y += A[:,s:e]·x for the columns in
|
|||
|
|
// [s, e). It is the whole body of the product on the calling
|
|||
|
|
// goroutine: the scatter walks the columns in ascending order, which
|
|||
|
|
// is the reduction order every row of the output sees.
|
|||
|
|
func cscMatVecColumns(s, e int, colStart, rowIdx []int, values, x, y []float64) {
|
|||
|
|
for j := s; j < e; j++ {
|
|||
|
|
xj := x[j]
|
|||
|
|
p0, p1 := colStart[j], colStart[j+1]
|
|||
|
|
rv := rowIdx[p0:p1:p1]
|
|||
|
|
vv := values[p0:p1:p1]
|
|||
|
|
for i, r := range rv {
|
|||
|
|
y[r] += vv[i] * xj
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// cscMatVecRowBlock computes y[s:e] = (A·x)[s:e] by walking every
|
|||
|
|
// column and storing only the entries whose row lands in [s, e). The
|
|||
|
|
// canonical form keeps a column's rows ascending, so one pair of end
|
|||
|
|
// rows decides whether the column can reach the block at all, and a
|
|||
|
|
// column that cannot costs two loads instead of its stored entries. The
|
|||
|
|
// reads are paid once per block, but a row's contributions still arrive
|
|||
|
|
// in ascending column order, exactly the order the serial walk gives
|
|||
|
|
// them, so the result is bit-identical to the serial product.
|
|||
|
|
func cscMatVecRowBlock(s, e int, colStart, rowIdx []int, values, x, y []float64) {
|
|||
|
|
for j := range len(colStart) - 1 {
|
|||
|
|
p0, p1 := colStart[j], colStart[j+1]
|
|||
|
|
if p0 == p1 {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
if rowIdx[p1-1] < s || rowIdx[p0] >= e {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
// The column's rows ascend, so the block's share of the column
|
|||
|
|
// is one contiguous span; the two binary searches find it and
|
|||
|
|
// the walk below carries no per-entry window test.
|
|||
|
|
seg := rowIdx[p0:p1:p1]
|
|||
|
|
lo, _ := slices.BinarySearch(seg, s)
|
|||
|
|
hi, _ := slices.BinarySearch(seg, e)
|
|||
|
|
xj := x[j]
|
|||
|
|
rv := rowIdx[p0+lo : p0+hi : p0+hi]
|
|||
|
|
vv := values[p0+lo : p0+hi : p0+hi]
|
|||
|
|
for i, r := range rv {
|
|||
|
|
y[r] += vv[i] * xj
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// cscMatVecMaxBlocks caps how many row blocks the split uses. The
|
|||
|
|
// split's cost per block is a full sweep over the column boundaries
|
|||
|
|
// whatever the block's share of entries, so blocks past four pay more
|
|||
|
|
// in sweeps than they win in parallel entries: measured on the banded
|
|||
|
|
// benchmark at 131,072 rows, four blocks hold the best time and
|
|||
|
|
// thirty-two regress past the serial product.
|
|||
|
|
const cscMatVecMaxBlocks = 4
|
|||
|
|
|
|||
|
|
// cscMatVecMaxWidth is the stored entries per column past which the
|
|||
|
|
// row split stays on the calling goroutine. A block's sweep reads every
|
|||
|
|
// column's boundaries, so a wide column's structure traffic repeats per
|
|||
|
|
// block while only its in-block entries turn into work: past about
|
|||
|
|
// sixteen entries per column the sweeps outweigh the parallel entries,
|
|||
|
|
// measured between a thirteen-wide band, which still wins, and a
|
|||
|
|
// seventeen-wide one, which does not.
|
|||
|
|
const cscMatVecMaxWidth = 16
|
|||
|
|
|
|||
|
|
// cscMatVecBlocks returns the row-block count the CSC matrix-vector
|
|||
|
|
// product splits into, one for the calling goroutine. The per-worker
|
|||
|
|
// share of stored entries must reach the work budget the CSR product
|
|||
|
|
// applies, the columns must stay narrow enough for the boundary sweeps
|
|||
|
|
// to pay, and more than one block must be available to split into.
|
|||
|
|
func cscMatVecBlocks(rows, cols, nnz int) int {
|
|||
|
|
w := min(engine.WorkersFor(rows), cscMatVecMaxBlocks)
|
|||
|
|
if w > 1 && nnz/w >= sparseMatVecWorkBudget && nnz <= cscMatVecMaxWidth*cols {
|
|||
|
|
return w
|
|||
|
|
}
|
|||
|
|
return 1
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// cscRowsAscending reports whether every column's row indices ascend,
|
|||
|
|
// the property the block walk's column skip and early exit rely on.
|
|||
|
|
// The constructors produce it, but the fields are exported and can
|
|||
|
|
// carry anything, so the split verifies it with one pass.
|
|||
|
|
func cscRowsAscending(colStart, rowIdx []int) bool {
|
|||
|
|
for j := range len(colStart) - 1 {
|
|||
|
|
for p := colStart[j] + 1; p < colStart[j+1]; p++ {
|
|||
|
|
if rowIdx[p] < rowIdx[p-1] {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// MatVec computes y = A·x for a dense vector x of length Cols. Rows of
|
|||
|
|
// the output are independent, so the row range splits across workers
|
|||
|
|
// once the split pays, and the block walk visits a row's contributions
|
|||
|
|
// in the same ascending column order the column walk gives them, which
|
|||
|
|
// keeps the split bit-identical to the serial product. The split also
|
|||
|
|
// verifies the canonical form with one pass, because its column skip
|
|||
|
|
// and early exit read a column's first and last row as its bounds: a
|
|||
|
|
// matrix whose columns do not ascend falls back to the serial product,
|
|||
|
|
// which is order-blind and answers it exactly as the previous release
|
|||
|
|
// did. Below the split the product runs on the calling goroutine,
|
|||
|
|
// column by column in a fixed order, which keeps the reduction order
|
|||
|
|
// deterministic. The constructors (CSCFromCOO, ToCSC, the factorisation
|
|||
|
|
// patterns) all produce the canonical form.
|
|||
|
|
func (c *SparseCSC) MatVec(x *core.Array) (*core.Array, error) {
|
|||
|
|
if x.Dtype() == core.Complex {
|
|||
|
|
return nil, base.Errf("MatVec: complex vectors are not supported")
|
|||
|
|
}
|
|||
|
|
if x.Len() != c.Cols {
|
|||
|
|
return nil, base.Errf("MatVec: vector length %d does not match %d columns", x.Len(), c.Cols)
|
|||
|
|
}
|
|||
|
|
out := core.New(core.Float, []int{c.Rows}...)
|
|||
|
|
xf := contiguousF64(x)
|
|||
|
|
yv := out.RawFloats()
|
|||
|
|
// The three structure slices are taken once: the scatter walks them
|
|||
|
|
// per stored entry, and the loop below holds no other state.
|
|||
|
|
colStart := c.ColStart
|
|||
|
|
rowIdx := c.RowIdx
|
|||
|
|
values := c.Values
|
|||
|
|
w := cscMatVecBlocks(c.Rows, c.Cols, len(values))
|
|||
|
|
if w > 1 && !cscRowsAscending(colStart, rowIdx) {
|
|||
|
|
w = 1
|
|||
|
|
}
|
|||
|
|
if w > 1 {
|
|||
|
|
chunk := (c.Rows + w - 1) / w
|
|||
|
|
var wg sync.WaitGroup
|
|||
|
|
for s := 0; s < c.Rows; s += chunk {
|
|||
|
|
e := min(s+chunk, c.Rows)
|
|||
|
|
wg.Go(func() {
|
|||
|
|
cscMatVecRowBlock(s, e, colStart, rowIdx, values, xf, yv)
|
|||
|
|
})
|
|||
|
|
}
|
|||
|
|
wg.Wait()
|
|||
|
|
return out, nil
|
|||
|
|
}
|
|||
|
|
cscMatVecColumns(0, c.Cols, colStart, rowIdx, values, xf, yv)
|
|||
|
|
return out, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// NNZ returns the count of stored non-zeros.
|
|||
|
|
func (c *SparseCSC) NNZ() int { return len(c.Values) }
|