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