374 lines
11 KiB
Go
374 lines
11 KiB
Go
// 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
|
||
}
|