Files
tensor/internal/core/arrayutil.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

2047 lines
60 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import "sourcedock.dev/petrbalvin/tensor/internal/base"
import (
"math"
"slices"
"sync/atomic"
)
// Array utilities, array creation, manipulation and conversion
// helpers: Linspace,
// Repeat, Tile, Flip, Roll, Unique, Argwhere, Astype, Item and Diag.
// copyMinPerWorker is the per-worker chunk floor for the kernels whose
// per-element work is a load and a store: the mirror, gather, repeated
// block and payload-conversion copies. Their chunks move bytes, so a
// spawned worker needs a few thousand elements before the chunk
// outweighs its own start-up; below that floor the copy runs on the
// calling goroutine.
const copyMinPerWorker = 1 << 12
// Linspace returns n evenly spaced values from start to stop inclusive.
// n = 1 yields the single value start; n = 0 yields an empty array.
func Linspace(start, stop float64, n int) (*Array, error) {
if n < 0 {
return nil, errf("Linspace: n must be zero or greater, got %d", n)
}
if n == 0 {
return FromFloats(nil, 0)
}
if n == 1 {
return FromFloats([]float64{start}, 1)
}
h := (stop - start) / float64(n-1)
vals := make([]float64, n)
for i := range n {
vals[i] = start + float64(i)*h
}
// Pin the endpoints exactly; the loop can drift on awkward ranges.
vals[0], vals[n-1] = start, stop
return FromFloats(vals, n)
}
// Repeat repeats each element of a `repeats` times along dim.
func Repeat(a *Array, repeats, dim int) (*Array, error) {
if repeats < 0 {
return nil, errf("Repeat: repeats must be zero or greater, got %d", repeats)
}
if dim < 0 || dim >= a.NDim() {
return nil, errf("Repeat: dimension %d out of range for shape %s", dim, shapeText(a.shape))
}
// Bound the extent before multiplying and the product before
// allocating: hostile repeats must be an error, not a wrapped
// extent or a tiny allocation paired with a huge shape.
if a.shape[dim] != 0 && repeats != 0 && a.shape[dim] > math.MaxInt/repeats {
return nil, errf("Repeat: repeats %d overflows dimension %d of length %d", repeats, dim, a.shape[dim])
}
newShape := a.Shape()
newShape[dim] *= repeats
total, _, terr := checkedDims(newShape)
if terr != nil {
return nil, base.WrapErr("Repeat", terr)
}
out := &Array{shape: newShape, dt: a.dt}
out.alloc(total)
// An empty result has nothing to fill: with repeats = 0 the extent
// collapses to zero while the row walk below would still visit the
// source rows and slice the empty destination payload, which
// panicked instead of answering the empty array (the repeat-zero
// test pins it).
if total == 0 {
return out, nil
}
src := a
if !src.isContiguous() {
src = src.materialise()
}
n := a.shape[dim]
outer, inner := 1, 1
for d := range dim {
outer *= a.shape[d]
}
for d := dim + 1; d < a.NDim(); d++ {
inner *= a.shape[d]
}
if inner == 1 {
// The repeated axis is innermost, so result element i is source
// element i/repeats: one flat pass per dtype, no block walk.
repeatEach(out, src, repeats)
return out, nil
}
// Every output row of `inner` elements is a whole copy of one source
// row; the repeat factor only decides which destination rows share it.
for o := range outer {
srcRow := o * n * inner
dstRow := o * n * repeats * inner
for j := range n {
dst := dstRow + j*repeats*inner
copyRun(out, src, dst, srcRow+j*inner, inner)
for k := 1; k < repeats; k++ {
copyRun(out, out, dst+k*inner, dst, inner)
}
}
}
return out, nil
}
// repeatEach fills dst with each element of src repeated r times, the
// mapping of a repeat along the innermost axis: dst[i] = src[i/r].
// repeatEach fills dst with each element of src repeated r times, the
// mapping of a repeat along the innermost axis: dst[i] = src[i/r]. The
// dispatch carries every element type on its own payload.
func repeatEach(dst, src *Array, r int) {
switch dst.dt {
case Int:
repeatInto(dst.ints, src.ints, r)
case Float16:
repeatInto(dst.halves, src.halves, r)
case Float32:
repeatInto(dst.floats32, src.floats32, r)
case Float:
repeatInto(dst.floats, src.floats, r)
case Complex:
repeatInto(dst.complexes, src.complexes, r)
case Bool:
repeatInto(dst.bools, src.bools, r)
case Int8:
repeatInto(dst.i8s, src.i8s, r)
case Uint8:
repeatInto(dst.u8s, src.u8s, r)
case Int16:
repeatInto(dst.i16s, src.i16s, r)
case Uint16:
repeatInto(dst.u16s, src.u16s, r)
case Int32:
repeatInto(dst.i32s, src.i32s, r)
case Uint32:
repeatInto(dst.u32s, src.u32s, r)
}
}
func repeatInto[T any](dst, src []T, r int) {
for i := range dst {
dst[i] = src[i/r]
}
}
// Tile repeats a whole array reps times per dimension. The repetition
// vector is right-aligned to the shape, and a shorter shape prepends
// size-1 dimensions. The result never aliases a.
func Tile(a *Array, reps ...int) (*Array, error) {
if len(reps) == 0 {
return Copy(a), nil
}
ndim := max(a.NDim(), len(reps))
shape := make([]int, ndim)
for d := range ndim {
ad := d - (ndim - a.NDim())
rd := d - (ndim - len(reps))
as := 1
if ad >= 0 {
as = a.shape[ad]
}
r := 1
if rd >= 0 {
r = reps[rd]
}
if r < 0 {
return nil, errf("Tile: negative repetition %d", r)
}
// Bound the extent before multiplying, exactly as Repeat does:
// a wrapped extent must be an error, not a negative allocation.
if as != 0 && r != 0 && as > math.MaxInt/r {
return nil, errf("Tile: repetition %d overflows a dimension of length %d", r, as)
}
shape[d] = as * r
}
total, _, terr := checkedDims(shape)
if terr != nil {
return nil, base.WrapErr("Tile", terr)
}
src := a
if !src.isContiguous() {
src = src.materialise()
}
// The source shape right-aligned to the result's rank; the leading
// size-1 dimensions the alignment adds do not move a flat index.
srcShape := make([]int, ndim)
for d := range ndim {
srcShape[d] = 1
if ad := d - (ndim - a.NDim()); ad >= 0 {
srcShape[d] = a.shape[ad]
}
}
// rep[d] is how many times dimension d repeats, right-aligned like
// the shape.
rep := make([]int, ndim)
last := -1
for d := range ndim {
rep[d] = 1
if rd := d - (ndim - len(reps)); rd >= 0 {
rep[d] = reps[rd]
}
if rep[d] != 1 {
last = d
}
}
out := &Array{shape: shape, dt: a.dt}
out.alloc(total)
if last == ndim-1 && ndim > 0 {
// The innermost dimension repeats, so a block of the result is no
// longer one run of the source; the coordinate walk fills it.
coord := make([]int, ndim)
for i := range out.Len() {
off := 0
for d := range ndim {
if srcShape[d] == 0 {
continue
}
off = off*srcShape[d] + coord[d]%srcShape[d]
}
out.setFrom(i, src, off)
advanceOdometer(coord, shape)
}
return out, nil
}
// Every dimension after the last repeat copies whole, so one output
// block of `run` elements is one contiguous run of the source.
run := 1
for d := last + 1; d < ndim; d++ {
run *= srcShape[d]
}
coord := make([]int, last+1)
blocks := 1
for d := 0; d <= last; d++ {
blocks *= shape[d]
}
for range blocks {
dstOff, srcOff := 0, 0
for d := 0; d <= last; d++ {
dstOff = dstOff*shape[d] + coord[d]
srcOff = srcOff*srcShape[d] + coord[d]%srcShape[d]
}
copyRun(out, src, dstOff*run, srcOff*run, run)
advanceOdometer(coord, shape)
}
return out, nil
}
// Flip reverses a along the given dimensions (all of them by default).
func Flip(a *Array, dims ...int) (*Array, error) {
if len(dims) == 0 {
dims = rangeN(a.NDim())
}
flipSet := make([]bool, a.NDim())
for _, d := range dims {
if d < 0 || d >= a.NDim() {
return nil, errf("Flip: dimension %d out of range for shape %s", d, shapeText(a.shape))
}
flipSet[d] = true
}
out := &Array{shape: a.Shape(), dt: a.dt}
out.alloc(a.Len())
src := a
if !src.isContiguous() {
src = src.materialise()
}
// Every dimension after the last flipped one keeps its order, so one
// run of `run` elements is copied whole and only the run's position
// reflects the reversal. The prefix order is walked with an odometer
// over the flipped coordinate, which folds into the destination run
// index.
last := -1
for _, d := range dims {
last = max(last, d)
}
run := 1
for d := last + 1; d < a.NDim(); d++ {
run *= a.shape[d]
}
if run == 1 {
// Nothing past the last flipped dimension survives the reversal,
// so the flipped axis itself becomes the reversed run: a row of
// shape[last] elements mirrors whole and the prefix odometer
// only decides where the row lands.
n := a.shape[last]
head := 1
for d := range last {
head *= a.shape[d]
}
row := make([]int, last)
for range head {
dstRow, srcRow := 0, 0
for d := range last {
c := row[d]
dst := c
if flipSet[d] {
dst = a.shape[d] - 1 - c
}
dstRow = dstRow*a.shape[d] + dst
srcRow = srcRow*a.shape[d] + c
}
reverseRows(out, src, dstRow*n, srcRow*n, n)
advanceOdometer(row, a.shape)
}
return out, nil
}
// Every dimension after the last flipped one keeps its order, so one
// run of `run` elements moves whole and only the run's position
// reflects the reversal.
coord := make([]int, last+1)
blocks := 1
for d := 0; d <= last; d++ {
blocks *= a.shape[d]
}
for range blocks {
dstOff := 0
for d := 0; d <= last; d++ {
c := coord[d]
if flipSet[d] {
c = a.shape[d] - 1 - c
}
dstOff = dstOff*a.shape[d] + c
}
copyRun(out, src, dstOff*run, blockFlat(coord, a.shape)*run, run)
advanceOdometer(coord, a.shape)
}
return out, nil
}
// reverseRows mirrors the run of src at srcOff into dst at dstOff, one
// row of a flip: the serial form the per-row walk needs, since spawning
// workers for every row would cost more than the row. The dispatch
// carries every element type on its own payload.
func reverseRows(dst, src *Array, dstOff, srcOff, run int) {
switch dst.dt {
case Int:
reverseIntoSerial(dst.ints[dstOff:dstOff+run], src.ints[srcOff:srcOff+run])
case Float16:
reverseIntoSerial(dst.halves[dstOff:dstOff+run], src.halves[srcOff:srcOff+run])
case Float32:
reverseIntoSerial(dst.floats32[dstOff:dstOff+run], src.floats32[srcOff:srcOff+run])
case Float:
reverseIntoSerial(dst.floats[dstOff:dstOff+run], src.floats[srcOff:srcOff+run])
case Complex:
reverseIntoSerial(dst.complexes[dstOff:dstOff+run], src.complexes[srcOff:srcOff+run])
case Bool:
reverseIntoSerial(dst.bools[dstOff:dstOff+run], src.bools[srcOff:srcOff+run])
case Int8:
reverseIntoSerial(dst.i8s[dstOff:dstOff+run], src.i8s[srcOff:srcOff+run])
case Uint8:
reverseIntoSerial(dst.u8s[dstOff:dstOff+run], src.u8s[srcOff:srcOff+run])
case Int16:
reverseIntoSerial(dst.i16s[dstOff:dstOff+run], src.i16s[srcOff:srcOff+run])
case Uint16:
reverseIntoSerial(dst.u16s[dstOff:dstOff+run], src.u16s[srcOff:srcOff+run])
case Int32:
reverseIntoSerial(dst.i32s[dstOff:dstOff+run], src.i32s[srcOff:srcOff+run])
case Uint32:
reverseIntoSerial(dst.u32s[dstOff:dstOff+run], src.u32s[srcOff:srcOff+run])
}
}
// blockFlat folds a full prefix of coordinates into its row-major flat
// index.
func blockFlat(coord, shape []int) int {
off := 0
for d, c := range coord {
off = off*shape[d] + c
}
return off
}
// Roll shifts a along dim by shift; values wrap around and a negative
// shift moves elements backward.
func Roll(a *Array, shift, dim int) (*Array, error) {
if dim < 0 || dim >= a.NDim() {
return nil, errf("Roll: dimension %d out of range for shape %s", dim, shapeText(a.shape))
}
n := a.shape[dim]
if n == 0 {
return Copy(a), nil
}
shift %= n
if shift < 0 {
shift += n
}
out := &Array{shape: a.Shape(), dt: a.dt}
out.alloc(a.Len())
src := a
if !src.isContiguous() {
src = src.materialise()
}
// out[c] = a[(c − shift) mod n] along the axis, so one outer block
// rotates as two runs: the axis tail after the shift, then the head.
// Every position off the axis keeps its order, so each run of `tail`
// elements moves whole.
tail := 1
for d := dim + 1; d < a.NDim(); d++ {
tail *= a.shape[d]
}
outer := 1
for d := range dim {
outer *= a.shape[d]
}
split := (n - shift) * tail
for o := range outer {
base := o * n * tail
copyRun(out, src, base+shift*tail, base, split)
copyRun(out, src, base, base+split, n*tail-split)
}
return out, nil
}
// Unique returns the sorted unique values of a (real arrays only); NaN
// counts once. The ordering rule is the same for every dtype, the one
// the Int branch has always read off Sort: values ascend in the dtype's
// natural order, equal adjacent values collapse to the first, and the
// output keeps the input dtype. The narrow integer class sorts in exact
// int64 space, where that rule is the Int rule verbatim; bool orders
// false before true. The float dtypes keep the Sort-backed walk, with
// NaN counted once through sameFloat.
func Unique(a *Array) (*Array, error) {
if a.dt == Complex {
return nil, errf("Unique: complex arrays have no ordering")
}
switch a.dt {
case Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32:
// Sort carries no narrow payload dispatch, so the narrow class
// walks its own exact int64 widening: ascending order and the
// adjacent dedupe are exactly what the Int branch reads off
// Sort, and the result rebuilds in the source dtype.
n := a.Len()
vals := make([]int64, n)
for i := range n {
vals[i] = a.intAt(i)
}
slices.Sort(vals)
uniq := make([]int64, 0, len(vals))
for i, v := range vals {
if i == 0 || v != vals[i-1] {
uniq = append(uniq, v)
}
}
switch a.dt {
case Bool:
out := make([]bool, len(uniq))
for i, v := range uniq {
out[i] = v != 0
}
return FromBools(out, len(out))
case Int8:
out := make([]int8, len(uniq))
for i, v := range uniq {
out[i] = int8(v)
}
return FromInt8s(out, len(out))
case Uint8:
out := make([]uint8, len(uniq))
for i, v := range uniq {
out[i] = uint8(v)
}
return FromUint8s(out, len(out))
case Int16:
out := make([]int16, len(uniq))
for i, v := range uniq {
out[i] = int16(v)
}
return FromInt16s(out, len(out))
case Uint16:
out := make([]uint16, len(uniq))
for i, v := range uniq {
out[i] = uint16(v)
}
return FromUint16s(out, len(out))
case Int32:
out := make([]int32, len(uniq))
for i, v := range uniq {
out[i] = int32(v)
}
return FromInt32s(out, len(out))
default:
out := make([]uint32, len(uniq))
for i, v := range uniq {
out[i] = uint32(v)
}
return FromUint32s(out, len(out))
}
}
sorted, err := Sort(a)
if err != nil {
return nil, err
}
switch a.dt {
case Int:
out := make([]int64, 0, len(sorted.ints))
for i, v := range sorted.ints {
if i == 0 || v != sorted.ints[i-1] {
out = append(out, v)
}
}
return FromInts(out, len(out))
case Float16:
out := make([]uint16, 0, len(sorted.halves))
for i, v := range sorted.halves {
if i == 0 || !sameFloat(HalfToFloat64(v), HalfToFloat64(sorted.halves[i-1])) {
out = append(out, v)
}
}
return HalvesFromArray(out, len(out))
case Float32:
out := make([]float32, 0, len(sorted.floats32))
for i, v := range sorted.floats32 {
if i == 0 || !sameFloat(float64(v), float64(sorted.floats32[i-1])) {
out = append(out, v)
}
}
return FromFloat32s(out, len(out))
default:
out := make([]float64, 0, len(sorted.floats))
for i, v := range sorted.floats {
if i == 0 || !sameFloat(v, sorted.floats[i-1]) {
out = append(out, v)
}
}
return FromFloats(out, len(out))
}
}
// sameFloat reports equality with NaN treated as equal to NaN: the
// dedupe rule for Unique.
func sameFloat(x, y float64) bool {
return x == y || (x != x && y != y)
}
// argwhereChunkBufMin is the first capacity of a chunk's coordinate
// buffer in Argwhere: large enough that a sparse chunk grows a handful
// of times, and capped by what the chunk could ever hold, so a small
// array pays a single exact allocation. It never guesses density:
// growth past it doubles, and the merge copies each chunk's coordinates
// into the exact output either way.
const argwhereChunkBufMin = 1024
// Argwhere returns the coordinates of every non-zero element as an
// (nnz, ndim) int array of non-zero coordinates.
func Argwhere(a *Array) (*Array, error) {
if a.dt == Complex {
return nil, errf("Argwhere: complex arrays have no notion of zero")
}
n := a.Len()
ndim := a.NDim()
chunk := n
chunks := 1
if n >= copyMinPerWorker {
w := workersFor(n)
chunk = (n + w - 1) / w
chunks = (n + chunk - 1) / chunk
}
// One input pass: every chunk walks its own range once, with its own
// odometer, appending its coordinates to a chunk-local buffer, so
// the parts concatenate into exactly the row-major order one walk
// produced. The merge then copies them in chunk order into the exact
// output, which is fully overwritten before anyone reads it.
parts := make([][]int64, chunks)
parallelMin(chunks, 1, func(s, e int) {
for c := s; c < e; c++ {
start := c * chunk
end := min(start+chunk, n)
buf := make([]int64, 0, min(argwhereChunkBufMin, (end-start)*ndim))
parts[c] = argwhereAppend(a, buf, start, end)
}
})
total := 0
for _, p := range parts {
total += len(p)
}
// Each part carries ndim values per non-zero element, so the row
// count is the value count over ndim.
out := &Array{shape: []int{total / ndim, ndim}, dt: Int}
out.alloc(total)
off := 0
for _, p := range parts {
copy(out.ints[off:off+len(p)], p)
off += len(p)
}
return out, nil
}
// argwhereAppend walks [start, end) with the chunk's own odometer and
// appends the coordinates of every non-zero element to buf, in
// row-major order. The seeding of the odometer from the flat start is
// what lets each chunk of a split walk produce its own part in order.
// The zero test is dispatched on the dtype once, and the loops are
// bounded by the array's own extent, never by a payload length: a
// rebased view carries a payload longer than its own Len.
func argwhereAppend(a *Array, buf []int64, start, end int) []int64 {
ndim := a.NDim()
coord := make([]int, ndim)
if start > 0 {
rest := start
for d := ndim - 1; d >= 0; d-- {
coord[d] = rest % a.shape[d]
rest /= a.shape[d]
}
}
appendCoord := func() {
for d := range ndim {
buf = append(buf, int64(coord[d]))
}
}
if !a.isContiguous() {
for i := start; i < end; i++ {
if !isZero(a, i) {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
return buf
}
switch a.dt {
case Int:
p := a.ints
for i := start; i < end; i++ {
if p[i] != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
case Float16:
p := a.halves
for i := start; i < end; i++ {
if p[i]&0x7FFF != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
case Float32:
p := a.floats32
for i := start; i < end; i++ {
if p[i] != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
case Float:
p := a.floats
for i := start; i < end; i++ {
if p[i] != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
case Bool:
// The mask reading: a true element is the non-zero one.
p := a.bools
for i := start; i < end; i++ {
if p[i] {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
case Int8:
p := a.i8s
for i := start; i < end; i++ {
if p[i] != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
case Uint8:
p := a.u8s
for i := start; i < end; i++ {
if p[i] != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
case Int16:
p := a.i16s
for i := start; i < end; i++ {
if p[i] != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
case Uint16:
p := a.u16s
for i := start; i < end; i++ {
if p[i] != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
case Int32:
p := a.i32s
for i := start; i < end; i++ {
if p[i] != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
case Uint32:
p := a.u32s
for i := start; i < end; i++ {
if p[i] != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
default:
// Complex: Argwhere's entry gate rejects complex arrays before
// the walk, so the arm exists for dispatch completeness; the
// value test matches the one the float arms carry.
p := a.complexes
for i := start; i < end; i++ {
if p[i] != 0 {
appendCoord()
}
advanceOdometer(coord, a.shape)
}
}
return buf
}
// Astype returns a copy of a converted to dt. The split is deliberate:
// the legacy conversions keep their historical cast semantics, the new
// narrow targets check their range. int to float and float to int
// convert like Go's casts; a float destination rounds through float64,
// so a float16 or float32 target narrows with HalfFromFloat64 and
// float32(v); complex to float keeps the real part, complex to int and
// complex to float16 are errors; real to complex adds a zero imaginary
// part. Narrowing into bool, int8, uint8, int16, uint16, int32 or
// uint32 is range-checked in the source's value space, and the first
// element the target cannot represent fails the call loudly: integer
// sources check exactly in int64; a float source must be finite,
// integral and inside the target's range, so a NaN, an infinity or a
// fraction is an error rather than a silent cast. Bool is the one
// target with no range to check: every source reaches it by the test
// against zero (a NaN reads true), and bool converts out as 0/1, exact
// everywhere including complex. Converting to the array's own dtype
// copies it unchanged, which the int route cannot express as a cast:
// float64 rounds above 2^53.
func Astype(a *Array, dt Dtype) (*Array, error) {
if a.dt == dt {
// Same dtype: a straight copy through cloneArray, fast and
// exact for every element type; the float64 detour a numeric
// route would take rounds above 2^53.
return a.cloneArray(), nil
}
out := &Array{shape: a.Shape(), dt: dt}
out.alloc(a.Len())
if dt == Float && a.dt == Complex {
// Complex to float keeps the real part, as documented.
parallelMin(a.Len(), copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.floats[i] = real(a.complexes[i])
}
})
return out, nil
}
if dt == Bool {
// Every source reaches bool by the test against zero, through
// the widening reader that resolves a view's strides; a NaN
// compares unequal to zero and reads true. No value can fail
// this conversion, so the walk cannot error.
n := a.Len()
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = a.boolAt(i)
}
})
return out, nil
}
if a.dt == Complex {
// Every remaining target sits below complex on the ladder and
// has no defined complex value space, the narrow numeric
// targets included: the historical loud refusal.
return nil, errf("Astype: cannot narrow complex to %s", dt)
}
src := a
if !src.isContiguous() {
// A strided source has no payload run to read; reduce it to a
// dense copy through the same accessor the walk used.
src = src.materialise()
}
// The source dispatch sits outside the element loop, so each
// destination runs one monomorphised conversion whose arithmetic is
// the widening floatAt performed followed by the destination's own
// cast, bit for bit.
n := src.Len()
switch src.dt {
case Int:
p := src.ints[:n]
switch dt {
case Float16:
d := out.halves
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = HalfFromFloat64(float64(p[i]))
}
})
case Float32:
d := out.floats32
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = float32(float64(p[i]))
}
})
case Float:
d := out.floats
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = float64(p[i])
}
})
case Int8, Uint8, Int16, Uint16, Int32, Uint32:
// Narrowing into a new narrow target: range-checked in the
// source's exact int64 value space.
if err := astypeNarrowFromInt(out, p); err != nil {
return nil, err
}
return out, nil
default:
d := out.complexes
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = complex(float64(p[i]), 0)
}
})
}
case Float16:
p := src.halves[:n]
switch dt {
case Int:
d := out.ints
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = int64(HalfToFloat64(p[i]))
}
})
case Float32:
d := out.floats32
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = float32(HalfToFloat64(p[i]))
}
})
case Float:
d := out.floats
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = HalfToFloat64(p[i])
}
})
case Int8, Uint8, Int16, Uint16, Int32, Uint32:
// Range-checked in the source's float64 value space: the
// half widens exactly, then the float rule applies.
if err := astypeNarrowFromFloat(out, floatPayload(src)); err != nil {
return nil, err
}
return out, nil
default:
d := out.complexes
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = complex(HalfToFloat64(p[i]), 0)
}
})
}
case Float32:
p := src.floats32[:n]
switch dt {
case Int:
d := out.ints
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = int64(float64(p[i]))
}
})
case Float16:
d := out.halves
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = HalfFromFloat64(float64(p[i]))
}
})
case Float:
d := out.floats
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = float64(p[i])
}
})
case Int8, Uint8, Int16, Uint16, Int32, Uint32:
// Range-checked in the source's float64 value space: the
// float32 widens exactly, then the float rule applies.
if err := astypeNarrowFromFloat(out, floatPayload(src)); err != nil {
return nil, err
}
return out, nil
default:
d := out.complexes
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = complex(float64(p[i]), 0)
}
})
}
case Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32:
// The narrow sources widen exactly into int64, the value space
// their conversions and range checks work in; bool reads 0/1.
vals := make([]int64, n)
for i := range n {
vals[i] = src.intAt(i)
}
switch dt {
case Int:
d := out.ints
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = vals[i]
}
})
case Float16:
d := out.halves
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
// The exact float64 widening followed by the
// destination's own nearest-narrowing cast; only the
// bool, int8 and uint8 value ranges stay exact in
// every half, wider integers round above 2048.
d[i] = HalfFromFloat64(float64(vals[i]))
}
})
case Float32:
d := out.floats32
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
// Exact through float64; float32 rounds above 2^24,
// which reaches the int32 and uint32 sources.
d[i] = float32(float64(vals[i]))
}
})
case Float:
d := out.floats
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
// Exact: every narrow source value fits float64.
d[i] = float64(vals[i])
}
})
case Complex:
d := out.complexes
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
// Exact: the element's exact float64 value with a
// zero imaginary part.
d[i] = complex(float64(vals[i]), 0)
}
})
case Int8, Uint8, Int16, Uint16, Int32, Uint32:
// Narrow to narrow: exact where the target contains the
// source's value set, a loud range error where it does not
// (int8 to uint8 of a negative value, say).
if err := astypeNarrowFromInt(out, vals); err != nil {
return nil, err
}
return out, nil
}
default:
// Float, the one source still unnamed.
p := src.floats[:n]
switch dt {
case Int:
d := out.ints
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = int64(p[i])
}
})
case Float16:
d := out.halves
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = HalfFromFloat64(p[i])
}
})
case Float32:
d := out.floats32
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = float32(p[i])
}
})
case Int8, Uint8, Int16, Uint16, Int32, Uint32:
// Range-checked in the source's float64 value space.
if err := astypeNarrowFromFloat(out, p); err != nil {
return nil, err
}
return out, nil
default:
d := out.complexes
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
d[i] = complex(p[i], 0)
}
})
}
}
return out, nil
}
// astypeNarrowFromInt writes vals, the source values in exact int64
// space, into out's narrow integer payload. Representability is checked
// in that space, and the lowest index whose value the target cannot
// hold fails the call with the Astype range error.
func astypeNarrowFromInt(out *Array, vals []int64) error {
n := len(vals)
switch out.dt {
case Int8:
return narrowWrite(out.i8s[:n], vals, math.MinInt8, math.MaxInt8, out.dt)
case Uint8:
return narrowWrite(out.u8s[:n], vals, 0, math.MaxUint8, out.dt)
case Int16:
return narrowWrite(out.i16s[:n], vals, math.MinInt16, math.MaxInt16, out.dt)
case Uint16:
return narrowWrite(out.u16s[:n], vals, 0, math.MaxUint16, out.dt)
case Int32:
return narrowWrite(out.i32s[:n], vals, math.MinInt32, math.MaxInt32, out.dt)
default:
return narrowWrite(out.u32s[:n], vals, 0, math.MaxUint32, out.dt)
}
}
// astypeNarrowFromFloat is astypeNarrowFromInt for float sources, whose
// value space is float64. A representable element is finite, integral
// and inside the target's range; a NaN, an infinity or a fraction fails
// the call instead of casting silently, and the error names the value
// in the source's own float64 space.
func astypeNarrowFromFloat(out *Array, vals []float64) error {
n := len(vals)
switch out.dt {
case Int8:
return narrowWriteF(out.i8s[:n], vals, math.MinInt8, math.MaxInt8, out.dt)
case Uint8:
return narrowWriteF(out.u8s[:n], vals, 0, math.MaxUint8, out.dt)
case Int16:
return narrowWriteF(out.i16s[:n], vals, math.MinInt16, math.MaxInt16, out.dt)
case Uint16:
return narrowWriteF(out.u16s[:n], vals, 0, math.MaxUint16, out.dt)
case Int32:
return narrowWriteF(out.i32s[:n], vals, math.MinInt32, math.MaxInt32, out.dt)
default:
return narrowWriteF(out.u32s[:n], vals, 0, math.MaxUint32, out.dt)
}
}
// narrowWrite stores vals in dst, failing the call when a value falls
// outside [lo, hi]. The walk parallelises like every conversion kernel,
// so the failure published is the lowest failing index: the first
// element in row-major order, whichever worker sees it. The stored cast
// is exact once the range holds.
func narrowWrite[T int8 | uint8 | int16 | uint16 | int32 | uint32](dst []T, vals []int64, lo, hi int64, dt Dtype) error {
n := len(vals)
var bad atomic.Int64
bad.Store(int64(n))
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
v := vals[i]
if v < lo || v > hi {
for {
cur := bad.Load()
if int64(i) >= cur || bad.CompareAndSwap(cur, int64(i)) {
break
}
}
continue
}
dst[i] = T(v)
}
})
if idx := int(bad.Load()); idx < n {
return errf("Astype: value %v at index %d does not fit %s", vals[idx], idx, dt)
}
return nil
}
// narrowWriteF is narrowWrite for float64 source values: representable
// means finite, integral and inside [lo, hi], and the stored cast is
// exact once those hold.
func narrowWriteF[T int8 | uint8 | int16 | uint16 | int32 | uint32](dst []T, vals []float64, lo, hi float64, dt Dtype) error {
n := len(vals)
var bad atomic.Int64
bad.Store(int64(n))
parallelMin(n, copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
v := vals[i]
if math.IsNaN(v) || math.IsInf(v, 0) || math.Trunc(v) != v || v < lo || v > hi {
for {
cur := bad.Load()
if int64(i) >= cur || bad.CompareAndSwap(cur, int64(i)) {
break
}
}
continue
}
dst[i] = T(v)
}
})
if idx := int(bad.Load()); idx < n {
return errf("Astype: value %v at index %d does not fit %s", vals[idx], idx, dt)
}
return nil
}
// astypeArray is Astype for the autograd engine; the float and
// float32 conversions it needs never error and return the input
// unchanged when the dtype already matches.
func astypeArray(a *Array, dt Dtype) (*Array, error) {
if a.dt == dt {
return a, nil
}
return Astype(a, dt)
}
// Item returns the single element of a 1-element array as a float64
// (the real part for complex arrays).
func Item(a *Array) (float64, error) {
if a.Len() != 1 {
return 0, errf("Item: needs a 1-element array, got shape %s", shapeText(a.shape))
}
if a.dt == Complex {
return real(a.complexes[0]), nil
}
return a.floatAt(0), nil
}
// Diag extracts the main diagonal of a 2-D array, or builds a diagonal
// matrix from a 1-D array.
func Diag(a *Array) (*Array, error) {
switch a.NDim() {
case 1:
n := a.Len()
out, err := Zeros(a.dt, n, n)
if err != nil {
return nil, err
}
for i := range n {
out.setFrom(i*n+i, a, i)
}
return out, nil
case 2:
return Diagonal(a, 0)
}
return nil, errf("Diag: needs a 1-D or 2-D array, got shape %s", shapeText(a.shape))
}
// All reports whether every element is non-zero (real arrays only; the
// mask semantics match Select: any non-zero value counts as true).
func All(a *Array) (bool, error) {
if a.dt == Complex {
return false, errf("All: complex arrays have no notion of zero")
}
return zeroScan(a, true, 1) == 0, nil
}
// Any reports whether at least one element is non-zero.
func Any(a *Array) (bool, error) {
if a.dt == Complex {
return false, errf("Any: complex arrays have no notion of zero")
}
return zeroScan(a, false, 1) > 0, nil
}
// CountNonzero returns the number of non-zero elements.
func CountNonzero(a *Array) (int, error) {
if a.dt == Complex {
return 0, errf("CountNonzero: complex arrays have no notion of zero")
}
return zeroScan(a, false, 0), nil
}
// zeroScan counts the elements of a that test as countZeros wants,
// stopping once it has seen stop of them (stop 0 counts them all). The
// zero test is dispatched on the dtype once rather than per element,
// and the walk stops as early as the old per-element loop did.
func zeroScan(a *Array, countZeros bool, stop int) int {
n := a.Len()
count := 0
hit := func() bool {
count++
return stop > 0 && count >= stop
}
if !a.isContiguous() {
for i := range n {
if (a.floatAt(i) == 0) == countZeros && hit() {
return count
}
}
return count
}
switch a.dt {
case Int:
for _, v := range a.ints[:n] {
if (v == 0) == countZeros && hit() {
return count
}
}
case Float16:
// The value test in bit space: clearing the sign bit folds -0.0
// onto +0.0 exactly as the float comparisons do, and a half's
// non-zero patterns all widen to a non-zero double.
for _, v := range a.halves[:n] {
if (v&0x7FFF == 0) == countZeros && hit() {
return count
}
}
case Float32:
for _, v := range a.floats32[:n] {
if (v == 0) == countZeros && hit() {
return count
}
}
case Float:
for _, v := range a.floats[:n] {
if (v == 0) == countZeros && hit() {
return count
}
}
case Bool:
// The mask reading: false is the zero element, true the
// non-zero one, the same test isZero carries.
for _, v := range a.bools[:n] {
if (!v) == countZeros && hit() {
return count
}
}
case Int8:
for _, v := range a.i8s[:n] {
if (v == 0) == countZeros && hit() {
return count
}
}
case Uint8:
for _, v := range a.u8s[:n] {
if (v == 0) == countZeros && hit() {
return count
}
}
case Int16:
for _, v := range a.i16s[:n] {
if (v == 0) == countZeros && hit() {
return count
}
}
case Uint16:
for _, v := range a.u16s[:n] {
if (v == 0) == countZeros && hit() {
return count
}
}
case Int32:
for _, v := range a.i32s[:n] {
if (v == 0) == countZeros && hit() {
return count
}
}
case Uint32:
for _, v := range a.u32s[:n] {
if (v == 0) == countZeros && hit() {
return count
}
}
default:
// Complex: All, Any and CountNonzero reject complex arrays at
// entry, so the arm exists for dispatch completeness; the value
// test matches the one the float arms carry.
for _, v := range a.complexes[:n] {
if (v == 0) == countZeros && hit() {
return count
}
}
}
return count
}
// zeros allocates a zeroed array, treating the constructor's error as
// unreachable for internally derived shapes.
func zeros(dt Dtype, shape []int) *Array {
a, _ := Zeros(dt, shape...)
return a
}
// Grid builds coordinate matrices from two 1-D vectors: the meshgrid
// pattern for evaluating functions over a 2-D domain. X varies along
// the columns and Y along the rows.
func Grid(a, b *Array) (xGrid, yGrid *Array, err error) {
if a.NDim() != 1 || b.NDim() != 1 {
return nil, nil, errf("Grid: needs two 1-D arrays, got %s and %s", shapeText(a.shape), shapeText(b.shape))
}
if a.dt == Complex || b.dt == Complex {
return nil, nil, errf("Grid: complex axes are not supported")
}
na, nb := a.Len(), b.Len()
xGrid = &Array{shape: []int{nb, na}, dt: Float}
yGrid = &Array{shape: []int{nb, na}, dt: Float}
xGrid.alloc(nb * na)
yGrid.alloc(nb * na)
// The x row is one copy per row, the y row one fill; the widening of
// the two 1-D sources is the exact one FloatAt reports.
av, bv := floatPayload(a), floatPayload(b)
parallelMin(nb, copyMinPerWorker, func(s, e int) {
for r := s; r < e; r++ {
copy(xGrid.floats[r*na:(r+1)*na], av)
yr := bv[r]
for c := range na {
yGrid.floats[r*na+c] = yr
}
}
})
return xGrid, yGrid, nil
}
// CrossProduct computes the vector cross product of two length-3
// vectors.
func CrossProduct(u, v *Array) (*Array, error) {
if u.dt == Complex || v.dt == Complex {
return nil, errf("CrossProduct: complex vectors are not supported")
}
if u.Len() != 3 || v.Len() != 3 {
return nil, errf("CrossProduct: needs two length-3 vectors, got %d and %d", u.Len(), v.Len())
}
out := zeros(Float, []int{3})
u1, u2, u3 := u.FloatAt(0), u.FloatAt(1), u.FloatAt(2)
v1, v2, v3 := v.FloatAt(0), v.FloatAt(1), v.FloatAt(2)
out.SetFloatAt(0, u2*v3-u3*v2)
out.SetFloatAt(1, u3*v1-u1*v3)
out.SetFloatAt(2, u1*v2-u2*v1)
return out, nil
}
// Integrate computes the definite integral of y over uniform spacing
// dx using the trapezoidal rule. The areas fold through the canonical
// partition, the accuracy the central-sum test measured against the
// exact referent; at and below one block the walk is the plain chain,
// bit for bit.
func Integrate(y *Array, dx float64) (float64, error) {
if y.dt == Complex {
return 0, errf("Integrate: complex samples are not supported")
}
n := y.Len()
if n < 2 {
return 0, errf("Integrate: needs at least two samples, got %d", n)
}
// The widening happens once, outside the accumulation: the addends
// and their order are the ones the accessor loop summed.
yv := floatPayload(y)
return integrateAreas(yv) * dx, nil
}
// integrateAreas sums the y samples' trapezoid areas through the
// canonical partition: the areas (one fewer than the samples) cut into
// the fixed blocks, one chain partial per block and the balanced tree
// over them. Area i reads samples i and i+1, the arithmetic the plain
// chain kept.
func integrateAreas(y []float64) float64 {
m := len(y) - 1
parts := foldParts(m)
if parts == 1 {
var total float64
for i := range m {
total += (y[i] + y[i+1]) / 2
}
return total
}
partials := make([]float64, parts)
for c := range parts {
lo, hi := c*m/parts, (c+1)*m/parts
var acc float64
for i := lo; i < hi; i++ {
acc += (y[i] + y[i+1]) / 2
}
partials[c] = acc
}
return treeSum(partials)
}
// CumulativeIntegrate returns the running trapezoidal integral of y
// over uniform spacing dx; the first element is zero.
func CumulativeIntegrate(y *Array, dx float64) (*Array, error) {
if y.dt == Complex {
return nil, errf("CumulativeIntegrate: complex samples are not supported")
}
n := y.Len()
yv := floatPayload(y)
out := zeros(Float, []int{n})
ov := out.floats
for i := 1; i < n; i++ {
ov[i] = ov[i-1] + (yv[i-1]+yv[i])/2*dx
}
return out, nil
}
// Interpolate evaluates the piecewise-linear interpolation of the
// points (xs[i], ys[i]) at each query position; queries outside the
// range clamp to the boundary values. xs need only be non-decreasing:
// a repeated knot gives a zero-width segment the walk skips, and a
// query landing exactly on it takes the left segment's upper end, so
// it returns the first of the repeated ys. Non-finite knots are
// refused, and so is a NaN query, which has no position to clamp to.
// The segment is located by searching the knots for the first one at
// or above the query, which is the segment the ascending walk stopped
// on; the per-point arithmetic is unchanged.
func Interpolate(xs, ys *Array, query *Array) (*Array, error) {
if xs.dt == Complex || ys.dt == Complex || query.dt == Complex {
return nil, errf("Interpolate: complex samples are not supported")
}
if xs.Len() != ys.Len() || xs.Len() < 2 {
return nil, errf("Interpolate: xs/ys must share length ≥ 2")
}
for k := range xs.Len() {
xk := xs.FloatAt(k)
if math.IsNaN(xk) || math.IsInf(xk, 0) {
return nil, errf("Interpolate: knot %d is not finite", k)
}
}
out := &Array{shape: query.Shape(), dt: Float}
out.alloc(query.Len())
n := xs.Len()
xv, yv, qv := floatPayload(xs), floatPayload(ys), floatPayload(query)
// The workers write disjoint output slots. A NaN query has no
// segment and aborts the call; the serial walk reported the first
// one in order, so the workers publish the smallest NaN index.
var nanIdx atomic.Int64
nanIdx.Store(int64(len(qv)) + 1)
parallelMin(len(qv), copyMinPerWorker, func(s, e int) {
ov := out.floats
for i := s; i < e; i++ {
q := qv[i]
if math.IsNaN(q) {
for {
cur := nanIdx.Load()
if int64(i) >= cur || nanIdx.CompareAndSwap(cur, int64(i)) {
break
}
}
continue
}
// The segment is located by bisecting the knots: the first
// knot at or above the query, which is the segment the
// ascending walk stopped on.
var lo int
if q <= xv[0] {
lo = 0
} else if q >= xv[n-1] {
lo = n - 2
} else {
lo0, hi0 := 1, n
for lo0 < hi0 {
mid := int(uint(lo0+hi0) >> 1)
if xv[mid] < q {
lo0 = mid + 1
} else {
hi0 = mid
}
}
lo = lo0 - 1
}
hi := lo + 1
x0, x1 := xv[lo], xv[hi]
y0, y1 := yv[lo], yv[hi]
t := 0.0
if x1 > x0 {
t = (q - x0) / (x1 - x0)
} else if q > x0 {
// A repeated knot at the range edge: the query sits past
// the zero-width segment, so it takes the upper end.
t = 1
}
if t < 0 {
t = 0
}
if t > 1 {
t = 1
}
// The result is float by construction, so the payload write needs
// no dtype dispatch.
ov[i] = y0 + t*(y1-y0)
}
})
if idx := int(nanIdx.Load()); idx <= len(qv) {
return nil, errf("Interpolate: query %d is NaN, which cannot be clamped", idx)
}
return out, nil
}
// EvaluatePolynomial evaluates coefficients (lowest power first) at
// the given points. Both operands are widened once and read off the
// payload slices, the same values the accessor calls returned without
// the per-element dispatch; the accumulation and the power walk are
// the ones the accessor loop kept.
func EvaluatePolynomial(coeffs, x *Array) (*Array, error) {
if coeffs.dt == Complex || x.dt == Complex {
return nil, errf("EvaluatePolynomial: complex inputs are not supported")
}
out := &Array{shape: x.Shape(), dt: Float}
out.alloc(x.Len())
cv, xv := floatPayload(coeffs), floatPayload(x)
for i := range xv {
pow := 1.0
var sum float64
for c := range cv {
sum += cv[c] * pow
pow *= xv[i]
}
out.floats[i] = sum
}
return out, nil
}
// MoveAxis moves an axis to a new position in the shape.
func MoveAxis(a *Array, from, to int) (*Array, error) {
if from < 0 || from >= a.NDim() || to < 0 || to >= a.NDim() {
return nil, errf("MoveAxis: axes %d to %d out of range for rank %d", from, to, a.NDim())
}
// The destination order: the remaining axes in sequence with from
// reinserted at to.
order := rangeN(a.NDim())
order = append(order[:from], order[from+1:]...)
order = append(order[:to], append([]int{from}, order[to:]...)...)
newShape := make([]int, a.NDim())
for k, d := range order {
newShape[k] = a.shape[d]
}
out := &Array{shape: newShape, dt: a.dt}
out.alloc(prodShape(newShape))
coord := make([]int, a.NDim())
srcCoord := make([]int, a.NDim())
for i := range a.Len() {
for d := range a.NDim() {
srcCoord[order[d]] = coord[d]
}
off := 0
for d := range a.NDim() {
off = off*a.shape[d] + srcCoord[d]
}
out.setFrom(i, a, off)
advanceOdometer(coord, newShape)
}
return out, nil
}
// floatPayload returns a's elements as a plain float64 slice: the
// payload itself for a contiguous float64 array, and otherwise a dense
// copy through the array's own accessor. Every value is exactly the one
// floatAt reports, so a strided view or a narrower dtype reads the same
// numbers; the searches that follow then index a slice instead of
// paying a dtype dispatch per probed element.
func floatPayload(a *Array) []float64 {
n := a.Len()
if a.dt == Float && a.isContiguous() {
return a.floats[:n]
}
out := make([]float64, n)
for i := range n {
out[i] = a.floatAt(i)
}
return out
}
// intPayload is floatPayload for the integer-class operands, whose
// comparisons must stay in int64: a contiguous int64 array aliases its
// payload, the fast path the int comparisons have always taken, and
// every other walk copies through intAt, which widens the whole integer
// class exactly (bool reads 0/1) and resolves a view's strides.
func intPayload(a *Array) []int64 {
n := a.Len()
if a.dt == Int && a.isContiguous() {
return a.ints[:n]
}
out := make([]int64, n)
for i := range n {
out[i] = a.intAt(i)
}
return out
}
// searchMinPerWorker is the per-worker needle floor for SearchSorted:
// a needle costs a bisection whose depth grows with the haystack, tens
// of nanoseconds at the sizes the search benchmarks reach, so a chunk of
// a few hundred needles already outweighs its worker's start-up while
// smaller chunks run on the calling goroutine.
const searchMinPerWorker = 512
// SearchSorted finds insertion positions for each needle in the sorted
// haystack so that order is preserved (rightmost rule: positions after
// equal elements). The haystack must be ascending; the position is the
// number of elements at or below the needle, found by bisecting the
// haystack rather than by walking the prefix, which selects the same
// index in a logarithmic number of comparisons. The whole integer class
// bisects its own payload natively in int64, where the widening of
// bool and every narrow integer is exact; float operands widen through
// floatAt, exact for every real dtype, because float64 would round
// int64 values onto one another above 2^53. The workers write disjoint
// output slots, so the split cannot move a single position.
func SearchSorted(haystack, needles *Array) (*Array, error) {
if haystack.dt == Complex || needles.dt == Complex {
return nil, errf("SearchSorted: complex arrays have no ordering")
}
n := haystack.Len()
out := &Array{shape: needles.Shape(), dt: Int}
out.alloc(needles.Len())
if n == 0 {
// Every position is zero, which is what the fresh payload holds.
return out, nil
}
// Both operands integer-class: compare natively in int64, because
// float64 would round the haystack edges and the needle onto one
// another above 2^53 and return the wrong insertion point.
if intClass(haystack.dt) && intClass(needles.dt) {
h, q := intPayload(haystack), intPayload(needles)
parallelMin(len(q), searchMinPerWorker, func(s, e int) {
oi := out.ints
for i := s; i < e; i++ {
// Bisection for the first element above the needle: lo
// ends at the count of elements at or below it, the
// rightmost position.
lo, hi := 0, n
for lo < hi {
mid := int(uint(lo+hi) >> 1)
if h[mid] <= q[i] {
lo = mid + 1
} else {
hi = mid
}
}
oi[i] = int64(lo)
}
})
return out, nil
}
h, q := floatPayload(haystack), floatPayload(needles)
parallelMin(len(q), searchMinPerWorker, func(s, e int) {
oi := out.ints
for i := s; i < e; i++ {
// <= (not <) keeps the position after equal elements, as the
// rightmost rule documented above promises. A NaN needle
// fails every comparison and lands at position zero.
lo, hi := 0, n
for lo < hi {
mid := int(uint(lo+hi) >> 1)
if h[mid] <= q[i] {
lo = mid + 1
} else {
hi = mid
}
}
oi[i] = int64(lo)
}
})
return out, nil
}
// AssignBins maps every value to its bin index given ascending bin
// edges: bin k covers [edges[k], edges[k+1]). Values below the first or
// above the last edge clamp to the outer bins, and a NaN value keeps
// the outermost bin, where the comparisons leave it. The whole integer
// class selects bins natively in int64, where the widening of bool and
// every narrow integer is exact; float operands widen through floatAt,
// because float64 would round an int value onto a bin edge above 2^53.
// The workers write disjoint output slots and every bin index is an
// exact integer selection, so the split cannot move a single bin.
func AssignBins(a *Array, edges *Array) (*Array, error) {
if a.dt == Complex || edges.dt == Complex {
return nil, errf("AssignBins: complex arrays are not supported")
}
m := edges.Len()
if m < 2 {
return nil, errf("AssignBins: edges need at least two values")
}
out := &Array{shape: a.Shape(), dt: Int}
out.alloc(a.Len())
// Both operands integer-class: select the bin natively in int64,
// because float64 would round the value onto a bin edge above 2^53
// and pick the wrong one.
if intClass(a.dt) && intClass(edges.dt) {
ev, av := intPayload(edges), intPayload(a)
parallelMin(len(av), copyMinPerWorker, func(s, e int) {
oi := out.ints
for i := s; i < e; i++ {
v := av[i]
bin := 0
if v >= ev[0] {
// The largest index whose edge sits at or below the
// value; the guard leaves below-range values in bin 0.
lo, hi := 0, m-1
for lo < hi {
mid := int(uint(lo+hi) >> 1)
if ev[mid] <= v {
lo = mid + 1
} else {
hi = mid
}
}
bin = lo - 1
}
oi[i] = int64(bin)
}
})
return out, nil
}
ev, av := floatPayload(edges), floatPayload(a)
parallelMin(len(av), copyMinPerWorker, func(s, e int) {
oi := out.ints
for i := s; i < e; i++ {
v := av[i]
bin := 0
// The guard keeps below-range values in bin 0; every value no
// edge compares at or below (a NaN) bisects down to lo zero
// and keeps the outermost bin, which is where the downward
// scan this replaces left it.
if !(v < ev[0]) {
lo, hi := 0, m-1
for lo < hi {
mid := int(uint(lo+hi) >> 1)
if ev[mid] <= v {
lo = mid + 1
} else {
hi = mid
}
}
bin = lo - 1
if lo == 0 {
bin = m - 2
}
}
oi[i] = int64(bin)
}
})
return out, nil
}
// centralSum folds the central cross sum Σ(x−mx)(y−my) through the
// canonical partition: the per-element arithmetic the plain chain
// kept, cut into the fixed blocks the folds use, one chain partial per
// block and the balanced tree over them. At and below one block the
// walk is the plain chain, bit for bit; past it the measured
// deviation from the exact referent drops by orders of magnitude (the
// central-sum test pins the numbers).
func centralSum(x, y []float64, mx, my float64) float64 {
n := len(x)
parts := foldParts(n)
if parts == 1 {
var acc float64
for i := range n {
acc += (x[i] - mx) * (y[i] - my)
}
return acc
}
partials := make([]float64, parts)
for c := range parts {
// The block's bounds are named locals: an unhoisted divided
// bound in the loop condition kept the walk off the fast path,
// three times the cost of the chain it replaced (measured).
lo, hi := c*n/parts, (c+1)*n/parts
var acc float64
for i := lo; i < hi; i++ {
acc += (x[i] - mx) * (y[i] - my)
}
partials[c] = acc
}
return treeSum(partials)
}
// centralMomentSums is centralSum carrying the two squared-deviation
// sums alongside the cross sum, the three Correlation needs. The
// partition and the per-element arithmetic are the same walk's.
func centralMomentSums(x, y []float64, mx, my float64) (num, dx2, dy2 float64) {
n := len(x)
parts := foldParts(n)
if parts == 1 {
for i := range n {
da := x[i] - mx
dbv := y[i] - my
num += da * dbv
dx2 += da * da
dy2 += dbv * dbv
}
return num, dx2, dy2
}
pn, px, py := make([]float64, parts), make([]float64, parts), make([]float64, parts)
for c := range parts {
lo, hi := c*n/parts, (c+1)*n/parts
var sn, sx, sy float64
for i := lo; i < hi; i++ {
da := x[i] - mx
dbv := y[i] - my
sn += da * dbv
sx += da * da
sy += dbv * dbv
}
pn[c], px[c], py[c] = sn, sx, sy
}
return treeSum(pn), treeSum(px), treeSum(py)
}
// Covariance computes the sample covariance of two equally sized
// 1-D samples (denominator n−1). The samples are widened once and the
// central sum folds through the canonical partition, the accuracy the
// central-sum test measured against the exact referent.
func Covariance(a, b *Array) (float64, error) {
if a.dt == Complex || b.dt == Complex {
return 0, errf("Covariance: complex samples are not supported")
}
if a.Len() != b.Len() || a.Len() < 2 {
return 0, errf("Covariance: samples must share length ≥ 2")
}
ma, err := Mean(a)
if err != nil {
return 0, err
}
mb, err := Mean(b)
if err != nil {
return 0, err
}
av, bv := floatPayload(a), floatPayload(b)
return centralSum(av, bv, ma, mb) / float64(a.Len()-1), nil
}
// Correlation computes the Pearson correlation coefficient of two
// 1-D samples: sum(dx·dy) / sqrt(sum dx² · sum dy²). The samples are
// widened once and the three central sums fold through the canonical
// partition, the accuracy the central-sum test measured against the
// exact referent.
func Correlation(a, b *Array) (float64, error) {
if a.dt == Complex || b.dt == Complex {
return 0, errf("Correlation: complex samples are not supported")
}
if a.Len() != b.Len() || a.Len() < 2 {
return 0, errf("Correlation: samples must share length ≥ 2")
}
ma, err := Mean(a)
if err != nil {
return 0, err
}
mb, err := Mean(b)
if err != nil {
return 0, err
}
av, bv := floatPayload(a), floatPayload(b)
num, da2, db2 := centralMomentSums(av, bv, ma, mb)
return num / math.Sqrt(da2*db2), nil
}
// Sign returns the sign of every element (-1, 0 or +1) as a float
// array; complex inputs error. A NaN element has no sign and reports 0.
// The elements are widened once and the sign read off the payload
// slice, the same values the accessor walk read without the
// per-element dispatch; the destination starts zeroed, so only the
// non-zero marks are written.
func Sign(a *Array) (*Array, error) {
if a.dt == Complex {
return nil, errf("Sign: complex arrays are not supported")
}
out := zeros(Float, a.Shape())
src := a
if !src.isContiguous() {
src = src.materialise()
}
vals := floatPayload(src)
d := out.floats
parallelMin(len(vals), elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
v := vals[i]
if v > 0 {
d[i] = 1
} else if v < 0 {
d[i] = -1
}
}
})
return out, nil
}
// IsNaN returns an int mask marking NaN elements. Complex arrays are
// not supported.
func IsNaN(a *Array) (*Array, error) {
if a.dt == Complex {
return nil, errf("IsNaN: complex arrays are not supported")
}
out := &Array{shape: a.Shape(), dt: Int}
out.alloc(a.Len())
// The widening happens once and the mask starts zeroed, so only the
// marking slots are written; the chunks are disjoint.
vals := floatPayload(a)
parallelMin(len(vals), copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
if vals[i] != vals[i] { // NaN != NaN
out.ints[i] = 1
}
}
})
return out, nil
}
// IsInf returns an int mask marking positive and negative infinity.
func IsInf(a *Array) (*Array, error) {
if a.dt == Complex {
return nil, errf("IsInf: complex arrays are not supported")
}
pos := math.Inf(1)
neg := math.Inf(-1)
out := &Array{shape: a.Shape(), dt: Int}
out.alloc(a.Len())
vals := floatPayload(a)
parallelMin(len(vals), copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
if v := vals[i]; v == pos || v == neg {
out.ints[i] = 1
}
}
})
return out, nil
}
// IsFinite returns an int mask marking finite values: neither NaN nor
// infinite.
func IsFinite(a *Array) (*Array, error) {
if a.dt == Complex {
return nil, errf("IsFinite: complex arrays are not supported")
}
pos := math.Inf(1)
neg := math.Inf(-1)
out := &Array{shape: a.Shape(), dt: Int}
out.alloc(a.Len())
vals := floatPayload(a)
parallelMin(len(vals), copyMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
if v := vals[i]; v == v && v != pos && v != neg {
out.ints[i] = 1
}
}
})
return out, nil
}
// LowerTriangle returns the lower triangular part of a square matrix as
// a copy with everything above the diagonal zeroed. Each row's kept
// prefix is one contiguous run, so a row copies whole rather than an
// element at a time.
func LowerTriangle(a *Array) (*Array, error) {
if a.NDim() != 2 || a.shape[0] != a.shape[1] {
return nil, errf("LowerTriangle: needs a square matrix, got %s", shapeText(a.shape))
}
n := a.shape[0]
out := ZerosLike(a)
if a.isContiguous() && out.isContiguous() {
for r := range n {
copyRun(out, a, r*n, r*n, r+1)
}
return out, nil
}
for r := range n {
for c := range n {
if c <= r {
out.setFrom(r*n+c, a, r*n+c)
}
}
}
return out, nil
}
// UpperTriangle returns the upper triangular part of a square matrix as
// a copy with everything below the diagonal zeroed. As in the lower
// builder, each row's kept suffix copies whole.
func UpperTriangle(a *Array) (*Array, error) {
if a.NDim() != 2 || a.shape[0] != a.shape[1] {
return nil, errf("UpperTriangle: needs a square matrix, got %s", shapeText(a.shape))
}
n := a.shape[0]
out := ZerosLike(a)
if a.isContiguous() && out.isContiguous() {
for r := range n {
copyRun(out, a, r*n+r, r*n+r, n-r)
}
return out, nil
}
for r := range n {
for c := range n {
if c >= r {
out.setFrom(r*n+c, a, r*n+c)
}
}
}
return out, nil
}
// copyRun copies run elements of src at srcOff to dst at dstOff. Both
// arrays carry the same dtype and are contiguous; the dispatch happens
// once per run rather than once per element, and the run itself is a
// single block copy. copyRun answers for every element type the package
// stores: bool, int8, uint8, int16, uint16, int32, uint32, int,
// float16, float32, float and complex, so the slice and join kernels
// take their block-copy path for every dtype.
func copyRun(dst, src *Array, dstOff, srcOff, run int) {
switch dst.dt {
case Int:
copy(dst.ints[dstOff:dstOff+run], src.ints[srcOff:srcOff+run])
case Float16:
copy(dst.halves[dstOff:dstOff+run], src.halves[srcOff:srcOff+run])
case Float32:
copy(dst.floats32[dstOff:dstOff+run], src.floats32[srcOff:srcOff+run])
case Float:
copy(dst.floats[dstOff:dstOff+run], src.floats[srcOff:srcOff+run])
case Complex:
copy(dst.complexes[dstOff:dstOff+run], src.complexes[srcOff:srcOff+run])
case Bool:
copy(dst.bools[dstOff:dstOff+run], src.bools[srcOff:srcOff+run])
case Int8:
copy(dst.i8s[dstOff:dstOff+run], src.i8s[srcOff:srcOff+run])
case Uint8:
copy(dst.u8s[dstOff:dstOff+run], src.u8s[srcOff:srcOff+run])
case Int16:
copy(dst.i16s[dstOff:dstOff+run], src.i16s[srcOff:srcOff+run])
case Uint16:
copy(dst.u16s[dstOff:dstOff+run], src.u16s[srcOff:srcOff+run])
case Int32:
copy(dst.i32s[dstOff:dstOff+run], src.i32s[srcOff:srcOff+run])
case Uint32:
copy(dst.u32s[dstOff:dstOff+run], src.u32s[srcOff:srcOff+run])
}
}
// prodShape returns the product of a shape slice; 1 for an empty slice
// (matches Reshape's behaviour for a 0-D tensor).
func prodShape(s []int) int {
p := 1
for _, v := range s {
p *= v
}
return p
}