// Copyright (c) 2026 Petr Balvín (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 }