// 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" ) 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 }