212 lines
7.5 KiB
Go
212 lines
7.5 KiB
Go
// 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) }
|