234 lines
7.3 KiB
Go
234 lines
7.3 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package core
|
||
|
|
|
||
|
|
// Indexing and slicing. Every read and write goes through a
|
||
|
|
// validated multi-dimensional index; updates are functional:
|
||
|
|
// WithInt, WithFloat and WithComplex return a new array and leave the
|
||
|
|
// receiver untouched. Slice returns a read-only view when the selection
|
||
|
|
// is contiguous and a copy otherwise; either way the source is never
|
||
|
|
// written through.
|
||
|
|
|
||
|
|
// IntAt returns the element at the given index as an int64; the array
|
||
|
|
// must be an integer-class dtype, bool included, and every widening is
|
||
|
|
// exact. A float or complex array is refused with the wording the tests
|
||
|
|
// pin.
|
||
|
|
func IntAt(a *Array, index ...int) (int64, error) {
|
||
|
|
if !intClass(a.dt) {
|
||
|
|
return 0, errf("IntAt: array is %s, not int", a.dt)
|
||
|
|
}
|
||
|
|
off, err := flatIndex(a.shape, index)
|
||
|
|
if err != nil {
|
||
|
|
return 0, err
|
||
|
|
}
|
||
|
|
return a.intAt(off), nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// FloatAt returns the element at the given index as a float64; every
|
||
|
|
// real dtype widens exactly, bool included. A complex array is refused
|
||
|
|
// with the wording the tests pin.
|
||
|
|
func FloatAt(a *Array, index ...int) (float64, error) {
|
||
|
|
if a.dt == Complex {
|
||
|
|
return 0, errf("FloatAt: array is %s, not float", a.dt)
|
||
|
|
}
|
||
|
|
off, err := flatIndex(a.shape, index)
|
||
|
|
if err != nil {
|
||
|
|
return 0, err
|
||
|
|
}
|
||
|
|
return a.floatAt(off), nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// ComplexAt returns the element at the given index; the array must be
|
||
|
|
// complex.
|
||
|
|
func ComplexAt(a *Array, index ...int) (complex128, error) {
|
||
|
|
if a.dt != Complex {
|
||
|
|
return 0, errf("ComplexAt: array is %s, not complex", a.dt)
|
||
|
|
}
|
||
|
|
off, err := flatIndex(a.shape, index)
|
||
|
|
if err != nil {
|
||
|
|
return 0, err
|
||
|
|
}
|
||
|
|
return a.complexes[a.physIndex(off)], nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// WithInt returns a new array with v stored at the given index in the
|
||
|
|
// receiver's own dtype; the receiver is unchanged. Every integer-class
|
||
|
|
// dtype is written, bool included: the store narrows with Go's
|
||
|
|
// conversion, the implicit-store cast the ladder has always carried, and
|
||
|
|
// bool stores v != 0. A float or complex receiver is refused with the
|
||
|
|
// wording the tests pin.
|
||
|
|
func WithInt(a *Array, v int64, index ...int) (*Array, error) {
|
||
|
|
if !intClass(a.dt) {
|
||
|
|
return nil, errf("WithInt: array is %s, not int", a.dt)
|
||
|
|
}
|
||
|
|
off, err := flatIndex(a.shape, index)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
out := a.cloneArray()
|
||
|
|
switch a.dt {
|
||
|
|
case Int:
|
||
|
|
out.ints[off] = v
|
||
|
|
case Bool:
|
||
|
|
out.bools[off] = v != 0
|
||
|
|
case Int8:
|
||
|
|
out.i8s[off] = int8(v)
|
||
|
|
case Uint8:
|
||
|
|
out.u8s[off] = uint8(v)
|
||
|
|
case Int16:
|
||
|
|
out.i16s[off] = int16(v)
|
||
|
|
case Uint16:
|
||
|
|
out.u16s[off] = uint16(v)
|
||
|
|
case Int32:
|
||
|
|
out.i32s[off] = int32(v)
|
||
|
|
default:
|
||
|
|
out.u32s[off] = uint32(v)
|
||
|
|
}
|
||
|
|
return out, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// WithFloat returns a new float array with v at the given index; the
|
||
|
|
// receiver is unchanged. It writes the float dtype only: the narrow
|
||
|
|
// float widths and the integer dtypes keep the pinned refusals.
|
||
|
|
func WithFloat(a *Array, v float64, index ...int) (*Array, error) {
|
||
|
|
if a.dt != Float {
|
||
|
|
return nil, errf("WithFloat: array is %s, not float", a.dt)
|
||
|
|
}
|
||
|
|
off, err := flatIndex(a.shape, index)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
out := a.cloneArray()
|
||
|
|
out.floats[off] = v
|
||
|
|
return out, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// WithComplex returns a new complex array with v at the given index; the
|
||
|
|
// receiver is unchanged. It writes the complex dtype only, the same
|
||
|
|
// kind-gate rule WithFloat carries.
|
||
|
|
func WithComplex(a *Array, v complex128, index ...int) (*Array, error) {
|
||
|
|
if a.dt != Complex {
|
||
|
|
return nil, errf("WithComplex: array is %s, not complex", a.dt)
|
||
|
|
}
|
||
|
|
off, err := flatIndex(a.shape, index)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
out := a.cloneArray()
|
||
|
|
out.complexes[off] = v
|
||
|
|
return out, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// Slice selects the half-open range [start, stop) along one dimension.
|
||
|
|
//
|
||
|
|
// When the selection is a contiguous region of the source, the slice
|
||
|
|
// spans whole rows of it (dim 0) or takes the dimension in full, the
|
||
|
|
// result is a read-only view sharing the source's storage: the payload
|
||
|
|
// is rebased to the first selected element and the shape narrowed, so
|
||
|
|
// payload[i] is still element i and every kernel treats the view as an
|
||
|
|
// ordinary dense array. Any other selection (an interior column range
|
||
|
|
// of a 2-D array, say) is copied, because giving it a strided view
|
||
|
|
// needs the kernel-boundary materialisation audit described in
|
||
|
|
// docs/ARCHITECTURE.md. Either way the receiver is never written
|
||
|
|
// through, and a view never aliases for writing: nothing in the
|
||
|
|
// library writes through an array it did not allocate, and optimiser
|
||
|
|
// updates route through materialised parameters.
|
||
|
|
func Slice(a *Array, dim, start, stop int) (*Array, error) {
|
||
|
|
if dim < 0 || dim >= len(a.shape) {
|
||
|
|
return nil, errf("Slice: dimension %d is out of range for shape %s", dim, shapeText(a.shape))
|
||
|
|
}
|
||
|
|
if start < 0 || stop > a.shape[dim] || start > stop {
|
||
|
|
return nil, errf("Slice: range [%d:%d] is out of range for dimension %d of size %d",
|
||
|
|
start, stop, dim, a.shape[dim])
|
||
|
|
}
|
||
|
|
newShape := a.Shape()
|
||
|
|
newShape[dim] = stop - start
|
||
|
|
// A strided source cannot be sliced by payload arithmetic, so
|
||
|
|
// reduce it to a dense array first; the copy path below and the
|
||
|
|
// view path both assume payload[i] is element i.
|
||
|
|
if !a.isContiguous() {
|
||
|
|
a = a.materialise()
|
||
|
|
}
|
||
|
|
// Physical offset of the first selected element, plus the block of
|
||
|
|
// elements that one index along dim stands for.
|
||
|
|
tail := 1
|
||
|
|
for d := dim + 1; d < a.NDim(); d++ {
|
||
|
|
tail *= a.shape[d]
|
||
|
|
}
|
||
|
|
head := 1
|
||
|
|
for d := range dim {
|
||
|
|
head *= a.shape[d]
|
||
|
|
}
|
||
|
|
if head == 1 || stop-start == a.shape[dim] {
|
||
|
|
view := &Array{shape: newShape, dt: a.dt}
|
||
|
|
view.rebase(a, start*tail)
|
||
|
|
return view, nil
|
||
|
|
}
|
||
|
|
total := 1
|
||
|
|
for _, d := range newShape {
|
||
|
|
total *= d
|
||
|
|
}
|
||
|
|
out := &Array{shape: newShape, dt: a.dt}
|
||
|
|
out.alloc(total)
|
||
|
|
// The selection is one contiguous run of (stop-start) blocks per
|
||
|
|
// position along the leading dimension, and the source above is
|
||
|
|
// dense, so each row moves whole rather than as gathered offsets.
|
||
|
|
// copyRun dispatches every dtype the package stores, the narrow
|
||
|
|
// payloads included.
|
||
|
|
run := (stop - start) * tail
|
||
|
|
stride := a.shape[dim] * tail
|
||
|
|
for r := range head {
|
||
|
|
copyRun(out, a, r*run, r*stride+start*tail, run)
|
||
|
|
}
|
||
|
|
return out, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// rebase points a's payload at element off of src's payload, sharing
|
||
|
|
// the storage. off must be a valid payload index of src; the elements
|
||
|
|
// beyond the view's own count stay invisible because Len is shape-based.
|
||
|
|
func (a *Array) rebase(src *Array, off int) {
|
||
|
|
switch a.dt {
|
||
|
|
case Int:
|
||
|
|
a.ints = src.ints[off:]
|
||
|
|
case Float16:
|
||
|
|
a.halves = src.halves[off:]
|
||
|
|
case Float32:
|
||
|
|
a.floats32 = src.floats32[off:]
|
||
|
|
case Float:
|
||
|
|
a.floats = src.floats[off:]
|
||
|
|
case Complex:
|
||
|
|
a.complexes = src.complexes[off:]
|
||
|
|
case Bool:
|
||
|
|
a.bools = src.bools[off:]
|
||
|
|
case Int8:
|
||
|
|
a.i8s = src.i8s[off:]
|
||
|
|
case Uint8:
|
||
|
|
a.u8s = src.u8s[off:]
|
||
|
|
case Int16:
|
||
|
|
a.i16s = src.i16s[off:]
|
||
|
|
case Uint16:
|
||
|
|
a.u16s = src.u16s[off:]
|
||
|
|
case Int32:
|
||
|
|
a.i32s = src.i32s[off:]
|
||
|
|
default:
|
||
|
|
a.u32s = src.u32s[off:]
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// flatIndex validates a multi-dimensional index against a shape and
|
||
|
|
// returns the row-major flat offset.
|
||
|
|
func flatIndex(shape []int, index []int) (int, error) {
|
||
|
|
if len(index) != len(shape) {
|
||
|
|
return 0, errf("index %v does not match the shape %s", index, shapeText(shape))
|
||
|
|
}
|
||
|
|
off := 0
|
||
|
|
for d, i := range index {
|
||
|
|
if i < 0 || i >= shape[d] {
|
||
|
|
return 0, errf("index %d is out of range for dimension %d of size %d", i, d, shape[d])
|
||
|
|
}
|
||
|
|
off = off*shape[d] + i
|
||
|
|
}
|
||
|
|
return off, nil
|
||
|
|
}
|