396 lines
13 KiB
Go
396 lines
13 KiB
Go
// 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]
|
|
}
|
|
}
|