Files
tensor/internal/core/sparse.go
T

374 lines
11 KiB
Go
Raw 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 core
import "sourcedock.dev/petrbalvin/tensor/internal/base"
// Sparse COO arrays. Used for embedding tables, attention
// masks, recommender features: anywhere the data is sparse enough that
// a dense representation would waste memory and bandwidth.
//
// A SparseCOO stores non-zero values as (indices, values) pairs; the
// full shape is part of the metadata. Operations preserve sparsity:
// multiplying a SparseCOO by a dense array stays sparse when the
// dense side does not introduce new non-zeros.
// SparseCOO is a coordinate-format sparse array. Indices has shape
// (nnz, ndim) and dtype int; values has shape (nnz,) and the
// element dtype of the array. Shape is the dense shape.
type SparseCOO struct {
Indices *Array
Values *Array
Shape []int
}
// NewSparseCOO creates a SparseCOO from explicit indices, values, and
// shape. Returns an error if either array is nil, if indices and values
// disagree on nnz, or if indices doesn't have the right rank.
func NewSparseCOO(indices, values *Array, shape []int) (*SparseCOO, error) {
// A nil array is what a caller holds after a constructor refused its
// arguments, so it is a refusal here rather than a nil-pointer panic
// on the fields read below.
if indices == nil {
return nil, errf("NewSparseCOO: indices must not be nil")
}
if values == nil {
return nil, errf("NewSparseCOO: values must not be nil")
}
if indices.dt != Int {
return nil, errf("NewSparseCOO: indices must be int, got %s", indices.dt)
}
if indices.NDim() != 2 {
return nil, errf("NewSparseCOO: indices must be 2-D (nnz, ndim), got shape %s", shapeText(indices.shape))
}
if indices.shape[1] != len(shape) {
return nil, errf("NewSparseCOO: indices dim %d does not match shape %v", indices.shape[1], shape)
}
if values.Len() != indices.shape[0] {
return nil, errf("NewSparseCOO: values len %d does not match nnz %d", values.Len(), indices.shape[0])
}
if narrowRefused(values.dt) {
return nil, errf("NewSparseCOO: dtype %s is not supported; convert with Astype", values.dt)
}
// The dense shape is metadata the arithmetic trusts (SpMatMul
// allocates from it directly), so it gets the same validation and
// private copy every constructor applies.
if _, _, serr := checkedDims(shape); serr != nil {
return nil, base.WrapErr("NewSparseCOO", serr)
}
sh := make([]int, len(shape))
copy(sh, shape)
return &SparseCOO{Indices: indices, Values: values, Shape: sh}, nil
}
// SparseFrom extracts a SparseCOO from a dense array by keeping only
// the non-zero elements. Renamed from `SparseFromDense` to mirror
// the symmetry with `SparseCOO.Dense`. The values array keeps the
// dense array's dtype, so the round trip preserves the element type.
func SparseFrom(dense *Array) (*SparseCOO, error) {
if dense.dt == Complex {
return nil, errf("SparseFrom: complex arrays are not supported")
}
if narrowRefused(dense.dt) {
return nil, errf("SparseFrom: dtype %s is not supported; convert with Astype", dense.dt)
}
shape := dense.Shape()
var coords []int64
for flat := range dense.Len() {
if !isZero(dense, flat) {
coord := make([]int, dense.NDim())
rem := flat
for d := range dense.NDim() {
stride := 1
for k := d + 1; k < dense.NDim(); k++ {
stride *= shape[k]
}
coord[d] = rem / stride
rem %= stride
}
for d := range dense.NDim() {
coords = append(coords, int64(coord[d]))
}
}
}
nnz := len(coords) / dense.NDim()
indices, err := FromInts(coords, nnz, dense.NDim())
if err != nil {
return nil, err
}
var values *Array
switch dense.dt {
case Int:
ints := make([]int64, nnz)
k := 0
for flat := range dense.Len() {
if !isZero(dense, flat) {
ints[k] = dense.ints[flat]
k++
}
}
values, err = FromInts(ints, nnz)
case Float16:
halves := make([]uint16, nnz)
k := 0
for flat := range dense.Len() {
if !isZero(dense, flat) {
halves[k] = dense.halves[flat]
k++
}
}
values, err = HalvesFromArray(halves, nnz)
case Float32:
f32 := make([]float32, nnz)
k := 0
for flat := range dense.Len() {
if !isZero(dense, flat) {
f32[k] = dense.floats32[flat]
k++
}
}
values, err = FromFloat32s(f32, nnz)
default:
floats := make([]float64, nnz)
k := 0
for flat := range dense.Len() {
if !isZero(dense, flat) {
floats[k] = dense.floats[flat]
k++
}
}
values, err = FromFloats(floats, nnz)
}
if err != nil {
return nil, err
}
return &SparseCOO{Indices: indices, Values: values, Shape: shape}, nil
}
// consistent validates the index/value pairing of a hand-built
// SparseCOO literal, whose exported fields the constructor's checks
// never saw. Every entry point that walks the stored coordinates
// calls this first, so a mismatched pair is an error, never an
// out-of-bounds panic.
func (s *SparseCOO) consistent(name string) error {
if s.Values.Len() != s.Indices.Shape()[0] {
return errf("%s: %d values for %d index rows", name, s.Values.Len(), s.Indices.Shape()[0])
}
return nil
}
// Dense materialises the sparse array as a dense *Array. Renamed
// from `ToDense` for symmetry with `SparseFrom`.
func (s *SparseCOO) Dense() (*Array, error) {
if err := s.consistent("Dense"); err != nil {
return nil, err
}
if narrowRefused(s.Values.dt) {
return nil, errf("Dense: dtype %s is not supported; convert with Astype", s.Values.dt)
}
out, err := Zeros(s.Values.dt, s.Shape...)
if err != nil {
return nil, err
}
strides := denseStrides(s.Shape)
for i := range s.Indices.shape[0] {
flat, err := s.flatEntry(i, strides, "Dense")
if err != nil {
return nil, err
}
switch s.Values.dt {
case Int:
out.ints[flat] = s.Values.ints[i]
case Float16:
out.halves[flat] = s.Values.halves[i]
case Float32:
out.floats32[flat] = s.Values.floats32[i]
case Float:
out.floats[flat] = s.Values.floats[i]
default:
out.complexes[flat] = s.Values.complexes[i]
}
}
return out, nil
}
// NNZ returns the number of non-zero entries.
func (s *SparseCOO) NNZ() int {
return s.Values.Len()
}
// SpMul multiplies a sparse array element-wise by a dense array of the
// same shape. Returns a dense array whose dtype follows the promotion
// ladder.
func SpMul(s *SparseCOO, dense *Array) (*Array, error) {
if err := s.consistent("SpMul"); err != nil {
return nil, err
}
if narrowRefused(s.Values.dt) {
return nil, errf("SpMul: dtype %s is not supported; convert with Astype", s.Values.dt)
}
if narrowRefused(dense.dt) {
return nil, errf("SpMul: dtype %s is not supported; convert with Astype", dense.dt)
}
if !sameShape(s.Shape, dense.shape) {
return nil, errf("SpMul: shape mismatch %v vs %s", s.Shape, shapeText(dense.shape))
}
dt := promote(s.Values.dt, dense.dt)
out, err := Zeros(dt, s.Shape...)
if err != nil {
return nil, err
}
strides := denseStrides(s.Shape)
for i := range s.Indices.shape[0] {
flat, err := s.flatEntry(i, strides, "SpMul")
if err != nil {
return nil, err
}
switch dt {
case Int:
out.ints[flat] = s.Values.ints[i] * dense.ints[flat]
case Float16:
// Computed in float64, narrowed once, like the float32
// branch.
out.halves[flat] = HalfFromFloat64(s.Values.FloatAt(i) * dense.FloatAt(flat))
case Float32:
out.floats32[flat] = float32(s.Values.FloatAt(i) * dense.FloatAt(flat))
case Float:
out.floats[flat] = s.Values.FloatAt(i) * dense.FloatAt(flat)
default:
out.complexes[flat] = s.Values.complexAt(i) * dense.complexAt(flat)
}
}
return out, nil
}
// SpAdd returns a dense array equal to the element-wise sum of two
// sparse arrays. The result is dense because addition can collapse
// zeros into non-zeros. Both arrays must have the same shape.
func SpAdd(a, b *SparseCOO) (*Array, error) {
for _, s := range []*SparseCOO{a, b} {
if narrowRefused(s.Values.dt) {
return nil, errf("SpAdd: dtype %s is not supported; convert with Astype", s.Values.dt)
}
}
if !sameShape(a.Shape, b.Shape) {
return nil, errf("SpAdd: shape mismatch %v vs %v", a.Shape, b.Shape)
}
da, err := a.Dense()
if err != nil {
return nil, err
}
db, err := b.Dense()
if err != nil {
return nil, err
}
return Add(da, db)
}
// SpMatMul multiplies a sparse matrix (n×k) by a dense matrix (k×m).
// Returns a dense n×m result. Only valid for 2-D sparse × 2-D dense.
// The result dtype follows the promotion ladder, like SpMul: int with
// int stays int, any float or complex operand promotes. Every stored
// coordinate is validated, so an out-of-range index is an error naming
// the entry, never an out-of-bounds panic.
func SpMatMul(s *SparseCOO, dense *Array) (*Array, error) {
if err := s.consistent("SpMatMul"); err != nil {
return nil, err
}
if s.Values.dt == Complex {
return nil, errf("SpMatMul: complex sparse is not supported")
}
if s.Values.dt == Float16 || dense.dt == Float16 {
// Like MatMul, the matrix product kernels are not offered for
// the half dtype yet; the refusal is loud, not a silent
// misread of the payload.
return nil, errf("SpMatMul: float16 is not supported; convert with Astype")
}
for _, op := range []Dtype{s.Values.dt, dense.dt} {
if narrowRefused(op) {
return nil, errf("SpMatMul: dtype %s is not supported; convert with Astype", op)
}
}
if len(s.Shape) != 2 || dense.NDim() != 2 {
return nil, errf("SpMatMul: needs 2-D sparse and 2-D dense, got %d-D and %d-D",
len(s.Shape), dense.NDim())
}
if s.Shape[1] != dense.shape[0] {
return nil, errf("SpMatMul: inner dim mismatch %d vs %d", s.Shape[1], dense.shape[0])
}
dt := promote(s.Values.dt, dense.dt)
n := s.Shape[0]
k := s.Shape[1]
m := dense.shape[1]
out := &Array{shape: []int{n, m}, dt: dt}
out.alloc(n * m)
strides := denseStrides(s.Shape)
// The Float32 branch accumulates its products in float64 and narrows
// once per stored coordinate, so one scratch row outlives the
// whole walk.
var acc []float64
if dt == Float32 {
acc = make([]float64, m)
}
for i := range s.Indices.shape[0] {
flat, err := s.flatEntry(i, strides, "SpMatMul")
if err != nil {
return nil, err
}
row := flat / k
col := flat % k
// Add s.Values[i] * dense[col, j] to out[row, j] for each j.
switch dt {
case Int:
for j := range m {
out.ints[row*m+j] += s.Values.ints[i] * dense.ints[col*m+j]
}
case Float32:
vf := s.Values.FloatAt(i)
for j := range m {
acc[j] = float64(out.floats32[row*m+j]) + vf*dense.FloatAt(col*m+j)
}
for j := range m {
out.floats32[row*m+j] = float32(acc[j])
}
case Float:
for j := range m {
out.floats[row*m+j] += s.Values.FloatAt(i) * dense.FloatAt(col*m+j)
}
default:
for j := range m {
out.complexes[row*m+j] += s.Values.complexAt(i) * dense.complexAt(col*m+j)
}
}
}
return out, nil
}
// denseStrides returns the row-major strides for a shape slice.
func denseStrides(shape []int) []int {
strides := make([]int, len(shape))
stride := 1
for i := len(shape) - 1; i >= 0; i-- {
strides[i] = stride
stride *= shape[i]
}
return strides
}
// flatEntry returns the flat dense offset of the i-th stored
// coordinate, checking every index component against the shape. name
// names the calling entry point in the error.
func (s *SparseCOO) flatEntry(i int, strides []int, name string) (int, error) {
flat := 0
for d := range s.Shape {
idx := int(s.Indices.ints[i*len(s.Shape)+d])
if idx < 0 || idx >= s.Shape[d] {
return 0, errf("%s: index [%d]=%d out of range for dim %d of size %d",
name, i, idx, d, s.Shape[d])
}
flat += idx * strides[d]
}
return flat, nil
}