Files
tensor/linalg/sparseops.go
T

151 lines
4.8 KiB
Go
Raw Permalink 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"
)
import "slices"
// Operations on the CSR sparse matrix beyond the matrix-vector
// product: dense and sparse products, and the transpose. The sparse
// product keeps its result sparse by construction, row by row, so
// chained sparse expressions never materialise a dense intermediate
// unless the caller asks for one.
// MatMulDense returns A·x for a dense matrix x of shape (Cols, k):
// every output column is one matrix-vector product against the same
// sparse structure, so the non-zeros are streamed once for the whole
// product.
func (c *SparseCSR) MatMulDense(x *core.Array) (*core.Array, error) {
if x.Dtype() == core.Complex {
return nil, base.Errf("MatMulDense: complex operands are not supported")
}
if x.NDim() != 2 || x.Shape()[0] != c.Cols {
return nil, base.Errf("MatMulDense: operand shape %s does not match (Rows, Cols) = (%d, %d)",
base.ShapeText(x.Shape()), c.Rows, c.Cols)
}
k := x.Shape()[1]
out := core.New(core.Float, []int{c.Rows, k}...)
xc := contiguousF64(x)
// Row-by-row with all k columns in the inner loop keeps the
// dense operand's rows hot; the output payload is taken once so the
// update is a plain store.
yc := out.RawFloats()
for i := range c.Rows {
rs, re := c.RowStart[i], c.RowStart[i+1]
yi := yc[i*k : i*k+k : i*k+k]
for p := rs; p < re; p++ {
v := c.Values[p]
j := c.ColIdx[p]
xj := xc[j*k : j*k+k : j*k+k]
for m := range yi {
yi[m] += v * xj[m]
}
}
}
return out, nil
}
// MatMulSparse returns the sparse product A·B, valid when A.Cols
// equals B.Rows. Each output row gathers the products of the row's
// non-zeros with the corresponding rows of B into an accumulator
// indexed by column, so the cost is the number of stored multiply-adds
// plus one sort of the columns the row touched, and the result stores
// only what survives.
func (c *SparseCSR) MatMulSparse(o *SparseCSR) (*SparseCSR, error) {
if c.Cols != o.Rows {
return nil, base.Errf("MatMulSparse: inner dimensions disagree, %d vs %d", c.Cols, o.Rows)
}
// An output row gathers one multiply-add per stored entry of A it
// reaches, but it can store no more entries than it has columns, so
// the row's bound is the smaller of the two and their sum bounds the
// result's non-zeros, which sizes the output in one pass. The bound
// is tight for a dense product and never exceeds it.
work := 0
for i := range c.Rows {
row := 0
for p := c.RowStart[i]; p < c.RowStart[i+1]; p++ {
r := c.ColIdx[p]
row += o.RowStart[r+1] - o.RowStart[r]
}
work += min(row, o.Cols)
}
out := &SparseCSR{Rows: c.Rows, Cols: o.Cols}
out.RowStart = make([]int, c.Rows+1)
out.ColIdx = make([]int, 0, work)
out.Values = make([]float64, 0, work)
// The accumulator is dense over the columns and addressed by the
// column index itself instead of by a hash of it. touched holds the
// columns the current row has contributed to, each once, so the sort
// below orders exactly the columns that were written; inAcc marks the
// ones the list already holds, which keeps a second contribution to
// the same column from listing it twice.
acc := make([]float64, o.Cols)
inAcc := make([]bool, o.Cols)
touched := make([]int, 0, 16)
for i := range c.Rows {
touched = touched[:0]
for p := c.RowStart[i]; p < c.RowStart[i+1]; p++ {
v := c.Values[p]
r := c.ColIdx[p]
rs, re := o.RowStart[r], o.RowStart[r+1]
oCols := o.ColIdx[rs:re:re]
oVals := o.Values[rs:re:re]
for q := range oCols {
j := oCols[q]
if !inAcc[j] {
inAcc[j] = true
touched = append(touched, j)
}
acc[j] += v * oVals[q]
}
}
slices.Sort(touched)
for _, j := range touched {
if v := acc[j]; v != 0 {
out.ColIdx = append(out.ColIdx, j)
out.Values = append(out.Values, v)
out.RowStart[i+1]++
}
acc[j] = 0
inAcc[j] = false
}
}
for i := range c.Rows {
out.RowStart[i+1] += out.RowStart[i]
}
return out, nil
}
// Transpose returns the transpose in CSR form, built by the standard
// counting pass: one sweep counts the entries per output row, the
// second places them.
func (c *SparseCSR) Transpose() *SparseCSR {
out := &SparseCSR{Rows: c.Cols, Cols: c.Rows}
out.RowStart = make([]int, c.Cols+1)
for _, j := range c.ColIdx {
out.RowStart[j+1]++
}
for i := range c.Cols {
out.RowStart[i+1] += out.RowStart[i]
}
next := make([]int, c.Cols)
copy(next, out.RowStart[:c.Cols])
out.ColIdx = make([]int, len(c.ColIdx))
out.Values = make([]float64, len(c.Values))
for i := range c.Rows {
for p := c.RowStart[i]; p < c.RowStart[i+1]; p++ {
j := c.ColIdx[p]
q := next[j]
out.ColIdx[q] = i
out.Values[q] = c.Values[p]
next[j]++
}
}
return out
}