Files
tensor/linalg/sparsecsr.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

212 lines
7.5 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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) }