// Copyright (c) 2026 Petr Balvín (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 } }