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