1564 lines
46 KiB
Go
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
|
||
|
|
}
|
||
|
|
}
|