Files
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

1564 lines
46 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"cmp"
"math"
"slices"
"sync/atomic"
)
// Axis-based reductions: the softmax, normalisation and loss
// primitive. Each reduces along one dimension and keeps the others in
// order; the result rank is NDim-1. Reducing the only dimension of a 1-D
// array is an error pointing at the global variant: tensor has no 0-d
// arrays.
// axisOp selects the fold an axis reduction applies to each line.
type axisOp uint8
const (
opSum axisOp = iota
opMean
opMin
opMax
)
// SumAxis returns the sums along the given dimension.
func SumAxis(a *Array, dim int) (*Array, error) {
return a.reduceAxis(dim, "SumAxis", opSum)
}
// MinAxis returns the smallest values along the given dimension; NaN
// elements never win, complex arrays have no ordering.
func MinAxis(a *Array, dim int) (*Array, error) {
return a.reduceAxis(dim, "MinAxis", opMin)
}
// MaxAxis returns the largest values along the given dimension; NaN
// elements never win, complex arrays have no ordering.
func MaxAxis(a *Array, dim int) (*Array, error) {
return a.reduceAxis(dim, "MaxAxis", opMax)
}
// MeanAxis returns the float means along the given dimension; complex
// arrays have no float mean.
func MeanAxis(a *Array, dim int) (*Array, error) {
if a.dt == Complex {
return nil, errf("MeanAxis: complex arrays have no float mean")
}
return a.reduceAxis(dim, "MeanAxis", opMean)
}
// reduceAxis folds every element into the accumulator slot addressed by
// the source coordinate with dim dropped. The walk is line based, like
// the Norm and scanDim kernels: a line is the run of a.shape[dim]
// elements that share every surviving coordinate, so the destination
// index collapses to b*stride + s and the per-element odometer
// disappears. Whole lines are handed to each worker, which makes every
// accumulator slot single-writer, so there is no merge phase; the
// fan-out is capped by splitCapped so a small fold never pays a
// core-count spawn bill.
// The float and complex sums fold each line through the canonical
// partition the global sums use: fixed blocks of the line, the block
// partials combined through the balanced tree. The partition follows
// from the line length alone, so a line's answer is the same bits
// whatever the worker split, and a single-line fold answers the bits of
// Sum over the same elements; against the single chain it replaced the
// tree holds full accuracy on lines long enough for a chain's roundings
// to pile up, measured against big.Float in the accuracy test. The
// integer sums stay plain chains: wrapping addition is exact under any
// grouping. Extrema seed each line from its first non-NaN element, NaN
// candidates never win (they fail every comparison), and a line whose
// every element is NaN never seeds and reports NaN. mean divides by the
// reduced dimension afterwards and accumulates in float regardless of
// the input dtype; everything else keeps its dtype. float16 and float32
// reductions fold in a float64 scratch and narrow once.
func (a *Array) reduceAxis(dim int, name string, op axisOp) (*Array, error) {
if dim < 0 || dim >= a.NDim() {
return nil, errf("%s: dimension %d is out of range for shape %s", name, dim, shapeText(a.shape))
}
if a.NDim() == 1 {
return nil, errf("%s: reducing the only dimension of a 1-D array, use the global variant", name)
}
if a.dt == Complex && op != opSum {
return nil, errf("%s: complex arrays have no ordering", name)
}
newShape := make([]int, 0, a.NDim()-1)
newShape = append(newShape, a.shape[:dim]...)
newShape = append(newShape, a.shape[dim+1:]...)
total := 1
for _, d := range newShape {
total *= d
}
acc := &Array{shape: newShape, dt: a.dt}
if op == opMean {
acc.dt = Float
} else if intClass(a.dt) && a.dt != Int {
// The scalar reductions answer Int scalars for the whole integer
// class; the axis folds mirror that with Int accumulators fed by
// exact widenings.
acc.dt = Int
}
// float16 and float32 reductions fold in a float64 scratch and
// narrow once; every other dtype folds straight into the
// accumulator.
scratch := acc
if (a.dt == Float16 || a.dt == Float32) && op != opMean {
scratch = &Array{shape: newShape, dt: Float}
}
scratch.alloc(total)
stride := 1
for k := dim + 1; k < a.NDim(); k++ {
stride *= a.shape[k]
}
line := a.shape[dim]
perLine := stride * line
lines := 0
if perLine > 0 {
lines = a.Len() / perLine
}
seed := op == opMin || op == opMax
// One line per accumulator slot, whole lines per worker: the slot at
// b*stride + s is written exactly once, by the worker that owns
// line b. The fan-out is capped by splitCapped so every worker
// carries at least reduceSplitFloor elements.
work := lines * perLine
switch a.dt {
case Int:
switch {
case seed:
foldExtremeInt(a.ints, scratch.ints, op == opMin, lines, work, perLine, stride, line)
case op == opMean:
foldMeanInt(a.ints, scratch.floats, lines, work, perLine, stride, line)
default:
foldSum(a.ints, scratch.ints, lines, work, perLine, stride, line)
}
case Bool:
switch {
case seed:
foldExtremeAxisBool(a.bools, scratch.ints, op == opMin, lines, work, perLine, stride, line)
case op == opMean:
foldMeanAxisBool(a.bools, scratch.floats, lines, work, perLine, stride, line)
default:
foldSumAxisBool(a.bools, scratch.ints, lines, work, perLine, stride, line)
}
case Int8:
axisFoldNarrow(a.i8s, scratch, op, seed, lines, work, perLine, stride, line)
case Uint8:
axisFoldNarrow(a.u8s, scratch, op, seed, lines, work, perLine, stride, line)
case Int16:
axisFoldNarrow(a.i16s, scratch, op, seed, lines, work, perLine, stride, line)
case Uint16:
axisFoldNarrow(a.u16s, scratch, op, seed, lines, work, perLine, stride, line)
case Int32:
axisFoldNarrow(a.i32s, scratch, op, seed, lines, work, perLine, stride, line)
case Uint32:
axisFoldNarrow(a.u32s, scratch, op, seed, lines, work, perLine, stride, line)
case Float16:
if seed {
foldExtremeHalf(a.halves, scratch.floats, op == opMin, lines, work, perLine, stride, line)
} else {
foldSumAxisF16(a.halves, scratch.floats, lines, work, perLine, stride, line)
}
case Float32:
if seed {
foldExtremeFloat(a.floats32, scratch.floats, op == opMin, lines, work, perLine, stride, line)
} else {
foldSumAxisF32(a.floats32, scratch.floats, lines, work, perLine, stride, line)
}
case Float:
if seed {
foldExtremeFloat(a.floats, scratch.floats, op == opMin, lines, work, perLine, stride, line)
} else {
foldSumAxis(a.floats, scratch.floats, lines, work, perLine, stride, line)
}
default:
foldSumAxis(a.complexes, scratch.complexes, lines, work, perLine, stride, line)
}
// A zero-length reduction dimension leaves every line empty, so the
// walk never runs and nothing can seed: each float slot reports the
// missing value NaN, and int slots keep their zeros. This is the
// same outcome the all-NaN fill rule produced, and only float
// accumulators report NaN because int elements always seed.
if seed && line == 0 {
// Only the float64 accumulator can land here unseeded: int
// elements always seed the first line, and every scratch this
// path allocates is the float64 one (a float32 or half input
// widens into it), so no narrower scratch exists to fill.
for d := range total {
if scratch.dt == Float {
scratch.floats[d] = math.NaN()
}
}
}
if op == opMean {
line := float64(a.shape[dim])
for k := range acc.floats {
acc.floats[k] /= line
}
return acc, nil
}
if scratch != acc {
if acc.dt == Float16 {
acc.halves = make([]uint16, total)
for k, v := range scratch.floats {
acc.halves[k] = HalfFromFloat64(v)
}
} else {
acc.floats32 = make([]float32, total)
for k, v := range scratch.floats {
acc.floats32[k] = float32(v)
}
}
}
return acc, nil
}
// reduceSplitFloor is the element count every spawned fold worker must
// carry before splitCapped grants it a goroutine: the folds stream about
// one element per instruction, so the spawn and its synchronisation only
// amortise above this many of them, and a smaller chunk costs more to
// schedule than to run (the axis sweep splits between 4,096 elements,
// where the lone walk wins outright, and 16,384, where the capped
// fan-out is ahead).
const reduceSplitFloor = 16_384
// foldSum adds every line's int64 elements into the slot the line
// shares, in ascending element order. Wrapping addition is associative,
// so no grouping can move a bit and the chain needs no partition.
func foldSum(src, dst []int64, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
// Contiguous lines: one slice expression bounds-checks the
// whole line and the range walk streams it.
for b := ls; b < le; b++ {
var m int64
for _, v := range src[b*perLine : b*perLine+line] {
m += v
}
dst[b] = m
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
var m int64
for off := range line {
m += src[base+off*stride+s]
}
dst[b*stride+s] = m
}
}
})
}
// foldBlock is one canonical block's fold over the line elements
// src[base+off*stride] for off in [0, count): four interleaved chains
// combined as ((s0+s1)+(s2+s3)), the pairing foldRange keeps. A stride
// of one visits the block in exactly foldRange's order, so a contiguous
// line's block value is the value reduce.go's block fold gives.
func foldBlock[N float64 | complex128](src []N, base, stride, count int) N {
var s0, s1, s2, s3 N
i := 0
for ; i+4 <= count; i += 4 {
o := base + i*stride
s0 += src[o]
s1 += src[o+stride]
s2 += src[o+2*stride]
s3 += src[o+3*stride]
}
for ; i < count; i++ {
s0 += src[base+i*stride]
}
return (s0 + s1) + (s2 + s3)
}
// lineFoldParts folds one line through the canonical partition: fixed
// block boundaries at c·line/foldParts(line), each block's partial from
// the block closure, the partials combined through the balanced tree.
// The line length alone picks the shape, so the value is the same bits
// whatever the worker split, and a single-block line is one block fold.
func lineFoldParts[N float64 | complex128](line int, block func(lo, hi int) N) N {
parts := foldParts(line)
if parts == 1 {
return block(0, line)
}
partials := make([]N, parts)
for c := range parts {
lo, hi := c*line/parts, (c+1)*line/parts
partials[c] = block(lo, hi)
}
return treeSum(partials)
}
// foldSumAxis folds every float64 or complex128 line through the
// canonical partition the global sums use: the line's blocks from
// foldBlock, the partials combined through treeSum. Against the single
// chain it replaced, the tree holds its accuracy on lines long enough
// for a chain's roundings to pile up and loses nothing on the short
// ones; on the large-plus-small counterpoint, where a chain's running
// total swallows the small elements outright, the tree's blocks keep
// them. Contiguous lines reuse reduce.go's own block folds, so a
// single-line fold answers Sum's exact bits.
func foldSumAxis[N float64 | complex128](src, dst []N, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
row := src[b*perLine : b*perLine+line]
dst[b] = lineFoldParts(line, func(lo, hi int) N {
return foldBlock(row, lo, 1, hi-lo)
})
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
dst[b*stride+s] = lineFoldParts(line, func(lo, hi int) N {
return foldBlock(src, base+lo*stride+s, stride, hi-lo)
})
}
}
})
}
// foldBlockWidenF32 is foldBlock over a float32 payload: every element
// widens exactly, so the chains see the values an accessor read would.
func foldBlockWidenF32(src []float32, base, stride, count int) float64 {
var s0, s1, s2, s3 float64
i := 0
for ; i+4 <= count; i += 4 {
o := base + i*stride
s0 += float64(src[o])
s1 += float64(src[o+stride])
s2 += float64(src[o+2*stride])
s3 += float64(src[o+3*stride])
}
for ; i < count; i++ {
s0 += float64(src[base+i*stride])
}
return (s0 + s1) + (s2 + s3)
}
// foldBlockWidenF16 is foldBlockWidenF32 for a uint16 half payload: the
// raw bits widen with HalfToFloat64, never cast as integers.
func foldBlockWidenF16(src []uint16, base, stride, count int) float64 {
var s0, s1, s2, s3 float64
i := 0
for ; i+4 <= count; i += 4 {
o := base + i*stride
s0 += HalfToFloat64(src[o])
s1 += HalfToFloat64(src[o+stride])
s2 += HalfToFloat64(src[o+2*stride])
s3 += HalfToFloat64(src[o+3*stride])
}
for ; i < count; i++ {
s0 += HalfToFloat64(src[base+i*stride])
}
return (s0 + s1) + (s2 + s3)
}
// foldSumAxisF32 is foldSumAxis for a float32 payload folded into a
// float64 scratch through the same canonical partition; every widening
// is exact, so the values match the accessor walk.
func foldSumAxisF32(src []float32, dst []float64, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
row := src[b*perLine : b*perLine+line]
dst[b] = lineFoldParts(line, func(lo, hi int) float64 {
return foldBlockWidenF32(row, lo, 1, hi-lo)
})
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
dst[b*stride+s] = lineFoldParts(line, func(lo, hi int) float64 {
return foldBlockWidenF32(src, base+lo*stride+s, stride, hi-lo)
})
}
}
})
}
// foldSumAxisF16 is foldSumAxisF32 for a half payload's raw bit
// patterns, widened with HalfToFloat64 exactly as they are read.
func foldSumAxisF16(src []uint16, dst []float64, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
row := src[b*perLine : b*perLine+line]
dst[b] = lineFoldParts(line, func(lo, hi int) float64 {
return foldBlockWidenF16(row, lo, 1, hi-lo)
})
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
dst[b*stride+s] = lineFoldParts(line, func(lo, hi int) float64 {
return foldBlockWidenF16(src, base+lo*stride+s, stride, hi-lo)
})
}
}
})
}
// foldMeanInt is foldSumWiden for an int payload: the widened elements
// accumulate in float64 exactly as the per-element walk did.
func foldMeanInt(src []int64, dst []float64, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
var m float64
for _, v := range src[b*perLine : b*perLine+line] {
m += float64(v)
}
dst[b] = m
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
var m float64
for off := range line {
m += float64(src[base+off*stride+s])
}
dst[b*stride+s] = m
}
}
})
}
// foldExtremeInt writes the smallest (wantMin) or largest element of
// every line. Int lines always seed: the first element opens the line
// and the rest compare against it, in ascending order, exactly as the
// per-element walk compared them.
func foldExtremeInt(src, dst []int64, wantMin bool, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
row := src[b*perLine : b*perLine+line]
m := row[0]
if wantMin {
for _, v := range row[1:] {
if v < m {
m = v
}
}
} else {
for _, v := range row[1:] {
if v > m {
m = v
}
}
}
dst[b] = m
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
m := src[base+s]
if wantMin {
for off := 1; off < line; off++ {
if v := src[base+off*stride+s]; v < m {
m = v
}
}
} else {
for off := 1; off < line; off++ {
if v := src[base+off*stride+s]; v > m {
m = v
}
}
}
dst[b*stride+s] = m
}
}
})
}
// foldExtremeFloat is foldExtremeInt for the float payloads, folded into
// a float64 scratch: the line seeds from its first non-NaN element, NaN
// candidates fail every comparison and can neither seed nor win, and a
// line whose every element is NaN never seeds and reports NaN. The seed
// scan and the compare walk visit the elements in the same order the
// single-loop walk did, so every selection is unchanged; splitting the
// two phases only lifts the per-element seeded test off the hot loop.
func foldExtremeFloat[F float32 | float64](src []F, dst []float64, wantMin bool, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
row := src[b*perLine : b*perLine+line]
i, m := 0, 0.0
for i < line {
if v := float64(row[i]); v == v {
m = v
break
}
i++
}
if i == line {
// Every element of the line is NaN.
dst[b] = math.NaN()
continue
}
i++ // step past the seed
if wantMin {
for _, v := range row[i:] {
if w := float64(v); w < m {
m = w
}
}
} else {
for _, v := range row[i:] {
if w := float64(v); w > m {
m = w
}
}
}
dst[b] = m
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
i, m := 0, 0.0
for i < line {
if v := float64(src[base+i*stride+s]); v == v {
m = v
break
}
i++
}
if i == line {
// Every element of the line is NaN.
dst[b*stride+s] = math.NaN()
continue
}
i++ // step past the seed
if wantMin {
for off := i; off < line; off++ {
if w := float64(src[base+off*stride+s]); w < m {
m = w
}
}
} else {
for off := i; off < line; off++ {
if w := float64(src[base+off*stride+s]); w > m {
m = w
}
}
}
dst[b*stride+s] = m
}
}
})
}
// foldExtremeHalf is foldExtremeFloat for a uint16 half payload: the
// raw bits widen with HalfToFloat64, never cast as integers, so the
// seed scan and the comparisons see exactly the values floatAt would
// hand over.
func foldExtremeHalf(src []uint16, dst []float64, wantMin bool, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
row := src[b*perLine : b*perLine+line]
i, m := 0, 0.0
for i < line {
if v := HalfToFloat64(row[i]); v == v {
m = v
break
}
i++
}
if i == line {
// Every element of the line is NaN.
dst[b] = math.NaN()
continue
}
i++ // step past the seed
if wantMin {
for _, v := range row[i:] {
if w := HalfToFloat64(v); w < m {
m = w
}
}
} else {
for _, v := range row[i:] {
if w := HalfToFloat64(v); w > m {
m = w
}
}
}
dst[b] = m
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
i, m := 0, 0.0
for i < line {
if v := HalfToFloat64(src[base+i*stride+s]); v == v {
m = v
break
}
i++
}
if i == line {
// Every element of the line is NaN.
dst[b*stride+s] = math.NaN()
continue
}
i++ // step past the seed
if wantMin {
for off := i; off < line; off++ {
if w := HalfToFloat64(src[base+off*stride+s]); w < m {
m = w
}
}
} else {
for off := i; off < line; off++ {
if w := HalfToFloat64(src[base+off*stride+s]); w > m {
m = w
}
}
}
dst[b*stride+s] = m
}
}
})
}
// axisFoldNarrow runs one narrow integer payload through reduceAxis's
// fold selection: sums and extrema widen exactly into the Int
// accumulator, means widen into the float64 scratch.
func axisFoldNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, scratch *Array, op axisOp, seed bool, lines, work, perLine, stride, line int) {
switch {
case seed:
foldExtremeAxisNarrow(src, scratch.ints, op == opMin, lines, work, perLine, stride, line)
case op == opMean:
foldMeanAxisNarrow(src, scratch.floats, lines, work, perLine, stride, line)
default:
foldSumAxisNarrow(src, scratch.ints, lines, work, perLine, stride, line)
}
}
// foldSumAxisNarrow is foldSum for a narrow integer payload folded into
// an int64 accumulator: every element widens exactly.
func foldSumAxisNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, dst []int64, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
var m int64
for _, v := range src[b*perLine : b*perLine+line] {
m += int64(v)
}
dst[b] = m
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
var m int64
for off := range line {
m += int64(src[base+off*stride+s])
}
dst[b*stride+s] = m
}
}
})
}
// foldSumAxisBool is foldSumAxisNarrow for a bool payload: each line
// counts its true elements.
func foldSumAxisBool(src []bool, dst []int64, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
var m int64
for _, v := range src[b*perLine : b*perLine+line] {
if v {
m++
}
}
dst[b] = m
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
var m int64
for off := range line {
if src[base+off*stride+s] {
m++
}
}
dst[b*stride+s] = m
}
}
})
}
// foldMeanAxisNarrow is foldMeanInt for a narrow integer payload: the
// widened elements accumulate in float64 exactly as the per-element walk
// did.
func foldMeanAxisNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, dst []float64, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
var m float64
for _, v := range src[b*perLine : b*perLine+line] {
m += float64(v)
}
dst[b] = m
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
var m float64
for off := range line {
m += float64(src[base+off*stride+s])
}
dst[b*stride+s] = m
}
}
})
}
// foldMeanAxisBool counts a bool line into the float64 mean scratch.
func foldMeanAxisBool(src []bool, dst []float64, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
var m float64
for _, v := range src[b*perLine : b*perLine+line] {
if v {
m++
}
}
dst[b] = m
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
var m float64
for off := range line {
if src[base+off*stride+s] {
m++
}
}
dst[b*stride+s] = m
}
}
})
}
// foldExtremeAxisNarrow is foldExtremeInt for a narrow integer payload:
// the line compares in the payload's own type, so no value ever meets a
// float64 rounding, and the winner widens exactly into the Int
// accumulator.
func foldExtremeAxisNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, dst []int64, wantMin bool, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
row := src[b*perLine : b*perLine+line]
m := row[0]
if wantMin {
for _, v := range row[1:] {
if v < m {
m = v
}
}
} else {
for _, v := range row[1:] {
if v > m {
m = v
}
}
}
dst[b] = int64(m)
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
m := src[base+s]
if wantMin {
for off := 1; off < line; off++ {
if v := src[base+off*stride+s]; v < m {
m = v
}
}
} else {
for off := 1; off < line; off++ {
if v := src[base+off*stride+s]; v > m {
m = v
}
}
}
dst[b*stride+s] = int64(m)
}
}
})
}
// boolToInt64 widens a bool to the 0/1 the int64 accumulators of the
// extrema family carry.
func boolToInt64(b bool) int64 {
if b {
return 1
}
return 0
}
// foldExtremeAxisBool is foldExtremeAxisNarrow for a bool payload:
// false below true, widened to the 0/1 the Int accumulator carries. Go
// orders no bool with < or >, so the strict improvement writes its own
// logic.
func foldExtremeAxisBool(src []bool, dst []int64, wantMin bool, lines, work, perLine, stride, line int) {
splitCapped(lines, work, reduceSplitFloor, func(ls, le int) {
if stride == 1 {
for b := ls; b < le; b++ {
row := src[b*perLine : b*perLine+line]
m := row[0]
if wantMin {
for _, v := range row[1:] {
if !v && m {
m = v
}
}
} else {
for _, v := range row[1:] {
if v && !m {
m = v
}
}
}
dst[b] = boolToInt64(m)
}
return
}
for b := ls; b < le; b++ {
base := b * perLine
for s := range stride {
m := src[base+s]
if wantMin {
for off := 1; off < line; off++ {
if v := src[base+off*stride+s]; !v && m {
m = v
}
}
} else {
for off := 1; off < line; off++ {
if v := src[base+off*stride+s]; v && !m {
m = v
}
}
}
dst[b*stride+s] = boolToInt64(m)
}
}
})
}
// ArgMax returns the index of the largest element of a 1-D array; NaN
// elements are skipped as missing. An empty or all-NaN array, a complex
// array, or a non-1-D shape is an error.
func ArgMax(a *Array) (int, error) {
return a.argExtreme("ArgMax", false)
}
// ArgMin returns the index of the smallest element of a 1-D array; NaN
// elements are skipped as missing.
func ArgMin(a *Array) (int, error) {
return a.argExtreme("ArgMin", true)
}
func (a *Array) argExtreme(name string, wantMin bool) (int, error) {
if a.NDim() != 1 {
return 0, errf("%s: needs a 1-D array, got shape %s", name, shapeText(a.shape))
}
if a.dt == Complex {
return 0, errf("%s: complex arrays have no ordering", name)
}
// Both dtype walks read the payload at the logical flat index, so a
// strided view is materialised first; a contiguous array is returned
// unchanged, keeping the hot path allocation-free.
a = a.materialise()
// An integer-class array compares in its own payload type: the reason
// the int64 walk exists is that floatAt rounds above 2^53, and the
// narrow widths and bool keep the same native-comparison contract.
if intClass(a.dt) {
n := a.Len()
switch a.dt {
case Int:
best := -1
for i := range n {
v := a.ints[i]
if best < 0 {
best = i
continue
}
b := a.ints[best]
if (wantMin && v < b) || (!wantMin && v > b) {
best = i
}
}
if best < 0 {
return 0, errf("%s: the array is empty", name)
}
return best, nil
case Bool:
return argExtremeIndexBool(a.bools, n, name, wantMin)
case Int8:
return argExtremeIndex(a.i8s, n, name, wantMin)
case Uint8:
return argExtremeIndex(a.u8s, n, name, wantMin)
case Int16:
return argExtremeIndex(a.i16s, n, name, wantMin)
case Uint16:
return argExtremeIndex(a.u16s, n, name, wantMin)
case Int32:
return argExtremeIndex(a.i32s, n, name, wantMin)
default:
return argExtremeIndex(a.u32s, n, name, wantMin)
}
}
best := -1
for i := range a.Len() {
v := a.floatAt(i)
if v != v {
continue
}
if best < 0 {
best = i
continue
}
b := a.floatAt(best)
if (wantMin && v < b) || (!wantMin && v > b) {
best = i
}
}
if best < 0 {
return 0, errf("%s: every element is NaN", name)
}
return best, nil
}
// argExtremeIndex walks the first n elements of a narrow integer payload
// comparing in the payload's own type; such a payload never holds a NaN,
// so the first element always seeds the walk.
func argExtremeIndex[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, n int, name string, wantMin bool) (int, error) {
if n == 0 {
return 0, errf("%s: the array is empty", name)
}
best := 0
for i := 1; i < n; i++ {
if (wantMin && src[i] < src[best]) || (!wantMin && src[i] > src[best]) {
best = i
}
}
return best, nil
}
// argExtremeIndexBool is argExtremeIndex for a bool payload: Go orders
// no bool with < or >, so the strict improvement writes its own logic,
// and a tie keeps the earlier index exactly as the integer walk does.
func argExtremeIndexBool(src []bool, n int, name string, wantMin bool) (int, error) {
if n == 0 {
return 0, errf("%s: the array is empty", name)
}
best := 0
for i := 1; i < n; i++ {
if (wantMin && !src[i] && src[best]) || (!wantMin && src[i] && !src[best]) {
best = i
}
}
return best, nil
}
// ArgMaxAxis returns the indices of the maximum values along the given
// dimension. The result is an Int array with the same shape as the
// receiver except that the chosen dimension is dropped: the same
// reduction shape SumAxis / MaxAxis / MinAxis produce. Reducing the only
// dimension of a 1-D array is an error; use ArgMax instead.
func ArgMaxAxis(a *Array, dim int) (*Array, error) {
return a.argExtremeAxis(dim, false)
}
// ArgMinAxis returns the indices of the minimum values along the given
// dimension.
func ArgMinAxis(a *Array, dim int) (*Array, error) {
return a.argExtremeAxis(dim, true)
}
// argExtremeAxis is the shared implementation behind ArgMaxAxis and
// ArgMinAxis. The result drops the chosen dimension, like SumAxis /
// MaxAxis, and each element is the position of the extreme along it.
// NaN elements are skipped as missing, mirroring the 1-D ArgMax: each
// line seeds from its first non-NaN element, and a line with no finite
// element at all is an error.
func (a *Array) argExtremeAxis(dim int, wantMin bool) (*Array, error) {
name := "ArgMaxAxis"
global := "ArgMax"
if wantMin {
name = "ArgMinAxis"
global = "ArgMin"
}
if dim < 0 || dim >= a.NDim() {
return nil, errf("%s: dimension %d out of range for shape %s", name, dim, shapeText(a.shape))
}
if a.dt == Complex {
return nil, errf("%s: complex arrays have no ordering", name)
}
if a.NDim() == 1 {
return nil, errf("%s: reducing the only dimension of a 1-D array, use %s", name, global)
}
if a.Len() == 0 {
return nil, errf("%s: an empty array has no result", name)
}
newShape := reduceShape(a.shape, dim)
out := &Array{shape: newShape, dt: Int}
total := 1
for _, d := range newShape {
total *= d
}
out.alloc(total)
stride := 1
for k := a.NDim() - 1; k > dim; k-- {
stride *= a.shape[k]
}
perLine := a.shape[dim] * stride
lines := a.Len() / perLine
lineLen := a.shape[dim]
var allNaN atomic.Bool
// The dtype dispatch sits outside the walk: elements come straight
// from the payload, int and float32 widening to float64 exactly.
switch a.dt {
case Int:
src := a.ints
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
for line := ls; line < le; line++ {
base := line * perLine
for post := range stride {
// Compared as int64: widening to float64 would round
// above 2^53 and pick the wrong element of two
// neighbours. An int is never NaN, so the first
// element always seeds.
best, bestVal := 0, src[base+post]
for k := 1; k < lineLen; k++ {
if v := src[base+k*stride+post]; wantMin {
if v < bestVal {
best, bestVal = k, v
}
} else if v > bestVal {
best, bestVal = k, v
}
}
out.ints[line*stride+post] = int64(best)
}
}
})
case Bool:
argExtremeAxisLineBool(a.bools, a.Len(), out, wantMin, lines, perLine, stride, lineLen)
case Int8:
argExtremeAxisLine(a.i8s, a.Len(), out, wantMin, lines, perLine, stride, lineLen)
case Uint8:
argExtremeAxisLine(a.u8s, a.Len(), out, wantMin, lines, perLine, stride, lineLen)
case Int16:
argExtremeAxisLine(a.i16s, a.Len(), out, wantMin, lines, perLine, stride, lineLen)
case Uint16:
argExtremeAxisLine(a.u16s, a.Len(), out, wantMin, lines, perLine, stride, lineLen)
case Int32:
argExtremeAxisLine(a.i32s, a.Len(), out, wantMin, lines, perLine, stride, lineLen)
case Uint32:
argExtremeAxisLine(a.u32s, a.Len(), out, wantMin, lines, perLine, stride, lineLen)
case Float16:
src := a.halves
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
for line := ls; line < le; line++ {
base := line * perLine
for post := range stride {
// Seed from the first non-NaN element; NaN candidates
// never enter the comparison. The widening is exact,
// so the half values compare exactly in float64.
best := -1
var bestVal float64
for k := range lineLen {
v := HalfToFloat64(src[base+k*stride+post])
if v != v { // NaN is skipped as missing
continue
}
if best < 0 {
best, bestVal = k, v
continue
}
if (wantMin && v < bestVal) || (!wantMin && v > bestVal) {
best, bestVal = k, v
}
}
if best < 0 {
// Every element of the line is NaN.
allNaN.Store(true)
continue
}
out.ints[line*stride+post] = int64(best)
}
}
})
case Float32:
src := a.floats32
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
for line := ls; line < le; line++ {
base := line * perLine
for post := range stride {
// Seed from the first non-NaN element; NaN candidates
// never enter the comparison.
best := -1
var bestVal float64
for k := range lineLen {
v := float64(src[base+k*stride+post])
if v != v { // NaN is skipped as missing
continue
}
if best < 0 {
best, bestVal = k, v
continue
}
if (wantMin && v < bestVal) || (!wantMin && v > bestVal) {
best, bestVal = k, v
}
}
if best < 0 {
// Every element of the line is NaN.
allNaN.Store(true)
continue
}
out.ints[line*stride+post] = int64(best)
}
}
})
default:
src := a.floats
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
for line := ls; line < le; line++ {
base := line * perLine
for post := range stride {
// Seed from the first non-NaN element; NaN candidates
// never enter the comparison.
best := -1
var bestVal float64
for k := range lineLen {
v := src[base+k*stride+post]
if v != v { // NaN is skipped as missing
continue
}
if best < 0 {
best, bestVal = k, v
continue
}
if (wantMin && v < bestVal) || (!wantMin && v > bestVal) {
best, bestVal = k, v
}
}
if best < 0 {
// Every element of the line is NaN.
allNaN.Store(true)
continue
}
out.ints[line*stride+post] = int64(best)
}
}
})
}
if allNaN.Load() {
return nil, errf("%s: every element along dimension %d is NaN", name, dim)
}
return out, nil
}
// argExtremeAxisLine is argExtremeAxis's walk for a narrow integer
// payload: each line seeds from its first element and the candidates
// compare in the payload's own type, so no value ever meets a float
// comparison; the winning position lands in the Int result unchanged.
func argExtremeAxisLine[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, n int, out *Array, wantMin bool, lines, perLine, stride, lineLen int) {
splitCapped(lines, n, reduceSplitFloor, func(ls, le int) {
for line := ls; line < le; line++ {
base := line * perLine
for post := range stride {
best, bestVal := 0, src[base+post]
for k := 1; k < lineLen; k++ {
if v := src[base+k*stride+post]; wantMin {
if v < bestVal {
best, bestVal = k, v
}
} else if v > bestVal {
best, bestVal = k, v
}
}
out.ints[line*stride+post] = int64(best)
}
}
})
}
// argExtremeAxisLineBool is argExtremeAxisLine for a bool payload: Go
// orders no bool with < or >, so the strict improvement writes its own
// logic, and a tie keeps the earlier position.
func argExtremeAxisLineBool(src []bool, n int, out *Array, wantMin bool, lines, perLine, stride, lineLen int) {
splitCapped(lines, n, reduceSplitFloor, func(ls, le int) {
for line := ls; line < le; line++ {
base := line * perLine
for post := range stride {
best, bestVal := 0, src[base+post]
for k := 1; k < lineLen; k++ {
v := src[base+k*stride+post]
if wantMin && !v && bestVal {
best, bestVal = k, v
} else if !wantMin && v && !bestVal {
best, bestVal = k, v
}
}
out.ints[line*stride+post] = int64(best)
}
}
})
}
// TopK returns the top-k values and their original indices along the
// given dimension, sorted descending by value. The result preserves
// the input shape except the chosen dimension is reduced to k. NaN
// elements never rank: they sort to the end of the output when the
// line holds fewer than k finite values. For a 1-D input it returns
// two 1-D arrays of length k.
func TopK(a *Array, k int, dim int) (values, indices *Array, err error) {
if a.dt == Complex {
return nil, nil, errf("TopK: complex arrays have no ordering")
}
if dim < 0 || dim >= a.NDim() {
return nil, nil, errf("TopK: dimension %d out of range for shape %s", dim, shapeText(a.shape))
}
if a.shape[dim] == 0 {
return nil, nil, errf("TopK: dimension %d of shape %s is empty", dim, shapeText(a.shape))
}
if a.Len() == 0 {
// A dimension of size zero elsewhere leaves the per-line stride at
// zero and the line count undefined; the same empty answer
// argExtremeAxis gives.
return nil, nil, errf("TopK: an empty array has no result")
}
if k < 0 {
return nil, nil, errf("TopK: k must be non-negative, got %d", k)
}
if k > a.shape[dim] {
return nil, nil, errf("TopK: k=%d exceeds dimension %d size %d", k, dim, a.shape[dim])
}
outShape := a.Shape()
outShape[dim] = k
vals := &Array{shape: outShape, dt: a.dt}
idxs := &Array{shape: outShape, dt: Int}
valsTotal := 1
for _, d := range outShape {
valsTotal *= d
}
vals.alloc(valsTotal)
idxs.alloc(valsTotal)
stride := 1
for kk := a.NDim() - 1; kk > dim; kk-- {
stride *= a.shape[kk]
}
perLine := a.shape[dim] * stride
totalLines := a.Len() / perLine
n := a.shape[dim]
// The candidate positions are the same identity list for every line
// and the loop only reads them, so one shared snapshot serves all
// workers.
idxSnap := make([]int, n)
for i := range n {
idxSnap[i] = i
}
// The split counts element visits, not elements: each line is
// snapshotted once and then scanned k election rounds, so the fold
// touches every element about k+1 times and the spawn floor applies
// to that count.
splitCapped(totalLines, a.Len()*(k+1), reduceSplitFloor, func(ls, le int) {
// Per-worker scratch reused across lines: the value snapshots are
// rewritten in full each pass and the elected flags cleared, so
// no state leaks between lines.
valsSnap := make([]float64, n)
intsSnap := make([]int64, n)
used := make([]bool, n)
// Repeated argmax costs k full scans per line, which is quadratic
// when k approaches n. Wide requests switch to one sort of the
// line under the same total order the elections implement:
// value descending, ties by first occurrence, NaN last.
sortPath := k*8 > n
var pairs []topkPair
if sortPath {
pairs = make([]topkPair, n)
}
for line := ls; line < le; line++ {
baseFlat := line * perLine
for post := range stride {
// Snapshot the line from the raw payload: the dtype
// dispatch sits outside the walk, float32 widens to
// float64 exactly, and an integer-class line is
// snapshotted as int64 because the float64 detour would
// round neighbours above 2^53 together (see
// argExtreme). Every narrow widening into int64 is
// exact, so the int64 ranking carries each payload's own
// order whatever the width.
switch a.dt {
case Int:
src := a.ints
for i := range n {
intsSnap[i] = src[baseFlat+i*stride+post]
}
case Bool:
src := a.bools
for i := range n {
intsSnap[i] = boolToInt64(src[baseFlat+i*stride+post])
}
case Int8:
topkSnapNarrow(a.i8s, intsSnap, baseFlat, stride, post)
case Uint8:
topkSnapNarrow(a.u8s, intsSnap, baseFlat, stride, post)
case Int16:
topkSnapNarrow(a.i16s, intsSnap, baseFlat, stride, post)
case Uint16:
topkSnapNarrow(a.u16s, intsSnap, baseFlat, stride, post)
case Int32:
topkSnapNarrow(a.i32s, intsSnap, baseFlat, stride, post)
case Uint32:
topkSnapNarrow(a.u32s, intsSnap, baseFlat, stride, post)
case Float16:
src := a.halves
for i := range n {
valsSnap[i] = HalfToFloat64(src[baseFlat+i*stride+post])
}
case Float32:
src := a.floats32
for i := range n {
valsSnap[i] = float64(src[baseFlat+i*stride+post])
}
default:
src := a.floats
for i := range n {
valsSnap[i] = src[baseFlat+i*stride+post]
}
}
if sortPath {
topkSortLine(a.dt, valsSnap, intsSnap, idxSnap, pairs, k,
vals, idxs, line, stride, post)
continue
}
clear(used)
// Repeated argmax over the finite values: NaN candidates
// are skipped, so they can never be elected, and a line
// with fewer than k finite values fills its remaining
// slots with NaN.
for outK := range k {
bestIdx := -1
var bestVal float64
var bestInt int64
for i := range n {
if used[i] {
continue
}
if intClass(a.dt) {
// Integer-class candidates compare in exact
// int64, never float64, and an integer is
// never NaN: the first one always seeds.
if v := intsSnap[i]; bestIdx < 0 || v > bestInt {
bestIdx, bestInt = i, v
}
continue
}
if v := valsSnap[i]; v == v && (bestIdx < 0 || v > bestVal) {
bestIdx = i
bestVal = v
}
}
// The output line keeps the (k, suffix) row-major
// order: the reduced dimension carries stride, not
// the suffix.
outIdx := line*k*stride + outK*stride + post
if bestIdx < 0 {
// No finite value left; the float payloads fill
// with NaN. An integer-class line always seeds, so
// its zero fill never surfaces.
if intClass(a.dt) {
topkStore(vals, outIdx, 0)
} else {
switch a.dt {
case Float16:
vals.halves[outIdx] = halfNaN
case Float32:
vals.floats32[outIdx] = float32(math.NaN())
default:
vals.floats[outIdx] = math.NaN()
}
}
idxs.ints[outIdx] = 0
continue
}
if intClass(a.dt) {
// The original payload element widened exactly,
// never a rounded float64 image.
topkStore(vals, outIdx, intsSnap[bestIdx])
} else {
switch a.dt {
case Float16:
vals.halves[outIdx] = HalfFromFloat64(valsSnap[bestIdx])
case Float32:
vals.floats32[outIdx] = float32(valsSnap[bestIdx])
default:
vals.floats[outIdx] = valsSnap[bestIdx]
}
}
idxs.ints[outIdx] = int64(idxSnap[bestIdx])
used[bestIdx] = true
}
}
}
})
return vals, idxs, nil
}
// topkPair is one line element of the sort-based TopK path.
type topkPair struct {
val float64
ival int64
idx int
}
// topkSnapNarrow snapshots one narrow integer line into the int64
// snapshot the elections rank: every widening is exact, so the int64
// comparison carries the payload's own order whatever the width.
func topkSnapNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, intsSnap []int64, baseFlat, stride, post int) {
for i := range intsSnap {
intsSnap[i] = int64(src[baseFlat+i*stride+post])
}
}
// topkStore writes an integer-class TopK value back into the values
// array's own payload: the implicit-store cast the ladder carries.
func topkStore(vals *Array, outIdx int, v int64) {
switch vals.dt {
case Int:
vals.ints[outIdx] = v
case Bool:
vals.bools[outIdx] = v != 0
case Int8:
vals.i8s[outIdx] = int8(v)
case Uint8:
vals.u8s[outIdx] = uint8(v)
case Int16:
vals.i16s[outIdx] = int16(v)
case Uint16:
vals.u16s[outIdx] = uint16(v)
case Int32:
vals.i32s[outIdx] = int32(v)
default:
vals.u32s[outIdx] = uint32(v)
}
}
// topkSortLine elects the top-k of one line by sorting it under the
// elections' own total order: value descending, ties by first
// occurrence, NaN last. The output slots, the NaN fill and the index
// reporting match the repeated-argmax path exactly, so the two paths
// are interchangeable for any k.
func topkSortLine(dt Dtype, valsSnap []float64, intsSnap []int64, idxSnap []int, pairs []topkPair, k int,
vals, idxs *Array, line, stride, post int) {
n := len(pairs)
for i := range n {
pairs[i] = topkPair{val: valsSnap[i], ival: intsSnap[i], idx: idxSnap[i]}
}
slices.SortFunc(pairs, func(a, b topkPair) int {
if intClass(dt) {
// The exact int64 snapshot: every narrow widening preserves
// the payload's own order, so this is the native comparison.
if c := cmp.Compare(b.ival, a.ival); c != 0 {
return c
}
return cmp.Compare(a.idx, b.idx)
}
aNaN, bNaN := math.IsNaN(a.val), math.IsNaN(b.val)
switch {
case aNaN && bNaN:
return cmp.Compare(a.idx, b.idx)
case aNaN:
return 1
case bNaN:
return -1
}
if c := cmp.Compare(b.val, a.val); c != 0 {
return c
}
return cmp.Compare(a.idx, b.idx)
})
for outK := range k {
outIdx := line*k*stride + outK*stride + post
if outK < n && (intClass(dt) || !math.IsNaN(pairs[outK].val)) {
if intClass(dt) {
topkStore(vals, outIdx, pairs[outK].ival)
} else {
switch dt {
case Float16:
vals.halves[outIdx] = HalfFromFloat64(pairs[outK].val)
case Float32:
vals.floats32[outIdx] = float32(pairs[outK].val)
default:
vals.floats[outIdx] = pairs[outK].val
}
}
idxs.ints[outIdx] = int64(pairs[outK].idx)
continue
}
// Fewer than k finite values: the float payloads fill with NaN;
// an integer-class line always seeds, so its zero fill never
// surfaces.
if intClass(dt) {
topkStore(vals, outIdx, 0)
} else {
switch dt {
case Float16:
vals.halves[outIdx] = halfNaN
case Float32:
vals.floats32[outIdx] = float32(math.NaN())
default:
vals.floats[outIdx] = math.NaN()
}
}
idxs.ints[outIdx] = 0
}
}