Files

396 lines
13 KiB
Go
Raw Permalink 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 "sync"
// Advanced indexing: Gather, Scatter, Nonzero, Take. The first three
// operate along a chosen dimension; Take gathers along the flat axis.
// Indices are int64 arrays; out-of-range indices are an error naming
// the offending position. Gather and Scatter read and write through the
// standard accessors, so a source may sit anywhere on the promotion
// ladder: complex works the same way as the others, except for the one
// narrowing the library rejects (complex into int or float32, Astype's
// rule).
// Gather returns a copy indexed by index along dim: out[i_0, …, i_dim,
// …, i_N] = self[i_0, …, index[i_…, j, …], …, i_N]. The output shape
// is the index shape (after the dim is dropped from self).
func Gather(src *Array, dim int, index *Array) (*Array, error) {
if index.dt != Int {
return nil, errf("Gather: index must be an int array, got %s", index.dt)
}
if dim < 0 || dim >= src.NDim() {
return nil, errf("Gather: dimension %d out of range for shape %s", dim, shapeText(src.shape))
}
// All non-dim dimensions of src must match the corresponding index dim.
if !gatherCompatible(src.shape, index.shape, dim) {
return nil, errf("Gather: shape mismatch %s vs %s on dim %d", shapeText(src.shape), shapeText(index.shape), dim)
}
out := &Array{shape: index.Shape(), dt: src.dt}
n := index.Len()
out.alloc(n)
if n == 0 {
return out, nil
}
// The source is read through the same accessors the element walk
// used, so a strided source materialises to exactly the values those
// reads produced; a contiguous source is its own payload.
ss := src.materialise()
// The workers write disjoint output slots. The serial walk reported
// the first offending position, so the workers publish the smallest
// one under the same lock-free-on-success rule Interpolate2D uses.
var mu sync.Mutex
firstPos := n + 1
var firstVal int
fail := func(pos, idx int) {
mu.Lock()
defer mu.Unlock()
if pos < firstPos {
firstPos, firstVal = pos, idx
}
}
ndim := index.NDim()
switch out.dt {
case Int:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.ints, ss.ints, index, ss.shape, s, e, dim, ndim, fail)
})
case Bool:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.bools, ss.bools, index, ss.shape, s, e, dim, ndim, fail)
})
case Int8:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.i8s, ss.i8s, index, ss.shape, s, e, dim, ndim, fail)
})
case Uint8:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.u8s, ss.u8s, index, ss.shape, s, e, dim, ndim, fail)
})
case Int16:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.i16s, ss.i16s, index, ss.shape, s, e, dim, ndim, fail)
})
case Uint16:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.u16s, ss.u16s, index, ss.shape, s, e, dim, ndim, fail)
})
case Int32:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.i32s, ss.i32s, index, ss.shape, s, e, dim, ndim, fail)
})
case Uint32:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.u32s, ss.u32s, index, ss.shape, s, e, dim, ndim, fail)
})
case Float16:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.halves, ss.halves, index, ss.shape, s, e, dim, ndim, fail)
})
case Float32:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.floats32, ss.floats32, index, ss.shape, s, e, dim, ndim, fail)
})
case Float:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.floats, ss.floats, index, ss.shape, s, e, dim, ndim, fail)
})
default:
parallelMin(n, copyMinPerWorker, func(s, e int) {
gatherFill(out.complexes, ss.complexes, index, ss.shape, s, e, dim, ndim, fail)
})
}
if firstPos <= n {
return nil, errf("Gather: index %d out of range for dimension %d of size %d at position %d",
firstVal, dim, src.shape[dim], firstPos)
}
return out, nil
}
// gatherFill writes dst[i] = src[off] for the worker's range of the
// index walk, where off folds the output's own coordinates with the
// gathered axis replaced by the index. An out-of-range index reports
// through fail and skips the write; a call that failed anywhere
// discards its output, so the skipped slot never surfaces. The odometer
// seeds from the flat start, so the chunks cover disjoint slots in the
// same coordinates one serial walk produced. The loops are bounded by
// the index array's own extent, never by a payload length.
func gatherFill[T any](dst, src []T, index *Array, srcShape []int, s, e, dim, ndim int, fail func(pos, idx int)) {
shape := index.shape
coord := make([]int, ndim)
if s > 0 {
rest := s
for d := ndim - 1; d >= 0; d-- {
coord[d] = rest % shape[d]
rest /= shape[d]
}
}
contig := index.isContiguous()
ip := index.ints
for i := s; i < e; i++ {
idx := int(ip[i])
if !contig {
idx = int(ip[index.physIndex(i)])
}
if idx < 0 || idx >= srcShape[dim] {
fail(i, idx)
advanceOdometer(coord, shape)
continue
}
// The source offset needs one fold per element: the output's own
// coordinates with the gathered axis replaced by the index.
off := 0
for d := range ndim {
c := coord[d]
if d == dim {
c = idx
}
off = off*srcShape[d] + c
}
dst[i] = src[off]
advanceOdometer(coord, shape)
}
}
// gatherCompatible reports whether src and index can combine on dim:
// every non-dim dimension must match in length.
func gatherCompatible(srcShape, indexShape []int, dim int) bool {
if len(srcShape) != len(indexShape) {
return false
}
for d := range srcShape {
if d == dim {
continue
}
if srcShape[d] != indexShape[d] {
return false
}
}
return true
}
// Scatter is the inverse of Gather: it writes src values into self at
// positions indexed by index along dim. It returns a new array; the
// receiver is unchanged. Unindexed positions in the result keep self's
// values.
func Scatter(self *Array, dim int, index *Array, src *Array) (*Array, error) {
if index.dt != Int {
return nil, errf("Scatter: index must be an int array, got %s", index.dt)
}
if dim < 0 || dim >= self.NDim() {
return nil, errf("Scatter: dimension %d out of range for shape %s", dim, shapeText(self.shape))
}
if !sameShape(index.shape, src.shape) {
return nil, errf("Scatter: index and src shapes must agree, got %s and %s", shapeText(index.shape), shapeText(src.shape))
}
if !gatherCompatible(self.shape, index.shape, dim) {
return nil, errf("Scatter: shape mismatch %s vs %s on dim %d", shapeText(self.shape), shapeText(index.shape), dim)
}
// The source is read through its own accessors, so it may sit above
// self on the promotion ladder; complex into int or float32 is the
// one narrowing the library rejects (Astype's rule), and it must be
// a loud error rather than a read of a payload the source lacks.
if !canStore(self.dt, src.dt) {
return nil, errf("Scatter: cannot store a %s source in a %s array", src.dt, self.dt)
}
// The output starts as a full clone of self in self's own dtype,
// narrow payloads included; the setConverted walk below is
// dtype-generic.
out := self.cloneArray()
dst := make([]int, index.NDim())
for i := range index.Len() {
idx := int(index.ints[index.physIndex(i)])
if idx < 0 || idx >= self.shape[dim] {
return nil, errf("Scatter: index %d out of range for dimension %d of size %d at position %d", idx, dim, self.shape[dim], i)
}
// The destination offset needs one fold per element: the source's
// own coordinates with the scatter axis replaced by the index.
off := 0
for d := range dst {
c := dst[d]
if d == dim {
c = idx
}
off = off*self.shape[d] + c
}
out.setConverted(off, src, i)
advanceOdometer(dst, index.shape)
}
return out, nil
}
// Nonzero returns the multi-dimensional indices of every non-zero
// element, grouped per dimension. The result has a.NDim() slices; the
// i-th entry in each slice is the coordinate of the i-th nonzero
// element in that dimension. A complex array is rejected (no notion of
// zero in the lattice; equality with the zero value is already
// supported by Eq).
func Nonzero(a *Array) ([][]int, error) {
if a.dt == Complex {
return nil, errf("Nonzero: complex arrays have no notion of zero")
}
perDim := make([][]int, a.NDim())
coord := make([]int, a.NDim())
for i := range a.Len() {
if !isZero(a, i) {
for d := range a.NDim() {
perDim[d] = append(perDim[d], coord[d])
}
}
advanceOdometer(coord, a.shape)
}
return perDim, nil
}
func isZero(a *Array, i int) bool {
// A strided array maps the logical index to a payload slot through
// its strides, so only a contiguous operand may read the payload at
// the logical position directly; the accessor path keeps Argwhere
// and Nonzero correct the day a public strided view exists.
if a.strides != nil {
switch a.dt {
case Int, Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32:
// intAt widens the integer class exactly and reads bool
// against zero.
return a.intAt(i) == 0
case Float16, Float32, Float:
return a.floatAt(i) == 0
}
return false
}
switch a.dt {
case Int:
return a.ints[i] == 0
case Bool:
return !a.bools[i]
case Int8:
return a.i8s[i] == 0
case Uint8:
return a.u8s[i] == 0
case Int16:
return a.i16s[i] == 0
case Uint16:
return a.u16s[i] == 0
case Int32:
return a.i32s[i] == 0
case Uint32:
return a.u32s[i] == 0
case Float16:
// The value test in bit space: clearing the sign bit folds -0.0
// onto +0.0 exactly as the float comparisons below do.
return a.halves[i]&0x7FFF == 0
case Float32:
return a.floats32[i] == 0
case Float:
return a.floats[i] == 0
}
return false
}
// Take returns a 1-D copy whose elements are self[indices[i]] for each
// i, flat-indexed. Negative indices are an error. The walk is bounded by
// the index array's element count, not by the length of the payload it
// shares: a Slice view's payload can be longer.
func Take(a *Array, indices *Array) (*Array, error) {
if indices.dt != Int {
return nil, errf("Take: indices must be an int array, got %s", indices.dt)
}
if indices.NDim() != 1 {
return nil, errf("Take: indices must be 1-D, got shape %s", shapeText(indices.shape))
}
n := indices.Len()
out := &Array{shape: []int{n}, dt: a.dt}
out.alloc(n)
if n == 0 {
return out, nil
}
src := a
if !src.isContiguous() {
src = src.materialise()
}
k := intPayload(indices)
// The workers write disjoint output slots and report the smallest
// offending position, the one the serial validation walk named.
bound := int64(a.Len())
var mu sync.Mutex
firstPos := n + 1
var firstVal int64
fail := func(pos int, idx int64) {
mu.Lock()
defer mu.Unlock()
if pos < firstPos {
firstPos, firstVal = pos, idx
}
}
switch out.dt {
case Int:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.ints, src.ints, k, s, e, bound, fail)
})
case Bool:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.bools, src.bools, k, s, e, bound, fail)
})
case Int8:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.i8s, src.i8s, k, s, e, bound, fail)
})
case Uint8:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.u8s, src.u8s, k, s, e, bound, fail)
})
case Int16:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.i16s, src.i16s, k, s, e, bound, fail)
})
case Uint16:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.u16s, src.u16s, k, s, e, bound, fail)
})
case Int32:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.i32s, src.i32s, k, s, e, bound, fail)
})
case Uint32:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.u32s, src.u32s, k, s, e, bound, fail)
})
case Float16:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.halves, src.halves, k, s, e, bound, fail)
})
case Float32:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.floats32, src.floats32, k, s, e, bound, fail)
})
case Float:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.floats, src.floats, k, s, e, bound, fail)
})
default:
parallelMin(n, copyMinPerWorker, func(s, e int) {
takeFill(out.complexes, src.complexes, k, s, e, bound, fail)
})
}
if firstPos <= n {
return nil, errf("Take: index %d out of range for flat size %d at position %d", firstVal, a.Len(), firstPos)
}
return out, nil
}
// takeFill writes dst[i] = src[k[i]] over the worker's range. An
// out-of-range index reports through fail and skips the write; a call
// that failed anywhere discards its output.
func takeFill[T any](dst, src []T, k []int64, s, e int, bound int64, fail func(pos int, idx int64)) {
for i := s; i < e; i++ {
v := k[i]
if v < 0 || v >= bound {
fail(i, v)
continue
}
dst[i] = src[v]
}
}