Files

2047 lines
60 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import "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
}