1495 lines
45 KiB
Go
1495 lines
45 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package core
|
||
|
||
import (
|
||
"fmt"
|
||
"math"
|
||
"strconv"
|
||
|
||
"sourcedock.dev/petrbalvin/tensor/internal/engine"
|
||
)
|
||
|
||
// Scalar boxes a single numeric result whose dtype is known only at
|
||
// runtime: Sum of an int array is an int, of a float array a
|
||
// float, of a complex array a complex. Int, Float and Complex unpack it;
|
||
// IsFloat and IsComplex say which.
|
||
type Scalar struct {
|
||
isFloat bool
|
||
isComplex bool
|
||
i int64
|
||
f float64
|
||
c complex128
|
||
}
|
||
|
||
// Int returns the scalar as an int64. It converts a float or complex
|
||
// scalar by truncation, mirroring Go's conversions; use IsFloat and
|
||
// IsComplex when the distinction matters.
|
||
func (s Scalar) Int() int64 {
|
||
if s.isFloat {
|
||
return int64(s.f)
|
||
}
|
||
if s.isComplex {
|
||
return int64(real(s.c))
|
||
}
|
||
return s.i
|
||
}
|
||
|
||
// Float returns the scalar as a float64, converting int and complex
|
||
// scalars (a complex scalar contributes its real part).
|
||
func (s Scalar) Float() float64 {
|
||
if s.isFloat {
|
||
return s.f
|
||
}
|
||
if s.isComplex {
|
||
return real(s.c)
|
||
}
|
||
return float64(s.i)
|
||
}
|
||
|
||
// Complex returns the scalar as a complex128, converting real scalars.
|
||
func (s Scalar) Complex() complex128 {
|
||
if s.isComplex {
|
||
return s.c
|
||
}
|
||
if s.isFloat {
|
||
return complex(s.f, 0)
|
||
}
|
||
return complex(float64(s.i), 0)
|
||
}
|
||
|
||
// IsFloat reports whether the scalar came from a float computation.
|
||
func (s Scalar) IsFloat() bool { return s.isFloat }
|
||
|
||
// IsComplex reports whether the scalar came from a complex computation.
|
||
func (s Scalar) IsComplex() bool { return s.isComplex }
|
||
|
||
// String renders the scalar with its dtype, as in "int 7" or
|
||
// "complex (4-2i)". The complex form is the one %v prints, matching
|
||
// Array.String: a single sign between the parts, never "+-".
|
||
func (s Scalar) String() string {
|
||
switch {
|
||
case s.isComplex:
|
||
return "complex " + fmt.Sprintf("%v", s.c)
|
||
case s.isFloat:
|
||
return "float " + strconv.FormatFloat(s.f, 'g', -1, 64)
|
||
default:
|
||
return "int " + strconv.FormatInt(s.i, 10)
|
||
}
|
||
}
|
||
|
||
// Sum returns the sum of all elements. An int sum wraps on overflow like
|
||
// Go's int64 arithmetic; the narrow integer widths widen
|
||
// exactly and accumulate in int64, answered as an Int scalar, and a bool
|
||
// sum counts its true elements into the same Int scalar; float16 and
|
||
// float32 sums accumulate in float64 and answer a float scalar;
|
||
// an empty array sums to zero.
|
||
//
|
||
// The fold is partitioned by length alone, never by the worker count, so
|
||
// the same input answers the same scalar on any machine and under any
|
||
// worker setting: the range is cut into fixed chunks, each chunk
|
||
// accumulates into four interleaved partials, and the chunk results
|
||
// combine through a balanced pairwise tree over the chunk indices
|
||
// (treeSum). The tree is the shape the distributed reduction shares: a
|
||
// range cut at chunk boundaries into shards answers the same bits as
|
||
// this fold, because shard partials and array chunks are entries of one
|
||
// tree combined by one function. The integer sum is exact under any
|
||
// order. The floating-point fold rounds differently from the single
|
||
// chain it replaced, not always in the same direction: the interleaved
|
||
// partials shorten each dependency chain and the tree adds a logarithmic
|
||
// number of combination roundings, which measured within a small factor
|
||
// of the chain at every length tried and better than it once the chain
|
||
// is long enough for its own roundings to accumulate. Only the array's
|
||
// own elements take part: a rebased view's payload may run past its
|
||
// element count, and those invisible tail slots never contribute.
|
||
func Sum(a *Array) Scalar {
|
||
n := a.Len()
|
||
switch a.dt {
|
||
case Int:
|
||
return Scalar{i: intFoldSum(a.ints[:n])}
|
||
case Bool:
|
||
// The bool sum counts the true elements, answered as the Int
|
||
// scalar every integer-class reduction answers.
|
||
return Scalar{i: boolFoldCount(a.bools[:n])}
|
||
case Int8:
|
||
return Scalar{i: intFoldSumNarrow(a.i8s[:n])}
|
||
case Uint8:
|
||
return Scalar{i: intFoldSumNarrow(a.u8s[:n])}
|
||
case Int16:
|
||
return Scalar{i: intFoldSumNarrow(a.i16s[:n])}
|
||
case Uint16:
|
||
return Scalar{i: intFoldSumNarrow(a.u16s[:n])}
|
||
case Int32:
|
||
return Scalar{i: intFoldSumNarrow(a.i32s[:n])}
|
||
case Uint32:
|
||
return Scalar{i: intFoldSumNarrow(a.u32s[:n])}
|
||
case Float16:
|
||
// The half sum accumulates in float64 and answers a float
|
||
// scalar, exactly as the float32 sum does: every
|
||
// widening is exact, so the fold sees the same addends an
|
||
// accessor walk would hand over.
|
||
return Scalar{isFloat: true, f: floatFoldSumF16(a.halves[:n])}
|
||
case Float32:
|
||
return Scalar{isFloat: true, f: floatFoldSumF32(a.floats32[:n])}
|
||
case Float:
|
||
return Scalar{isFloat: true, f: floatFoldSum(a.floats[:n])}
|
||
default:
|
||
return Scalar{isComplex: true, c: complexFoldSum(a.complexes[:n])}
|
||
}
|
||
}
|
||
|
||
// foldChunk is the element count one partial fold carries. It is a
|
||
// constant of the arithmetic: the partition follows from the length
|
||
// alone, never from the worker count, so the same input answers the
|
||
// same scalar on any machine and under any worker setting.
|
||
const foldChunk = 1 << 16
|
||
|
||
// foldParts is the number of chunks a length is cut into, bounded so the
|
||
// partial table stays small for very long arrays.
|
||
func foldParts(n int) int {
|
||
if n <= foldChunk {
|
||
return 1
|
||
}
|
||
parts := min((n+foldChunk-1)/foldChunk, 1<<12)
|
||
return parts
|
||
}
|
||
|
||
// FoldParts reports how many blocks the canonical reduction partition
|
||
// cuts a range of n elements into. The partition is a function of the
|
||
// length alone, so it is the same on every machine and under any worker
|
||
// setting; the spmd package cuts distributed data on these boundaries,
|
||
// which is what makes a sharded reduction compose into the single-array
|
||
// fold's exact bits.
|
||
func FoldParts(n int) int { return foldParts(n) }
|
||
|
||
// FoldBoundary reports the index where block c of the canonical
|
||
// partition of n elements begins; block c spans [FoldBoundary(n, c),
|
||
// FoldBoundary(n, c+1)). Block 0 begins at 0 and FoldBoundary(n,
|
||
// FoldParts(n)) is n.
|
||
func FoldBoundary(n, c int) int { return c * n / foldParts(n) }
|
||
|
||
// TreeSum combines reduction partials through the balanced pairwise
|
||
// tree the folds combine their chunk partials with: the node over a
|
||
// range splits at its midpoint and adds the two halves' nodes, left
|
||
// before right. The value of any contiguous range of partials is one
|
||
// node of the tree, so partials gathered from a sharded cut combine
|
||
// into the single-array fold's exact bits; the spmd package feeds it
|
||
// the block values the shards folded.
|
||
func TreeSum[N int64 | float64 | complex128](vals []N) N { return treeSum(vals) }
|
||
|
||
// FoldRange is one canonical block's fold over a float64 payload: the
|
||
// four interleaved chains combined as a balanced pair. Shards fold
|
||
// their own blocks with it, so a block's value is the same bits
|
||
// wherever its elements live.
|
||
func FoldRange(src []float64) float64 { return foldRange(src) }
|
||
|
||
// FoldRangeF32 is FoldRange over a float32 payload; every widening is
|
||
// exact.
|
||
func FoldRangeF32(src []float32) float64 { return foldRangeF32(src) }
|
||
|
||
// FoldRangeF16 is FoldRange over a half payload's raw bit patterns;
|
||
// every widening is exact.
|
||
func FoldRangeF16(src []uint16) float64 { return foldRangeF16(src) }
|
||
|
||
// FoldRangeC128 is FoldRange over a complex payload.
|
||
func FoldRangeC128(src []complex128) complex128 { return foldRangeC128(src) }
|
||
|
||
// ExtremeRange is the serial extremum rule on one canonical block of a
|
||
// float64 payload: seed past the leading NaNs, keep the first strictly
|
||
// better value, and report whether the block holds a candidate at all.
|
||
func ExtremeRange(src []float64, greater bool) (float64, bool) { return floatExtreme(src, greater) }
|
||
|
||
// ExtremeRangeF32 is ExtremeRange over a float32 payload; every
|
||
// widening is exact.
|
||
func ExtremeRangeF32(src []float32, greater bool) (float64, bool) {
|
||
return float32Extreme(src, greater)
|
||
}
|
||
|
||
// ExtremeRangeF16 is ExtremeRange over a half payload's raw bit
|
||
// patterns; every widening is exact.
|
||
func ExtremeRangeF16(src []uint16, greater bool) (float64, bool) { return halfExtreme(src, greater) }
|
||
|
||
// CombineExtrema combines the blocks' extrema with the serial walk's
|
||
// strict comparison in index order: a tie keeps the earlier block's
|
||
// value, a block with no candidate contributes nothing, and an array
|
||
// with no candidate anywhere answers the last block's fallback, which
|
||
// is the whole array's last element. The extrema the spmd shards fold
|
||
// combine through it, so a sharded extremum is the single-array
|
||
// extremum bit for bit.
|
||
func CombineExtrema(vals []float64, oks []bool, greater bool) float64 {
|
||
return combineExtreme(vals, oks, greater)
|
||
}
|
||
|
||
// FloatScalar boxes a float64 as the scalar the float reductions
|
||
// answer.
|
||
func FloatScalar(f float64) Scalar { return Scalar{isFloat: true, f: f} }
|
||
|
||
// IntScalar boxes an int64 as the scalar the integer-class reductions
|
||
// answer.
|
||
func IntScalar(i int64) Scalar { return Scalar{i: i} }
|
||
|
||
// ComplexScalar boxes a complex128 as the scalar the complex
|
||
// reductions answer.
|
||
func ComplexScalar(c complex128) Scalar { return Scalar{isComplex: true, c: c} }
|
||
|
||
// TreeProd combines partial products through the balanced midpoint
|
||
// tree, multiplying the left half before the right: partials gathered
|
||
// from a sharded cut combine into the single-array product's exact
|
||
// bits. The float32 partials multiply natively in float32, the way the
|
||
// float32 product fold keeps.
|
||
func TreeProd[N int64 | float64 | float32](vals []N) N { return treeProd(vals) }
|
||
|
||
// TreeProdHalf combines half-precision partial products: the values
|
||
// carry exact half bits in float64 and every combine narrows through
|
||
// half, the per-step rounding the half product fold keeps.
|
||
func TreeProdHalf(vals []float64) float64 { return treeProdHalf(vals) }
|
||
|
||
// FoldProd is one canonical block's product over a float64 payload.
|
||
// Shards fold their own blocks with it, so a block's product is the
|
||
// same bits wherever its elements live.
|
||
func FoldProd(src []float64) float64 { return foldProdRange(src) }
|
||
|
||
// FoldProdF32 is FoldProd over a float32 payload, multiplied natively
|
||
// in float32.
|
||
func FoldProdF32(src []float32) float32 { return foldProdRangeF32(src) }
|
||
|
||
// FoldProdF16 is FoldProd over a half payload's raw bit patterns, with
|
||
// the per-step half rounding the line product keeps; the answer is an
|
||
// exact half value carried in float64.
|
||
func FoldProdF16(src []uint16) float64 { return foldProdRangeF16(src) }
|
||
|
||
// FoldProdI64 is FoldProd over an int64 payload; the wrapping product
|
||
// is exact under any grouping.
|
||
func FoldProdI64(src []int64) int64 { return foldProdRangeI64(src) }
|
||
|
||
// FoldNormPower is one canonical block's sum of |v|^p over a float64
|
||
// payload, with the per-element arithmetic the norm fold keeps.
|
||
func FoldNormPower(src []float64, p float64) float64 { return foldNormPowerRange(src, p) }
|
||
|
||
// FoldNormPowerF32 is FoldNormPower over a float32 payload; every
|
||
// widening is exact.
|
||
func FoldNormPowerF32(src []float32, p float64) float64 { return foldNormPowerRangeF32(src, p) }
|
||
|
||
// FoldNormPowerF16 is FoldNormPower over a half payload's raw bit
|
||
// patterns; every widening is exact.
|
||
func FoldNormPowerF16(src []uint16, p float64) float64 { return foldNormPowerRangeF16(src, p) }
|
||
|
||
// FoldNormPowerI64 is FoldNormPower over an int64 payload.
|
||
func FoldNormPowerI64(src []int64, p float64) float64 { return foldNormPowerRangeI64(src, p) }
|
||
|
||
// NormRoot closes a power sum into the norm: Sqrt for p = 2, the sum
|
||
// for p = 1, Pow of the sum for every other exponent.
|
||
func NormRoot(sum, p float64) float64 { return normRoot(sum, p) }
|
||
|
||
// FoldDot is one canonical block's dot product over float64 payloads
|
||
// of equal length.
|
||
func FoldDot(x, y []float64) float64 { return foldDotRange(x, y) }
|
||
|
||
// FoldDotF32 is FoldDot over float32 payloads; every product is exact
|
||
// in float64.
|
||
func FoldDotF32(x, y []float32) float64 { return foldDotRangeF32(x, y) }
|
||
|
||
// FoldDotF16 is FoldDot over half payloads' raw bit patterns; every
|
||
// product is exact in float64.
|
||
func FoldDotF16(x, y []uint16) float64 { return foldDotRangeF16(x, y) }
|
||
|
||
// FoldDotC128 is FoldDot over complex payloads.
|
||
func FoldDotC128(x, y []complex128) complex128 { return foldDotRangeC128(x, y) }
|
||
|
||
// FoldDotI64 is FoldDot over int64 payloads: the wrapping products and
|
||
// the wrapping sums are exact under any grouping.
|
||
func FoldDotI64(x, y []int64) int64 {
|
||
var s int64
|
||
for i := range x {
|
||
s += x[i] * y[i]
|
||
}
|
||
return s
|
||
}
|
||
|
||
// foldRange runs one chunk's fold over a float payload: four interleaved
|
||
// chains, so the adds of a long array overlap instead of queueing on one
|
||
// adder, combined as ((s0+s1)+(s2+s3)), the pairing a balanced tree
|
||
// gives. Each chain carries a quarter of the chunk, so the dependency
|
||
// chain is four times shorter and the combination costs three roundings;
|
||
// against the single chain it replaced the total error measured within a
|
||
// small factor either way, so this form is chosen for the throughput and
|
||
// not for a claimed accuracy win.
|
||
func foldRange(src []float64) float64 {
|
||
var s0, s1, s2, s3 float64
|
||
i := 0
|
||
for ; i+4 <= len(src); i += 4 {
|
||
s0 += src[i]
|
||
s1 += src[i+1]
|
||
s2 += src[i+2]
|
||
s3 += src[i+3]
|
||
}
|
||
for ; i < len(src); i++ {
|
||
s0 += src[i]
|
||
}
|
||
return (s0 + s1) + (s2 + s3)
|
||
}
|
||
|
||
// treeSum combines chunk partials with a balanced pairwise tree over
|
||
// their indices: the node over a range splits at its midpoint and adds
|
||
// the two halves' nodes, left before right. The shape depends on nothing
|
||
// but the partial count, and the value of any contiguous range of
|
||
// partials is one node of the tree. That property is what lets a
|
||
// reduction cut at chunk boundaries into shards reproduce this fold
|
||
// exactly, for any number of shards: the spmd package gathers the shard
|
||
// partials and combines them with this same function.
|
||
func treeSum[N int64 | float64 | complex128](vals []N) N {
|
||
return treeSumRange(vals, 0, len(vals))
|
||
}
|
||
|
||
// treeSumRange is one node of the treeSum tree over the partials [l, r).
|
||
func treeSumRange[N int64 | float64 | complex128](vals []N, l, r int) N {
|
||
if r-l == 1 {
|
||
return vals[l]
|
||
}
|
||
m := l + (r-l)/2
|
||
return treeSumRange(vals, l, m) + treeSumRange(vals, m, r)
|
||
}
|
||
|
||
// treeProd combines partial products through the balanced midpoint
|
||
// tree treeSum uses, multiplying the left half before the right. The
|
||
// shape depends on nothing but the partial count, so partial products
|
||
// gathered from a sharded cut combine into the single-array product's
|
||
// exact bits; the spmd package feeds it the block values the shards
|
||
// folded.
|
||
func treeProd[N int64 | float64 | float32](vals []N) N {
|
||
return treeProdRange(vals, 0, len(vals))
|
||
}
|
||
|
||
func treeProdRange[N int64 | float64 | float32](vals []N, l, r int) N {
|
||
if r-l == 1 {
|
||
return vals[l]
|
||
}
|
||
m := l + (r-l)/2
|
||
return treeProdRange(vals, l, m) * treeProdRange(vals, m, r)
|
||
}
|
||
|
||
// treeProdHalf is treeProd for a half-precision product: the partials
|
||
// carry exact half values and every combine narrows through half, the
|
||
// per-step rounding the half product fold keeps.
|
||
func treeProdHalf(vals []float64) float64 {
|
||
return treeProdHalfRange(vals, 0, len(vals))
|
||
}
|
||
|
||
func treeProdHalfRange(vals []float64, l, r int) float64 {
|
||
if r-l == 1 {
|
||
return vals[l]
|
||
}
|
||
m := l + (r-l)/2
|
||
return halfRound(treeProdHalfRange(vals, l, m) * treeProdHalfRange(vals, m, r))
|
||
}
|
||
|
||
// halfRound narrows through half and back, the per-step rounding the
|
||
// half precision folds keep.
|
||
func halfRound(f float64) float64 {
|
||
return HalfToFloat64(HalfFromFloat64(f))
|
||
}
|
||
|
||
// foldProdRangeI64 is one canonical block's product over an int64
|
||
// payload: wrapping multiplication is associative, so the grouping
|
||
// cannot change the value.
|
||
func foldProdRangeI64(src []int64) int64 {
|
||
m := int64(1)
|
||
for _, v := range src {
|
||
m *= v
|
||
}
|
||
return m
|
||
}
|
||
|
||
// foldProdRange is one canonical block's product over a float64
|
||
// payload: a single chain from one, the association the product fold
|
||
// keeps inside a block.
|
||
func foldProdRange(src []float64) float64 {
|
||
m := 1.0
|
||
for _, v := range src {
|
||
m *= v
|
||
}
|
||
return m
|
||
}
|
||
|
||
// foldProdRangeF32 is one canonical block's product over a float32
|
||
// payload, multiplied natively in float32.
|
||
func foldProdRangeF32(src []float32) float32 {
|
||
m := float32(1)
|
||
for _, v := range src {
|
||
m *= v
|
||
}
|
||
return m
|
||
}
|
||
|
||
// foldProdRangeF16 is one canonical block's product over a half
|
||
// payload: the running product narrows through half before every
|
||
// combine, exactly the per-step rounding the line fold keeps, and the
|
||
// answer is an exact half value carried in float64.
|
||
func foldProdRangeF16(src []uint16) float64 {
|
||
m := 1.0
|
||
for _, v := range src {
|
||
m = halfRound(m * HalfToFloat64(v))
|
||
}
|
||
return m
|
||
}
|
||
|
||
// foldNormPowerRange is one canonical block's sum of |v|^p over a
|
||
// float64 payload, with the same per-element arithmetic the norm fold
|
||
// keeps: the bare absolute for p = 1, the squared absolute for p = 2,
|
||
// and Pow of the absolute for every other exponent.
|
||
func foldNormPowerRange(src []float64, p float64) float64 {
|
||
switch {
|
||
case p == 1:
|
||
var acc float64
|
||
for _, v := range src {
|
||
acc += math.Abs(v)
|
||
}
|
||
return acc
|
||
case p == 2:
|
||
var acc float64
|
||
for _, v := range src {
|
||
w := math.Abs(v)
|
||
acc += w * w
|
||
}
|
||
return acc
|
||
default:
|
||
var acc float64
|
||
for _, v := range src {
|
||
acc += math.Pow(math.Abs(v), p)
|
||
}
|
||
return acc
|
||
}
|
||
}
|
||
|
||
// foldNormPowerRangeF32 is foldNormPowerRange over a float32 payload;
|
||
// every widening is exact.
|
||
func foldNormPowerRangeF32(src []float32, p float64) float64 {
|
||
switch {
|
||
case p == 1:
|
||
var acc float64
|
||
for _, v := range src {
|
||
acc += math.Abs(float64(v))
|
||
}
|
||
return acc
|
||
case p == 2:
|
||
var acc float64
|
||
for _, v := range src {
|
||
w := math.Abs(float64(v))
|
||
acc += w * w
|
||
}
|
||
return acc
|
||
default:
|
||
var acc float64
|
||
for _, v := range src {
|
||
acc += math.Pow(math.Abs(float64(v)), p)
|
||
}
|
||
return acc
|
||
}
|
||
}
|
||
|
||
// foldNormPowerRangeF16 is foldNormPowerRange over a half payload's
|
||
// raw bit patterns; every widening is exact.
|
||
func foldNormPowerRangeF16(src []uint16, p float64) float64 {
|
||
switch {
|
||
case p == 1:
|
||
var acc float64
|
||
for _, v := range src {
|
||
acc += math.Abs(HalfToFloat64(v))
|
||
}
|
||
return acc
|
||
case p == 2:
|
||
var acc float64
|
||
for _, v := range src {
|
||
w := math.Abs(HalfToFloat64(v))
|
||
acc += w * w
|
||
}
|
||
return acc
|
||
default:
|
||
var acc float64
|
||
for _, v := range src {
|
||
acc += math.Pow(math.Abs(HalfToFloat64(v)), p)
|
||
}
|
||
return acc
|
||
}
|
||
}
|
||
|
||
// foldNormPowerRangeI64 is foldNormPowerRange over an int64 payload.
|
||
func foldNormPowerRangeI64(src []int64, p float64) float64 {
|
||
switch {
|
||
case p == 1:
|
||
var acc float64
|
||
for _, v := range src {
|
||
acc += math.Abs(float64(v))
|
||
}
|
||
return acc
|
||
case p == 2:
|
||
var acc float64
|
||
for _, v := range src {
|
||
w := math.Abs(float64(v))
|
||
acc += w * w
|
||
}
|
||
return acc
|
||
default:
|
||
var acc float64
|
||
for _, v := range src {
|
||
acc += math.Pow(math.Abs(float64(v)), p)
|
||
}
|
||
return acc
|
||
}
|
||
}
|
||
|
||
// normRoot closes a power sum into the norm: Sqrt for p = 2, the sum
|
||
// itself for p = 1, and Pow of the sum for every other exponent, the
|
||
// closing the norm fold keeps.
|
||
func normRoot(sum, p float64) float64 {
|
||
switch {
|
||
case p == 2:
|
||
return math.Sqrt(sum)
|
||
case p == 1:
|
||
return sum
|
||
default:
|
||
return math.Pow(sum, 1/p)
|
||
}
|
||
}
|
||
|
||
// foldDotRange is one canonical block's dot product over float64
|
||
// payloads of equal length: four interleaved product chains, combined
|
||
// as a balanced pair.
|
||
func foldDotRange(x, y []float64) float64 {
|
||
var s0, s1, s2, s3 float64
|
||
i := 0
|
||
for ; i+4 <= len(x); i += 4 {
|
||
s0 += x[i] * y[i]
|
||
s1 += x[i+1] * y[i+1]
|
||
s2 += x[i+2] * y[i+2]
|
||
s3 += x[i+3] * y[i+3]
|
||
}
|
||
for ; i < len(x); i++ {
|
||
s0 += x[i] * y[i]
|
||
}
|
||
return (s0 + s1) + (s2 + s3)
|
||
}
|
||
|
||
// foldDotRangeF32 is foldDotRange over float32 payloads; every product
|
||
// is exact in float64.
|
||
func foldDotRangeF32(x, y []float32) float64 {
|
||
var s0, s1, s2, s3 float64
|
||
i := 0
|
||
for ; i+4 <= len(x); i += 4 {
|
||
s0 += float64(x[i]) * float64(y[i])
|
||
s1 += float64(x[i+1]) * float64(y[i+1])
|
||
s2 += float64(x[i+2]) * float64(y[i+2])
|
||
s3 += float64(x[i+3]) * float64(y[i+3])
|
||
}
|
||
for ; i < len(x); i++ {
|
||
s0 += float64(x[i]) * float64(y[i])
|
||
}
|
||
return (s0 + s1) + (s2 + s3)
|
||
}
|
||
|
||
// foldDotRangeF16 is foldDotRange over half payloads' raw bit
|
||
// patterns; every product is exact in float64.
|
||
func foldDotRangeF16(x, y []uint16) float64 {
|
||
var s0, s1, s2, s3 float64
|
||
i := 0
|
||
for ; i+4 <= len(x); i += 4 {
|
||
s0 += HalfToFloat64(x[i]) * HalfToFloat64(y[i])
|
||
s1 += HalfToFloat64(x[i+1]) * HalfToFloat64(y[i+1])
|
||
s2 += HalfToFloat64(x[i+2]) * HalfToFloat64(y[i+2])
|
||
s3 += HalfToFloat64(x[i+3]) * HalfToFloat64(y[i+3])
|
||
}
|
||
for ; i < len(x); i++ {
|
||
s0 += HalfToFloat64(x[i]) * HalfToFloat64(y[i])
|
||
}
|
||
return (s0 + s1) + (s2 + s3)
|
||
}
|
||
|
||
// foldDotRangeC128 is foldDotRange over complex payloads.
|
||
func foldDotRangeC128(x, y []complex128) complex128 {
|
||
var s0, s1, s2, s3 complex128
|
||
i := 0
|
||
for ; i+4 <= len(x); i += 4 {
|
||
s0 += x[i] * y[i]
|
||
s1 += x[i+1] * y[i+1]
|
||
s2 += x[i+2] * y[i+2]
|
||
s3 += x[i+3] * y[i+3]
|
||
}
|
||
for ; i < len(x); i++ {
|
||
s0 += x[i] * y[i]
|
||
}
|
||
return (s0 + s1) + (s2 + s3)
|
||
}
|
||
|
||
// floatFoldSum sums a float64 payload over the fixed partition.
|
||
func floatFoldSum(src []float64) float64 {
|
||
parts := foldParts(len(src))
|
||
if parts == 1 {
|
||
return foldRange(src)
|
||
}
|
||
partials := make([]float64, parts)
|
||
engine.Parallel(parts, func(cs, ce int) {
|
||
for c := cs; c < ce; c++ {
|
||
partials[c] = foldRange(src[c*len(src)/parts : (c+1)*len(src)/parts])
|
||
}
|
||
})
|
||
return treeSum(partials)
|
||
}
|
||
|
||
// foldRangeF32 is one block's fold over a float32 payload: four
|
||
// interleaved float64 chains, every widening exact.
|
||
func foldRangeF32(part []float32) float64 {
|
||
var s0, s1, s2, s3 float64
|
||
i := 0
|
||
for ; i+4 <= len(part); i += 4 {
|
||
s0 += float64(part[i])
|
||
s1 += float64(part[i+1])
|
||
s2 += float64(part[i+2])
|
||
s3 += float64(part[i+3])
|
||
}
|
||
for ; i < len(part); i++ {
|
||
s0 += float64(part[i])
|
||
}
|
||
return (s0 + s1) + (s2 + s3)
|
||
}
|
||
|
||
// foldRangeF16 is one block's fold over a half payload's raw bit
|
||
// patterns, widening each element exactly as it is read.
|
||
func foldRangeF16(part []uint16) float64 {
|
||
var s0, s1, s2, s3 float64
|
||
i := 0
|
||
for ; i+4 <= len(part); i += 4 {
|
||
s0 += HalfToFloat64(part[i])
|
||
s1 += HalfToFloat64(part[i+1])
|
||
s2 += HalfToFloat64(part[i+2])
|
||
s3 += HalfToFloat64(part[i+3])
|
||
}
|
||
for ; i < len(part); i++ {
|
||
s0 += HalfToFloat64(part[i])
|
||
}
|
||
return (s0 + s1) + (s2 + s3)
|
||
}
|
||
|
||
// floatFoldSumF32 sums a float32 payload: every widening to float64 is
|
||
// exact, so the fold sees the values an accessor walk would hand over.
|
||
func floatFoldSumF32(src []float32) float64 {
|
||
parts := foldParts(len(src))
|
||
return sumOverParts(len(src), parts, func(c int) float64 {
|
||
return foldRangeF32(src[c*len(src)/parts : (c+1)*len(src)/parts])
|
||
})
|
||
}
|
||
|
||
// floatFoldSumF16 sums a half-precision payload, widening each element
|
||
// exactly as it is read.
|
||
func floatFoldSumF16(src []uint16) float64 {
|
||
parts := foldParts(len(src))
|
||
return sumOverParts(len(src), parts, func(c int) float64 {
|
||
return foldRangeF16(src[c*len(src)/parts : (c+1)*len(src)/parts])
|
||
})
|
||
}
|
||
|
||
// sumOverParts runs the per-chunk fold and combines the results through
|
||
// the balanced partial tree: the only order-dependent step, and its shape
|
||
// depends on nothing but the chunk count.
|
||
func sumOverParts(n, parts int, fold func(c int) float64) float64 {
|
||
if parts == 1 {
|
||
return fold(0)
|
||
}
|
||
partials := make([]float64, parts)
|
||
engine.Parallel(parts, func(cs, ce int) {
|
||
for c := cs; c < ce; c++ {
|
||
partials[c] = fold(c)
|
||
}
|
||
})
|
||
return treeSum(partials)
|
||
}
|
||
|
||
// intFoldSum sums an int64 payload: wrapping addition is associative, so
|
||
// the partition cannot change the value and a plain chunked scan needs no
|
||
// interleaved partials.
|
||
func intFoldSum(src []int64) int64 {
|
||
parts := foldParts(len(src))
|
||
if parts == 1 {
|
||
var s int64
|
||
for _, v := range src {
|
||
s += v
|
||
}
|
||
return s
|
||
}
|
||
partials := make([]int64, parts)
|
||
engine.Parallel(parts, func(cs, ce int) {
|
||
for c := cs; c < ce; c++ {
|
||
var s int64
|
||
for _, v := range src[c*len(src)/parts : (c+1)*len(src)/parts] {
|
||
s += v
|
||
}
|
||
partials[c] = s
|
||
}
|
||
})
|
||
var total int64
|
||
for _, v := range partials {
|
||
total += v
|
||
}
|
||
return total
|
||
}
|
||
|
||
// intFoldSumNarrow sums a narrow integer payload into int64 through the
|
||
// fixed partition intFoldSum uses; every widening is exact, so the sum
|
||
// stays machine-independent and the accumulation order changes nothing.
|
||
func intFoldSumNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T) int64 {
|
||
parts := foldParts(len(src))
|
||
if parts == 1 {
|
||
var s int64
|
||
for _, v := range src {
|
||
s += int64(v)
|
||
}
|
||
return s
|
||
}
|
||
partials := make([]int64, parts)
|
||
engine.Parallel(parts, func(cs, ce int) {
|
||
for c := cs; c < ce; c++ {
|
||
var s int64
|
||
for _, v := range src[c*len(src)/parts : (c+1)*len(src)/parts] {
|
||
s += int64(v)
|
||
}
|
||
partials[c] = s
|
||
}
|
||
})
|
||
var total int64
|
||
for _, v := range partials {
|
||
total += v
|
||
}
|
||
return total
|
||
}
|
||
|
||
// boolFoldCount counts the true elements of a bool payload through the
|
||
// same fixed partition intFoldSum uses.
|
||
func boolFoldCount(src []bool) int64 {
|
||
parts := foldParts(len(src))
|
||
if parts == 1 {
|
||
var s int64
|
||
for _, v := range src {
|
||
if v {
|
||
s++
|
||
}
|
||
}
|
||
return s
|
||
}
|
||
partials := make([]int64, parts)
|
||
engine.Parallel(parts, func(cs, ce int) {
|
||
for c := cs; c < ce; c++ {
|
||
var s int64
|
||
for _, v := range src[c*len(src)/parts : (c+1)*len(src)/parts] {
|
||
if v {
|
||
s++
|
||
}
|
||
}
|
||
partials[c] = s
|
||
}
|
||
})
|
||
var total int64
|
||
for _, v := range partials {
|
||
total += v
|
||
}
|
||
return total
|
||
}
|
||
|
||
// foldRangeC128 is one block's fold over a complex payload: four
|
||
// interleaved chains, combined as a balanced pair.
|
||
func foldRangeC128(part []complex128) complex128 {
|
||
var s0, s1, s2, s3 complex128
|
||
i := 0
|
||
for ; i+4 <= len(part); i += 4 {
|
||
s0 += part[i]
|
||
s1 += part[i+1]
|
||
s2 += part[i+2]
|
||
s3 += part[i+3]
|
||
}
|
||
for ; i < len(part); i++ {
|
||
s0 += part[i]
|
||
}
|
||
return (s0 + s1) + (s2 + s3)
|
||
}
|
||
|
||
// complexFoldSum sums a complex payload the way the float fold sums a
|
||
// real one: fixed chunks, four interleaved partials, partials combined
|
||
// through the balanced partial tree.
|
||
func complexFoldSum(src []complex128) complex128 {
|
||
parts := foldParts(len(src))
|
||
if parts == 1 {
|
||
return foldRangeC128(src)
|
||
}
|
||
partials := make([]complex128, parts)
|
||
engine.Parallel(parts, func(cs, ce int) {
|
||
for c := cs; c < ce; c++ {
|
||
partials[c] = foldRangeC128(src[c*len(src)/parts : (c+1)*len(src)/parts])
|
||
}
|
||
})
|
||
return treeSum(partials)
|
||
}
|
||
|
||
// Min returns the smallest element; an empty array, or a complex array
|
||
// (no ordering), is an error.
|
||
func Min(a *Array) (Scalar, error) {
|
||
return a.reduceOrder("Min", false)
|
||
}
|
||
|
||
// Max returns the largest element; an empty array, or a complex array
|
||
// (no ordering), is an error.
|
||
func Max(a *Array) (Scalar, error) {
|
||
return a.reduceOrder("Max", true)
|
||
}
|
||
|
||
// Mean returns the arithmetic mean of all elements as a float64,
|
||
// computed in float64 even for int arrays (true division); an empty or
|
||
// complex array is an error.
|
||
//
|
||
// The sum runs through Sum, so every ordinary input rounds exactly as
|
||
// the standalone Sum does. The one exception is the input whose largest
|
||
// magnitude could carry the running total past float64's range: there
|
||
// the plain sum overflows to an Inf the final division cannot undo
|
||
// (Mean of two MaxFloat64s answered +Inf). Such an input takes a scaled
|
||
// two-pass form instead: the elements are summed divided by the largest
|
||
// magnitude and the result rescaled, which stays finite. The scale is
|
||
// chosen by the pre-pass below, so any input below the overflow
|
||
// threshold, and every int array whose whole range sits far below
|
||
// float64's ceiling, keeps the plain path and its digest.
|
||
func Mean(a *Array) (float64, error) {
|
||
if a.dt == Complex {
|
||
return 0, errf("Mean: complex arrays have no float mean")
|
||
}
|
||
if a.Len() == 0 {
|
||
return 0, errf("Mean: an empty array has no mean")
|
||
}
|
||
n := a.Len()
|
||
fn := float64(n)
|
||
// Every widening in this pre-pass is exact, and NaN fails the
|
||
// comparison, so a NaN payload never triggers the scaled path and
|
||
// reaches the caller through Sum's fold as before. The walk reads the
|
||
// payload directly, which is the value the accessor would hand over.
|
||
maxAbs := a.maxAbs(n)
|
||
if maxAbs > math.MaxFloat64/fn && !math.IsInf(maxAbs, 0) {
|
||
var acc float64
|
||
for i := range n {
|
||
acc += a.floatAt(i) / maxAbs
|
||
}
|
||
// The rescale divides by n first: multiplying acc by maxAbs
|
||
// before the division could overflow again, which is the very
|
||
// failure this path exists to avoid.
|
||
return acc / fn * maxAbs, nil
|
||
}
|
||
return Sum(a).Float() / fn, nil
|
||
}
|
||
|
||
// maxAbs returns the largest magnitude among the first n elements, the
|
||
// seed pre-pass the scaled mean path needs. The dtype dispatch sits
|
||
// outside the walk and float16 and float32 widen exactly, so every
|
||
// value is the one an accessor read would hand over, at two
|
||
// instructions per element instead of a switch on the dtype.
|
||
//
|
||
// A magnitude maximum is associative and commutative, and it has no sign
|
||
// to lose: |+0| and |−0| are the same zero, so unlike the signed extrema
|
||
// the answer cannot depend on the visit order. The walk therefore takes
|
||
// the same fixed partition the sums use and combines the partial maxima
|
||
// with max, which makes it both parallel and deterministic.
|
||
func (a *Array) maxAbs(n int) float64 {
|
||
if a.strides != nil {
|
||
parts := foldParts(n)
|
||
if parts == 1 {
|
||
return a.maxAbsSerial(n)
|
||
}
|
||
partials := make([]float64, parts)
|
||
engine.Parallel(parts, func(cs, ce int) {
|
||
for c := cs; c < ce; c++ {
|
||
lo, hi := c*n/parts, (c+1)*n/parts
|
||
m := 0.0
|
||
for i := lo; i < hi; i++ {
|
||
if w := math.Abs(a.floatAt(i)); w > m {
|
||
m = w
|
||
}
|
||
}
|
||
partials[c] = m
|
||
}
|
||
})
|
||
m := 0.0
|
||
for _, v := range partials {
|
||
m = math.Max(m, v)
|
||
}
|
||
return m
|
||
}
|
||
switch a.dt {
|
||
case Int:
|
||
return foldMaxAbsParts(n, func(lo, hi int) float64 {
|
||
m := 0.0
|
||
for _, v := range a.ints[lo:hi] {
|
||
if w := math.Abs(float64(v)); w > m {
|
||
m = w
|
||
}
|
||
}
|
||
return m
|
||
})
|
||
case Float16:
|
||
return foldMaxAbsParts(n, func(lo, hi int) float64 {
|
||
m := 0.0
|
||
for _, v := range a.halves[lo:hi] {
|
||
if w := math.Abs(HalfToFloat64(v)); w > m {
|
||
m = w
|
||
}
|
||
}
|
||
return m
|
||
})
|
||
case Float32:
|
||
return foldMaxAbsParts(n, func(lo, hi int) float64 {
|
||
m := 0.0
|
||
for _, v := range a.floats32[lo:hi] {
|
||
if w := math.Abs(float64(v)); w > m {
|
||
m = w
|
||
}
|
||
}
|
||
return m
|
||
})
|
||
case Float:
|
||
return foldMaxAbsParts(n, func(lo, hi int) float64 {
|
||
m := 0.0
|
||
for _, v := range a.floats[lo:hi] {
|
||
if w := math.Abs(v); w > m {
|
||
m = w
|
||
}
|
||
}
|
||
return m
|
||
})
|
||
case Bool:
|
||
return foldMaxAbsParts(n, func(lo, hi int) float64 {
|
||
for _, v := range a.bools[lo:hi] {
|
||
if v {
|
||
return 1
|
||
}
|
||
}
|
||
return 0
|
||
})
|
||
case Int8:
|
||
return maxAbsNarrow(a.i8s, n)
|
||
case Uint8:
|
||
return maxAbsNarrow(a.u8s, n)
|
||
case Int16:
|
||
return maxAbsNarrow(a.i16s, n)
|
||
case Uint16:
|
||
return maxAbsNarrow(a.u16s, n)
|
||
case Int32:
|
||
return maxAbsNarrow(a.i32s, n)
|
||
case Uint32:
|
||
return maxAbsNarrow(a.u32s, n)
|
||
}
|
||
// Complex never reaches here: Mean rejects it before the guard.
|
||
return 0
|
||
}
|
||
|
||
// maxAbsNarrow is maxAbs's chunk walk for a narrow integer payload:
|
||
// every widening is exact, so the magnitude sees the accessor value.
|
||
func maxAbsNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, n int) float64 {
|
||
return foldMaxAbsParts(n, func(lo, hi int) float64 {
|
||
m := 0.0
|
||
for _, v := range src[lo:hi] {
|
||
if w := math.Abs(float64(v)); w > m {
|
||
m = w
|
||
}
|
||
}
|
||
return m
|
||
})
|
||
}
|
||
|
||
// foldMaxAbsParts runs the magnitude maximum over the fixed partition of
|
||
// n elements. The chunk closure walks a payload slice, never a function
|
||
// per element, and the partial maxima combine with max, so the partition
|
||
// changes nothing.
|
||
func foldMaxAbsParts(n int, chunk func(lo, hi int) float64) float64 {
|
||
parts := foldParts(n)
|
||
if parts == 1 {
|
||
return chunk(0, n)
|
||
}
|
||
partials := make([]float64, parts)
|
||
engine.Parallel(parts, func(cs, ce int) {
|
||
for c := cs; c < ce; c++ {
|
||
partials[c] = chunk(c*n/parts, (c+1)*n/parts)
|
||
}
|
||
})
|
||
m := 0.0
|
||
for _, v := range partials {
|
||
m = math.Max(m, v)
|
||
}
|
||
return m
|
||
}
|
||
|
||
// maxAbsSerial is the strided fallback for a length below the partition
|
||
// floor.
|
||
func (a *Array) maxAbsSerial(n int) float64 {
|
||
var m float64
|
||
for i := range n {
|
||
if v := math.Abs(a.floatAt(i)); v > m {
|
||
m = v
|
||
}
|
||
}
|
||
return m
|
||
}
|
||
|
||
// Dot returns the dot product of two 1-D arrays of equal length.
|
||
// Integer-class pairs produce an int scalar accumulated in int64 (the
|
||
// products wrap like the int64 kernel's); bool pairs are refused: bool
|
||
// carries no arithmetic. Float32 operands accumulate in float64 and
|
||
// produce a float scalar; any float64 operand promotes the
|
||
// result to float, any complex operand to complex.
|
||
func Dot(a, b *Array) (Scalar, error) {
|
||
if a.NDim() != 1 || b.NDim() != 1 {
|
||
return Scalar{}, errf("Dot: needs 1-D arrays, got shapes %s and %s",
|
||
shapeText(a.shape), shapeText(b.shape))
|
||
}
|
||
if a.Len() != b.Len() {
|
||
return Scalar{}, errf("Dot: length mismatch %d vs %d", a.Len(), b.Len())
|
||
}
|
||
n := a.Len()
|
||
dt := promote(a.dt, b.dt)
|
||
if dt == Bool {
|
||
return Scalar{}, errf("Dot: bool arrays have no arithmetic")
|
||
}
|
||
switch {
|
||
case dt == Int && a.dt == Int && b.dt == Int:
|
||
return Scalar{i: intFoldDot(a.ints[:n], b.ints[:n])}, nil
|
||
case intClass(dt):
|
||
// Every integer-class pair widens exactly into int64; the
|
||
// products wrap exactly as the int64 kernel's do, and wrapping
|
||
// addition is associative, so the fixed partition changes
|
||
// nothing.
|
||
if a.dt == b.dt && a.dt != Bool && a.isContiguous() && b.isContiguous() {
|
||
switch a.dt {
|
||
case Int:
|
||
return Scalar{i: intFoldDot(a.ints[:n], b.ints[:n])}, nil
|
||
case Int8:
|
||
return Scalar{i: narrowFoldDot(a.i8s[:n], b.i8s[:n])}, nil
|
||
case Uint8:
|
||
return Scalar{i: narrowFoldDot(a.u8s[:n], b.u8s[:n])}, nil
|
||
case Int16:
|
||
return Scalar{i: narrowFoldDot(a.i16s[:n], b.i16s[:n])}, nil
|
||
case Uint16:
|
||
return Scalar{i: narrowFoldDot(a.u16s[:n], b.u16s[:n])}, nil
|
||
case Int32:
|
||
return Scalar{i: narrowFoldDot(a.i32s[:n], b.i32s[:n])}, nil
|
||
case Uint32:
|
||
return Scalar{i: narrowFoldDot(a.u32s[:n], b.u32s[:n])}, nil
|
||
}
|
||
}
|
||
parts := foldParts(n)
|
||
return Scalar{i: foldPartsOver(n, parts, func(c int) int64 {
|
||
lo, hi := c*n/parts, (c+1)*n/parts
|
||
var s int64
|
||
for i := lo; i < hi; i++ {
|
||
s += a.intAt(i) * b.intAt(i)
|
||
}
|
||
return s
|
||
})}, nil
|
||
}
|
||
switch dt {
|
||
case Float16:
|
||
var s float64
|
||
if a.dt == Float16 && b.dt == Float16 {
|
||
// Both payloads hold half bit patterns at their flat index,
|
||
// so the kernel streams the raw slices; every product is
|
||
// exact in float64 either way.
|
||
as, bs := a.halves[:n], b.halves[:n]
|
||
s = foldDotF16(as, bs)
|
||
} else {
|
||
for i := range n {
|
||
s += a.floatAt(i) * b.floatAt(i)
|
||
}
|
||
}
|
||
return Scalar{isFloat: true, f: s}, nil
|
||
case Float32:
|
||
var s float64
|
||
if a.dt == Float32 && b.dt == Float32 {
|
||
// Both payloads hold float32 at their flat index, so the
|
||
// kernel streams the raw slices; every product is exact in
|
||
// float64 either way, so the values match accessor reads.
|
||
as, bs := a.floats32[:n], b.floats32[:n]
|
||
s = foldDotF32(as, bs)
|
||
} else {
|
||
for i := range n {
|
||
s += a.floatAt(i) * b.floatAt(i)
|
||
}
|
||
}
|
||
return Scalar{isFloat: true, f: s}, nil
|
||
case Float:
|
||
var s float64
|
||
if a.dt == Float && b.dt == Float {
|
||
as, bs := a.floats[:n], b.floats[:n]
|
||
s = foldDotF64(as, bs)
|
||
} else {
|
||
for i := range n {
|
||
s += a.floatAt(i) * b.floatAt(i)
|
||
}
|
||
}
|
||
return Scalar{isFloat: true, f: s}, nil
|
||
default:
|
||
var s complex128
|
||
if a.dt == Complex && b.dt == Complex {
|
||
as, bs := a.complexes[:n], b.complexes[:n]
|
||
s = foldDotC128(as, bs)
|
||
} else {
|
||
for i := range n {
|
||
s += a.complexAt(i) * b.complexAt(i)
|
||
}
|
||
}
|
||
return Scalar{isComplex: true, c: s}, nil
|
||
}
|
||
}
|
||
|
||
// reduceOrder walks the array keeping the smallest (greater=false) or
|
||
// largest (greater=true) element per dtype. An empty array is an error,
|
||
// and so is a complex array: it has no ordering. The comparisons are
|
||
// inlined per dtype: a closure per element would dominate the walk.
|
||
// reduceOrder walks the array keeping the smallest (greater=false) or
|
||
// largest (greater=true) element per dtype. An empty array is an error,
|
||
// and so is a complex array: it has no ordering.
|
||
//
|
||
// The walk is partitioned by length alone, never by the worker count, so
|
||
// the same input answers the same element on any machine: the range is
|
||
// cut into fixed chunks, each chunk applies the serial rule to its own
|
||
// range, and the chunks combine in index order with the same strict
|
||
// comparison. A tie therefore keeps the earlier chunk's element and, in a
|
||
// chunk, the earlier element's, which is what the serial walk kept; the
|
||
// rules for a NaN (never a candidate) and an all-NaN array (the last
|
||
// element, where the seed walk ends) are preserved exactly, zeros
|
||
// included.
|
||
func (a *Array) reduceOrder(name string, greater bool) (Scalar, error) {
|
||
if a.dt == Complex {
|
||
return Scalar{}, errf("%s: complex arrays have no ordering", name)
|
||
}
|
||
if a.Len() == 0 {
|
||
return Scalar{}, errf("%s: an empty array has no %s", name, name)
|
||
}
|
||
n := a.Len()
|
||
switch a.dt {
|
||
case Int:
|
||
return Scalar{i: foldExtreme(n, greater, func(lo, hi int) (int64, bool) {
|
||
is := a.ints[lo:hi]
|
||
best := is[0]
|
||
if greater {
|
||
for _, v := range is[1:] {
|
||
if v > best {
|
||
best = v
|
||
}
|
||
}
|
||
} else {
|
||
for _, v := range is[1:] {
|
||
if v < best {
|
||
best = v
|
||
}
|
||
}
|
||
}
|
||
return best, true
|
||
})}, nil
|
||
case Bool:
|
||
return Scalar{i: foldExtremeBool(a.bools, n, greater)}, nil
|
||
case Int8:
|
||
return Scalar{i: foldExtremeNarrow(a.i8s, n, greater)}, nil
|
||
case Uint8:
|
||
return Scalar{i: foldExtremeNarrow(a.u8s, n, greater)}, nil
|
||
case Int16:
|
||
return Scalar{i: foldExtremeNarrow(a.i16s, n, greater)}, nil
|
||
case Uint16:
|
||
return Scalar{i: foldExtremeNarrow(a.u16s, n, greater)}, nil
|
||
case Int32:
|
||
return Scalar{i: foldExtremeNarrow(a.i32s, n, greater)}, nil
|
||
case Uint32:
|
||
return Scalar{i: foldExtremeNarrow(a.u32s, n, greater)}, nil
|
||
}
|
||
// The float folds read the payload slices directly: float16 and
|
||
// float32 widen exactly, so the raw walk sees the same values the
|
||
// accessor would hand over, in the same order.
|
||
var f float64
|
||
switch a.dt {
|
||
case Float16:
|
||
f = foldExtreme(n, greater, func(lo, hi int) (float64, bool) {
|
||
return halfExtreme(a.halves[lo:hi], greater)
|
||
})
|
||
case Float32:
|
||
f = foldExtreme(n, greater, func(lo, hi int) (float64, bool) {
|
||
return float32Extreme(a.floats32[lo:hi], greater)
|
||
})
|
||
default:
|
||
f = foldExtreme(n, greater, func(lo, hi int) (float64, bool) {
|
||
return floatExtreme(a.floats[lo:hi], greater)
|
||
})
|
||
}
|
||
return Scalar{isFloat: true, f: f}, nil
|
||
}
|
||
|
||
// floatExtreme is the serial fold's rule applied to one chunk of a float64
|
||
// payload: skip the leading NaNs to seed, then keep the first strictly
|
||
// better value. A chunk holding no non-NaN reports its last element, where
|
||
// the serial seed walk would have stopped, and says there is no candidate.
|
||
func floatExtreme(src []float64, greater bool) (float64, bool) {
|
||
best, k := src[0], 1
|
||
for math.IsNaN(best) && k < len(src) {
|
||
best = src[k]
|
||
k++
|
||
}
|
||
if math.IsNaN(best) {
|
||
return src[len(src)-1], false
|
||
}
|
||
if greater {
|
||
for _, v := range src[k:] {
|
||
if v > best {
|
||
best = v
|
||
}
|
||
}
|
||
} else {
|
||
for _, v := range src[k:] {
|
||
if v < best {
|
||
best = v
|
||
}
|
||
}
|
||
}
|
||
return best, true
|
||
}
|
||
|
||
// float32Extreme is floatExtreme over a float32 payload; every widening is
|
||
// exact, so the comparisons see the values an accessor read would.
|
||
func float32Extreme(src []float32, greater bool) (float64, bool) {
|
||
best, k := float64(src[0]), 1
|
||
for math.IsNaN(best) && k < len(src) {
|
||
best = float64(src[k])
|
||
k++
|
||
}
|
||
if math.IsNaN(best) {
|
||
return float64(src[len(src)-1]), false
|
||
}
|
||
if greater {
|
||
for _, v := range src[k:] {
|
||
if w := float64(v); w > best {
|
||
best = w
|
||
}
|
||
}
|
||
} else {
|
||
for _, v := range src[k:] {
|
||
if w := float64(v); w < best {
|
||
best = w
|
||
}
|
||
}
|
||
}
|
||
return best, true
|
||
}
|
||
|
||
// halfExtreme is floatExtreme over a half-precision payload.
|
||
func halfExtreme(src []uint16, greater bool) (float64, bool) {
|
||
best, k := HalfToFloat64(src[0]), 1
|
||
for math.IsNaN(best) && k < len(src) {
|
||
best = HalfToFloat64(src[k])
|
||
k++
|
||
}
|
||
if math.IsNaN(best) {
|
||
return HalfToFloat64(src[len(src)-1]), false
|
||
}
|
||
if greater {
|
||
for _, v := range src[k:] {
|
||
if w := HalfToFloat64(v); w > best {
|
||
best = w
|
||
}
|
||
}
|
||
} else {
|
||
for _, v := range src[k:] {
|
||
if w := HalfToFloat64(v); w < best {
|
||
best = w
|
||
}
|
||
}
|
||
}
|
||
return best, true
|
||
}
|
||
|
||
// foldExtreme runs a chunk fold over the fixed partition of n elements
|
||
// and combines the partials with the serial walk's rule, so the answer
|
||
// is the serial walk's element whatever the worker count. A chunk with
|
||
// no candidate contributes nothing; an array with no candidate at all
|
||
// is all NaN and answers the last element, which is what the final
|
||
// chunk stored.
|
||
func foldExtreme[N int64 | float64](n int, greater bool, chunk func(lo, hi int) (N, bool)) N {
|
||
parts := foldParts(n)
|
||
vals := make([]N, parts)
|
||
oks := make([]bool, parts)
|
||
if parts == 1 {
|
||
vals[0], oks[0] = chunk(0, n)
|
||
} else {
|
||
engine.Parallel(parts, func(cs, ce int) {
|
||
for c := cs; c < ce; c++ {
|
||
vals[c], oks[c] = chunk(c*n/parts, (c+1)*n/parts)
|
||
}
|
||
})
|
||
}
|
||
return combineExtreme(vals, oks, greater)
|
||
}
|
||
|
||
// combineExtreme is foldExtreme's combination: the partials compare in
|
||
// index order with the strict rule the serial walk used, so a tie keeps
|
||
// the earlier block's value; a block with no candidate contributes
|
||
// nothing; a whole with no candidate answers the last block's stored
|
||
// fallback.
|
||
func combineExtreme[N int64 | float64](vals []N, oks []bool, greater bool) N {
|
||
best, have := vals[0], false
|
||
for c := range len(vals) {
|
||
if !oks[c] {
|
||
continue
|
||
}
|
||
if !have || (greater && vals[c] > best) || (!greater && vals[c] < best) {
|
||
best, have = vals[c], true
|
||
}
|
||
}
|
||
if have {
|
||
return best
|
||
}
|
||
return vals[len(vals)-1]
|
||
}
|
||
|
||
// foldExtremeNarrow applies reduceOrder's chunk rule to a narrow integer
|
||
// payload: the comparison runs in the payload's own type, so no value
|
||
// ever meets a float64 rounding, and the winner widens exactly into the
|
||
// int64 the fold combines.
|
||
func foldExtremeNarrow[T int8 | uint8 | int16 | uint16 | int32 | uint32](src []T, n int, greater bool) int64 {
|
||
return foldExtreme(n, greater, func(lo, hi int) (int64, bool) {
|
||
is := src[lo:hi]
|
||
best := is[0]
|
||
if greater {
|
||
for _, v := range is[1:] {
|
||
if v > best {
|
||
best = v
|
||
}
|
||
}
|
||
} else {
|
||
for _, v := range is[1:] {
|
||
if v < best {
|
||
best = v
|
||
}
|
||
}
|
||
}
|
||
return int64(best), true
|
||
})
|
||
}
|
||
|
||
// foldExtremeBool is foldExtremeNarrow for a bool payload: false below
|
||
// true, widened to the 0/1 the Int scalar carries. Go orders no bool
|
||
// with < or >, so the strict improvement writes its own logic.
|
||
func foldExtremeBool(src []bool, n int, greater bool) int64 {
|
||
return foldExtreme(n, greater, func(lo, hi int) (int64, bool) {
|
||
is := src[lo:hi]
|
||
best := is[0]
|
||
if greater {
|
||
for _, v := range is[1:] {
|
||
if v && !best {
|
||
best = v
|
||
}
|
||
}
|
||
} else {
|
||
for _, v := range is[1:] {
|
||
if !v && best {
|
||
best = v
|
||
}
|
||
}
|
||
}
|
||
if best {
|
||
return 1, true
|
||
}
|
||
return 0, true
|
||
})
|
||
}
|
||
|
||
// foldPartsOver runs a per-chunk fold over the fixed partition of n
|
||
// elements and combines the results through the balanced partial tree:
|
||
// the only ordering step, and its shape depends on the chunk count
|
||
// alone. The chunk boundaries c·n/parts are a function of the length and
|
||
// the part count, so the total is the same on any machine and under any
|
||
// worker setting.
|
||
func foldPartsOver[N int64 | float64 | complex128](n, parts int, fold func(c int) N) N {
|
||
if parts == 1 {
|
||
return fold(0)
|
||
}
|
||
partials := make([]N, parts)
|
||
engine.Parallel(parts, func(cs, ce int) {
|
||
for c := cs; c < ce; c++ {
|
||
partials[c] = fold(c)
|
||
}
|
||
})
|
||
return treeSum(partials)
|
||
}
|
||
|
||
// intFoldDot is the integer dot product: wrapping addition is
|
||
// associative, so the partition cannot change the value.
|
||
func intFoldDot(x, y []int64) int64 {
|
||
n := len(x)
|
||
parts := foldParts(n)
|
||
return foldPartsOver(n, parts, func(c int) int64 {
|
||
lo, hi := c*n/parts, (c+1)*n/parts
|
||
var s int64
|
||
for i := lo; i < hi; i++ {
|
||
s += x[i] * y[i]
|
||
}
|
||
return s
|
||
})
|
||
}
|
||
|
||
// narrowFoldDot is the same-dtype narrow integer dot product: the
|
||
// products run in int64 over the exact widenings, exactly what the
|
||
// accessor fold's intAt reads produce, so the partition and the value
|
||
// are the accessor fold's own.
|
||
func narrowFoldDot[T int8 | uint8 | int16 | uint16 | int32 | uint32](x, y []T) int64 {
|
||
n := len(x)
|
||
parts := foldParts(n)
|
||
return foldPartsOver(n, parts, func(c int) int64 {
|
||
lo, hi := c*n/parts, (c+1)*n/parts
|
||
var s int64
|
||
for i := lo; i < hi; i++ {
|
||
s += int64(x[i]) * int64(y[i])
|
||
}
|
||
return s
|
||
})
|
||
}
|
||
|
||
// foldDotF64 is the float64 dot product: four interleaved product chains
|
||
// per chunk, combined as a balanced pair.
|
||
func foldDotF64(x, y []float64) float64 {
|
||
n := len(x)
|
||
parts := foldParts(n)
|
||
return foldPartsOver(n, parts, func(c int) float64 {
|
||
lo, hi := c*n/parts, (c+1)*n/parts
|
||
return foldDotRange(x[lo:hi], y[lo:hi])
|
||
})
|
||
}
|
||
|
||
// foldDotF32 is the float32 dot product: every product is exact in
|
||
// float64, so the accumulation sees the values an accessor read would.
|
||
func foldDotF32(x, y []float32) float64 {
|
||
n := len(x)
|
||
parts := foldParts(n)
|
||
return foldPartsOver(n, parts, func(c int) float64 {
|
||
lo, hi := c*n/parts, (c+1)*n/parts
|
||
return foldDotRangeF32(x[lo:hi], y[lo:hi])
|
||
})
|
||
}
|
||
|
||
// foldDotF16 is the half-precision dot product, widening each element
|
||
// exactly as it is read.
|
||
func foldDotF16(x, y []uint16) float64 {
|
||
n := len(x)
|
||
parts := foldParts(n)
|
||
return foldPartsOver(n, parts, func(c int) float64 {
|
||
lo, hi := c*n/parts, (c+1)*n/parts
|
||
return foldDotRangeF16(x[lo:hi], y[lo:hi])
|
||
})
|
||
}
|
||
|
||
// foldDotC128 is the complex dot product with the same partition.
|
||
func foldDotC128(x, y []complex128) complex128 {
|
||
n := len(x)
|
||
parts := foldParts(n)
|
||
return foldPartsOver(n, parts, func(c int) complex128 {
|
||
lo, hi := c*n/parts, (c+1)*n/parts
|
||
return foldDotRangeC128(x[lo:hi], y[lo:hi])
|
||
})
|
||
}
|