151 lines
4.8 KiB
Go
151 lines
4.8 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"
|
||
|
|
)
|
||
|
|
|
||
|
|
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
|
||
|
|
}
|