Files
tensor/linalg/sparsecsr.go
T

212 lines
7.5 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 (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"sourcedock.dev/petrbalvin/tensor/internal/engine"
)
// The compressed sparse row view of a COO matrix: the format the
// iterative solvers and the Lanczos eigensolver run on.
type SparseCSR struct {
RowStart []int
ColIdx []int
Values []float64
Rows int
Cols int
}
// contiguousF64 returns an array's values as a float64 slice with no
// dtype switch and no strides test per element: the payload itself when
// the array is a contiguous float64 one, and a widening copy otherwise.
// Either way the values are the ones FloatAt returns.
func contiguousF64(a *core.Array) []float64 {
if a.Dtype() == core.Float && !a.Strided() {
return a.RawFloats()[:a.Len()]
}
out := make([]float64, a.Len())
for i := range out {
out[i] = a.FloatAt(i)
}
return out
}
// sparseMatVecWorkBudget is the number of stored multiply-adds one
// worker must carry before a row split pays for the goroutine that
// carries it. A stored entry costs a gathered load, a multiply and an
// add, which measures a little under a nanosecond on the banded sweep,
// so a worker's start-up is worth about this many of them. Measured on
// that sweep: 12,286 stored entries split across the worker count
// measure 10 µs against the calling goroutine's 10 µs, 24,574 measure
// 20 µs against 28 µs, and 393,214 measure 126 µs against 441 µs.
const sparseMatVecWorkBudget = 1 << 9
// sparseMatVecSplit reports whether a sparse matrix-vector product with
// the given shape spreads its rows across the workers. The fan-out is
// worth it only once a worker's share of the stored entries reaches the
// work budget, and the worker count bounds that share, so a matrix whose
// product misses the budget stays on the calling goroutine.
func sparseMatVecSplit(rows, nnz int) bool {
w := engine.WorkersFor(rows)
return w > 1 && nnz/w >= sparseMatVecWorkBudget
}
// csrMatVecRange computes y[i] = A[i,:]·x for the rows in [s, e) of a
// matrix in compressed row form. It is the whole body of the product, so
// the serial path calls it directly and the parallel path calls it per
// chunk: taking the structure slices as arguments rather than as
// captured variables keeps the serial path free of the closure a row
// split needs, which the solvers call once per iteration.
func csrMatVecRange(s, e int, rowStart, colIdx []int, values, x, y []float64) {
for i := s; i < e; i++ {
rs, re := rowStart[i], rowStart[i+1]
vs := values[rs:re:re]
cs := colIdx[rs:re:re]
sum := 0.0
for p := range vs {
sum += vs[p] * x[cs[p]]
}
y[i] = sum
}
}
// cooEntry is one coordinate entry on the way from COO into a
// compressed form.
type cooEntry struct {
row, col int
val float64
}
// countingPlace stably places src into dst ordered by key: the key
// function returns the sort key of an entry, count is scratch holding at
// least the largest key plus two entries, and entries that share a key
// keep the order they had in src, because each one is appended to the
// end of its key's run as the input is walked.
func countingPlace[E any](dst, src []E, count []int, key func(E) int) {
clear(count)
for i := range src {
count[key(src[i])+1]++
}
for i := range len(count) - 1 {
count[i+1] += count[i]
}
for i := range src {
k := key(src[i])
dst[count[k]] = src[i]
count[k]++
}
}
// stableOrderEntries orders entries by (row, col) and returns the slice
// holding the result. Two stable counting passes reach the order a
// stable comparison sort by (row, col) reaches, in linear time: the
// first places by column, the second by row, so entries that share both
// coordinates keep their input order and their values accumulate in it.
func stableOrderEntries(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.col })
countingPlace(entries, buf, count, func(e cooEntry) int { return e.row })
return entries
}
// CSRFromCOO converts a core.SparseCOO to CSR form, summing duplicate
// coordinates the way COO semantics require and dropping explicit
// zeros. core.Complex values are refused.
func CSRFromCOO(s *core.SparseCOO) (*SparseCSR, error) {
if s.Values.Dtype() == core.Complex {
return nil, base.Errf("CSRFromCOO: complex sparse matrices are not supported")
}
if len(s.Shape) != 2 {
return nil, base.Errf("CSRFromCOO: 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("CSRFromCOO: 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 (row, col) so duplicates merge and rows are contiguous,
// keeping equal coordinates in COO order.
sorted := stableOrderEntries(entries, rows, cols)
csr := &SparseCSR{Rows: rows, Cols: cols}
csr.RowStart = make([]int, rows+1)
csr.ColIdx = make([]int, 0, nnz)
csr.Values = make([]float64, 0, nnz)
// Merge duplicates first (already adjacent after the sort) so the CSR
// is canonical: one entry per (row, col), NNZ counts unique
// coordinates, and downstream consumers may assume sorted, unique
// column indices per row. Merging comes before the zero drop on
// purpose: dropping first would let a later duplicate of a dropped
// coordinate accumulate onto whatever slot happens to sit last in
// Values, which is a different coordinate, or no slot at all.
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
}
csr.ColIdx = append(csr.ColIdx, e.col)
csr.Values = append(csr.Values, v)
csr.RowStart[e.row+1]++
}
for i := range rows {
csr.RowStart[i+1] += csr.RowStart[i]
}
return csr, nil
}
// MatVec computes y = A·x for a dense vector x of length Cols.
// Output rows are independent, so the row range splits across
// workers once a worker's share of the stored entries pays for the
// split, and runs on the calling goroutine below that.
func (c *SparseCSR) 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: every worker's rows walk
// them per stored entry, and the row loop holds no other state.
rowStart := c.RowStart
colIdx := c.ColIdx
values := c.Values
if sparseMatVecSplit(c.Rows, len(values)) {
engine.ParallelMin(c.Rows, 1, func(s, e int) {
csrMatVecRange(s, e, rowStart, colIdx, values, xf, yv)
})
return out, nil
}
csrMatVecRange(0, c.Rows, rowStart, colIdx, values, xf, yv)
return out, nil
}
// NNZ returns the count of stored non-zeros.
func (c *SparseCSR) NNZ() int { return len(c.Values) }