Files
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

374 lines
11 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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
}