880 lines
26 KiB
Go
880 lines
26 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package core
|
|
|
|
import "math"
|
|
|
|
// Extended reductions: prefix scans (CumSum, CumProd) along one
|
|
// dimension, product along one dimension (Prod), and the Lp norm along
|
|
// one dimension (Norm). The cumulative scans return the same dtype as
|
|
// the input; int scans wrap on overflow (consistent with the rest of
|
|
// the library). Norm always returns float64 (the result
|
|
// even for int inputs: Lp norms need a real-valued magnitude).
|
|
|
|
// CumSum returns the cumulative sum along dim; the result has the same
|
|
// shape as a.
|
|
func CumSum(a *Array, dim int) (*Array, error) {
|
|
return a.scanDim(dim, "CumSum", false)
|
|
}
|
|
|
|
// CumProd returns the cumulative product along dim.
|
|
func CumProd(a *Array, dim int) (*Array, error) {
|
|
return a.scanDim(dim, "CumProd", true)
|
|
}
|
|
|
|
// Prod returns the product along dim; keepDim preserves the reduced
|
|
// dimension as size 1.
|
|
func Prod(a *Array, dim int, keepDim bool) (*Array, error) {
|
|
// The one-dimensional global product is the shape the sharded
|
|
// reductions mirror: it folds the canonical blocks where they lie
|
|
// and combines them through the same balanced tree, so a sharded
|
|
// product and this product are one computation. Lines of higher
|
|
// ranks keep the sequential walk below.
|
|
if a.NDim() == 1 && dim == 0 && a.isContiguous() {
|
|
out, err := prodGlobal1D(a)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if keepDim {
|
|
return keepReducedDim(out, a.shape, dim), nil
|
|
}
|
|
return out, nil
|
|
}
|
|
out, err := a.reduceDimProd(dim, "Prod")
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if keepDim {
|
|
return keepReducedDim(out, a.shape, dim), nil
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
// prodGlobal1D folds the whole one-dimensional array through the
|
|
// canonical partition: each block multiplies into one partial and the
|
|
// partials combine through the balanced tree, the exact shape the spmd
|
|
// shards reproduce. The integer product is exact under any grouping;
|
|
// the floating products round differently from the single chain at
|
|
// lengths past one block, measured against the exact referent in the
|
|
// accuracy test's bound.
|
|
func prodGlobal1D(a *Array) (*Array, error) {
|
|
if narrowRefused(a.dt) {
|
|
return nil, errf("Prod: dtype %s is not supported; convert with Astype", a.dt)
|
|
}
|
|
if a.dt == Complex {
|
|
return nil, errf("Prod: complex arrays have no real-valued product")
|
|
}
|
|
n := a.Len()
|
|
out := &Array{shape: []int{1}, dt: a.dt}
|
|
out.alloc(1)
|
|
parts := foldParts(n)
|
|
switch a.dt {
|
|
case Int:
|
|
ints := make([]int64, parts)
|
|
for c := range parts {
|
|
ints[c] = foldProdRangeI64(a.ints[c*n/parts : (c+1)*n/parts])
|
|
}
|
|
out.ints[0] = treeProd(ints)
|
|
case Float16:
|
|
// Each block value is an exact half carried in float64, and the
|
|
// tree combines narrow through half, the per-step rounding the
|
|
// line product keeps.
|
|
vals := make([]float64, parts)
|
|
for c := range parts {
|
|
vals[c] = foldProdRangeF16(a.halves[c*n/parts : (c+1)*n/parts])
|
|
}
|
|
out.halves[0] = HalfFromFloat64(treeProdHalf(vals))
|
|
case Float32:
|
|
// The partials multiply natively in float32, the way the line
|
|
// product keeps.
|
|
vals := make([]float32, parts)
|
|
for c := range parts {
|
|
vals[c] = foldProdRangeF32(a.floats32[c*n/parts : (c+1)*n/parts])
|
|
}
|
|
out.floats32[0] = treeProd(vals)
|
|
case Float:
|
|
vals := make([]float64, parts)
|
|
for c := range parts {
|
|
vals[c] = foldProdRange(a.floats[c*n/parts : (c+1)*n/parts])
|
|
}
|
|
out.floats[0] = treeProd(vals)
|
|
default:
|
|
return nil, errf("Prod: dtype %s is not supported; convert with Astype", a.dt)
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
// Norm returns the Lp norm along dim: (sum |x|^p)^(1/p). p must be a
|
|
// positive number; p == math.Inf returns the max-abs, where NaN
|
|
// elements never win and a line whose every element is NaN reports
|
|
// NaN, exactly as MinAxis and MaxAxis fold their extrema. NaN is
|
|
// rejected like any other invalid p, because every comparison against
|
|
// NaN is false and it would otherwise slip through the positivity
|
|
// test. The result is always float64, dtype promoted.
|
|
//
|
|
// Lines are independent and each line writes only its own destination
|
|
// slot, so the walk splits across workers with the same ascending
|
|
// addend order per line, and the fan-out is capped so a small norm
|
|
// never pays a spawn bill. The general exponent calls math.Pow per
|
|
// element, so its floor is far lower than the plain folds'.
|
|
func Norm(a *Array, p float64, dim int, keepDim bool) (*Array, error) {
|
|
if math.IsNaN(p) || (p <= 0 && !math.IsInf(p, 1)) {
|
|
return nil, errf("Norm: p must be positive, got %v", p)
|
|
}
|
|
if a.dt == Complex {
|
|
return nil, errf("Norm: complex arrays have no real-valued norm")
|
|
}
|
|
if narrowRefused(a.dt) {
|
|
// The narrow dtypes stay refused for the Norm family.
|
|
return nil, errf("Norm: dtype %s is not supported; convert with Astype", a.dt)
|
|
}
|
|
if dim < 0 || dim >= a.NDim() {
|
|
return nil, errf("Norm: dimension %d out of range for shape %s", dim, shapeText(a.shape))
|
|
}
|
|
// The one-dimensional global norm for a finite p is the shape the
|
|
// sharded reductions mirror: the power sums fold the canonical
|
|
// blocks where they lie and combine through the same balanced tree,
|
|
// so a sharded norm and this norm are one computation. The infinity
|
|
// norm is a maximum, not a sum, and keeps the line walk below.
|
|
if a.NDim() == 1 && dim == 0 && a.isContiguous() && !math.IsInf(p, 1) {
|
|
out, err := normGlobal1D(a, p)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if keepDim {
|
|
return keepReducedDim(out, a.shape, dim), nil
|
|
}
|
|
return out, nil
|
|
}
|
|
outShape := reduceShape(a.shape, dim)
|
|
out := &Array{shape: outShape, dt: Float}
|
|
total := 1
|
|
for _, d := range outShape {
|
|
total *= d
|
|
}
|
|
out.alloc(total)
|
|
stride := 1
|
|
for k := dim + 1; k < a.NDim(); k++ {
|
|
stride *= a.shape[k]
|
|
}
|
|
// The walk goes line by line: a line is the run of a.shape[dim]
|
|
// elements that share every surviving coordinate. Each line is
|
|
// gathered and folded through the canonical partition the global
|
|
// norm uses: fixed blocks, the block partials combined through the
|
|
// balanced tree. The partition follows from the line length alone,
|
|
// so a single-line norm answers normGlobal1D's exact bits whatever
|
|
// the shape or the worker split.
|
|
line := a.shape[dim]
|
|
perLine := stride * line
|
|
lines := 0
|
|
if perLine > 0 {
|
|
lines = a.Len() / perLine
|
|
}
|
|
pInf := math.IsInf(p, 1)
|
|
if pInf && line == 0 {
|
|
// An empty reduced dimension has no maximum: MinAxis and
|
|
// MaxAxis report NaN for it, and the infinity norm agrees
|
|
// rather than reporting the untouched zero.
|
|
for i := range out.floats {
|
|
out.floats[i] = math.NaN()
|
|
}
|
|
return out, nil
|
|
}
|
|
if pInf {
|
|
// The infinity norm is a maximum, not a sum: the extremum walk
|
|
// below keeps the line order, which no grouping can move.
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
switch a.dt {
|
|
case Int:
|
|
src := a.ints
|
|
for b := ls; b < le; b++ {
|
|
base := b * perLine
|
|
for s := range stride {
|
|
dst := b*stride + s
|
|
acc := out.floats[dst]
|
|
for off := range line {
|
|
if v := math.Abs(float64(src[base+off*stride+s])); v > acc {
|
|
acc = v
|
|
}
|
|
}
|
|
out.floats[dst] = acc
|
|
}
|
|
}
|
|
case Float16:
|
|
src := a.halves
|
|
for b := ls; b < le; b++ {
|
|
base := b * perLine
|
|
for s := range stride {
|
|
dst := b*stride + s
|
|
i, m := 0, 0.0
|
|
for i < line {
|
|
if v := math.Abs(HalfToFloat64(src[base+i*stride+s])); v == v {
|
|
m = v
|
|
break
|
|
}
|
|
i++
|
|
}
|
|
if i == line {
|
|
out.floats[dst] = math.NaN()
|
|
continue
|
|
}
|
|
i++
|
|
for off := i; off < line; off++ {
|
|
if v := math.Abs(HalfToFloat64(src[base+off*stride+s])); v > m {
|
|
m = v
|
|
}
|
|
}
|
|
out.floats[dst] = m
|
|
}
|
|
}
|
|
case Float32:
|
|
src := a.floats32
|
|
for b := ls; b < le; b++ {
|
|
base := b * perLine
|
|
for s := range stride {
|
|
dst := b*stride + s
|
|
i, m := 0, 0.0
|
|
for i < line {
|
|
if v := math.Abs(float64(src[base+i*stride+s])); v == v {
|
|
m = v
|
|
break
|
|
}
|
|
i++
|
|
}
|
|
if i == line {
|
|
out.floats[dst] = math.NaN()
|
|
continue
|
|
}
|
|
i++
|
|
for off := i; off < line; off++ {
|
|
if v := math.Abs(float64(src[base+off*stride+s])); v > m {
|
|
m = v
|
|
}
|
|
}
|
|
out.floats[dst] = m
|
|
}
|
|
}
|
|
default:
|
|
src := a.floats
|
|
for b := ls; b < le; b++ {
|
|
base := b * perLine
|
|
for s := range stride {
|
|
dst := b*stride + s
|
|
i, m := 0, 0.0
|
|
for i < line {
|
|
if v := math.Abs(src[base+i*stride+s]); v == v {
|
|
m = v
|
|
break
|
|
}
|
|
i++
|
|
}
|
|
if i == line {
|
|
out.floats[dst] = math.NaN()
|
|
continue
|
|
}
|
|
i++
|
|
for off := i; off < line; off++ {
|
|
if v := math.Abs(src[base+off*stride+s]); v > m {
|
|
m = v
|
|
}
|
|
}
|
|
out.floats[dst] = m
|
|
}
|
|
}
|
|
}
|
|
})
|
|
if keepDim {
|
|
return keepReducedDim(out, a.shape, dim), nil
|
|
}
|
|
return out, nil
|
|
}
|
|
pOne := p == 1
|
|
pTwo := p == 2
|
|
// One canonical power-sum fold serves every finite exponent: the
|
|
// per-element arithmetic the mode picks is the one foldNormPower
|
|
// keeps, the blocks are the canonical ones and treeSum combines the
|
|
// partials. The common exponents multiply instead of calling Pow,
|
|
// bit-identically to the general path. The dtype and the mode are
|
|
// picked once per worker segment, so the fold reads the payload
|
|
// where the elements live with no per-element call and no gather
|
|
// scratch, and the partial table is shared by the segment's lines.
|
|
const (
|
|
normP1 = iota
|
|
normP2
|
|
normPGeneral
|
|
)
|
|
mode := normPGeneral
|
|
switch {
|
|
case pOne:
|
|
mode = normP1
|
|
case pTwo:
|
|
mode = normP2
|
|
}
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
parts := foldParts(line)
|
|
partials := make([]float64, parts)
|
|
var foldRange func(base, s, lo, hi int) float64
|
|
switch a.dt {
|
|
case Int:
|
|
src := a.ints
|
|
switch mode {
|
|
case normP1:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
acc += math.Abs(float64(src[base+off*stride+s]))
|
|
}
|
|
return acc
|
|
}
|
|
case normP2:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
v := math.Abs(float64(src[base+off*stride+s]))
|
|
acc += v * v
|
|
}
|
|
return acc
|
|
}
|
|
default:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
acc += math.Pow(math.Abs(float64(src[base+off*stride+s])), p)
|
|
}
|
|
return acc
|
|
}
|
|
}
|
|
case Float16:
|
|
src := a.halves
|
|
switch mode {
|
|
case normP1:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
acc += math.Abs(HalfToFloat64(src[base+off*stride+s]))
|
|
}
|
|
return acc
|
|
}
|
|
case normP2:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
v := math.Abs(HalfToFloat64(src[base+off*stride+s]))
|
|
acc += v * v
|
|
}
|
|
return acc
|
|
}
|
|
default:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
acc += math.Pow(math.Abs(HalfToFloat64(src[base+off*stride+s])), p)
|
|
}
|
|
return acc
|
|
}
|
|
}
|
|
case Float32:
|
|
src := a.floats32
|
|
switch mode {
|
|
case normP1:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
acc += math.Abs(float64(src[base+off*stride+s]))
|
|
}
|
|
return acc
|
|
}
|
|
case normP2:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
v := math.Abs(float64(src[base+off*stride+s]))
|
|
acc += v * v
|
|
}
|
|
return acc
|
|
}
|
|
default:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
acc += math.Pow(math.Abs(float64(src[base+off*stride+s])), p)
|
|
}
|
|
return acc
|
|
}
|
|
}
|
|
default:
|
|
src := a.floats
|
|
switch mode {
|
|
case normP1:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
acc += math.Abs(src[base+off*stride+s])
|
|
}
|
|
return acc
|
|
}
|
|
case normP2:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
v := math.Abs(src[base+off*stride+s])
|
|
acc += v * v
|
|
}
|
|
return acc
|
|
}
|
|
default:
|
|
foldRange = func(base, s, lo, hi int) float64 {
|
|
var acc float64
|
|
for off := lo; off < hi; off++ {
|
|
acc += math.Pow(math.Abs(src[base+off*stride+s]), p)
|
|
}
|
|
return acc
|
|
}
|
|
}
|
|
}
|
|
for b := ls; b < le; b++ {
|
|
base := b * perLine
|
|
for s := range stride {
|
|
if parts == 1 {
|
|
out.floats[b*stride+s] = foldRange(base, s, 0, line)
|
|
continue
|
|
}
|
|
for c := range parts {
|
|
partials[c] = foldRange(base, s, c*line/parts, (c+1)*line/parts)
|
|
}
|
|
out.floats[b*stride+s] = treeSum(partials)
|
|
}
|
|
}
|
|
})
|
|
switch {
|
|
case pTwo:
|
|
for k := range out.floats {
|
|
out.floats[k] = math.Sqrt(out.floats[k])
|
|
}
|
|
case !pOne:
|
|
for k := range out.floats {
|
|
out.floats[k] = math.Pow(out.floats[k], 1/p)
|
|
}
|
|
}
|
|
if keepDim {
|
|
return keepReducedDim(out, a.shape, dim), nil
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
// scanDim walks a row-major, maintaining a per-line running value for
|
|
// dim: mul selects the product (CumProd) over the sum (CumSum). The
|
|
// combine happens in registers per line slot; each output element still
|
|
// receives the same combine sequence the off-then-s walk produced, so
|
|
// results are bit-identical, with no per-element closure call. Whole
|
|
// lines go to whole workers and each line carries its own running value,
|
|
// so the split moves no bit; the fan-out is capped so a small scan stays
|
|
// on the calling goroutine. The float64 sum carry is Neumaier
|
|
// compensated, the accuracy the compensated-scan test measured against
|
|
// the exact referent; the product and every other dtype keep the plain
|
|
// chain, and the compensation runs per line in element order, so the
|
|
// split still moves no bit.
|
|
func (a *Array) scanDim(dim int, name string, mul bool) (*Array, error) {
|
|
if dim < 0 || dim >= a.NDim() {
|
|
return nil, errf("%s: dimension %d out of range for shape %s", name, dim, shapeText(a.shape))
|
|
}
|
|
if narrowRefused(a.dt) {
|
|
// The reduction surface offers Sum, Mean, Min, Max and
|
|
// the arg extremes; the scans refuse the narrow dtypes.
|
|
return nil, errf("%s: dtype %s is not supported; convert with Astype", name, a.dt)
|
|
}
|
|
if a.shape[dim] == 0 {
|
|
return nil, errf("%s: dimension %d of shape %s is empty", name, dim, shapeText(a.shape))
|
|
}
|
|
out := &Array{shape: a.Shape(), dt: a.dt}
|
|
out.alloc(a.Len())
|
|
// A zero trailing dimension leaves no payload to scan: the result
|
|
// is the empty array, and the per-line division below would
|
|
// otherwise divide zero by zero.
|
|
if a.Len() == 0 {
|
|
return out, nil
|
|
}
|
|
stride := 1
|
|
for k := dim + 1; k < a.NDim(); k++ {
|
|
stride *= a.shape[k]
|
|
}
|
|
line := a.shape[dim]
|
|
perLine := stride * line
|
|
lines := a.Len() / perLine
|
|
switch a.dt {
|
|
case Int:
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
for blk := ls * perLine; blk < le*perLine; blk += perLine {
|
|
for s := range stride {
|
|
pos := blk + s
|
|
acc := a.ints[pos]
|
|
out.ints[pos] = acc
|
|
for off := 1; off < line; off++ {
|
|
pos += stride
|
|
if mul {
|
|
acc *= a.ints[pos]
|
|
} else {
|
|
acc += a.ints[pos]
|
|
}
|
|
out.ints[pos] = acc
|
|
}
|
|
}
|
|
}
|
|
})
|
|
case Float16:
|
|
// Compute in float64 for accuracy; each step narrows the carry
|
|
// to half before combining, exactly as the float32 path narrows
|
|
// to float32.
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
for blk := ls * perLine; blk < le*perLine; blk += perLine {
|
|
for s := range stride {
|
|
pos := blk + s
|
|
acc := HalfToFloat64(a.halves[pos])
|
|
out.halves[pos] = HalfFromFloat64(acc)
|
|
for off := 1; off < line; off++ {
|
|
pos += stride
|
|
if mul {
|
|
acc = HalfToFloat64(HalfFromFloat64(acc)) * HalfToFloat64(a.halves[pos])
|
|
} else {
|
|
acc = HalfToFloat64(HalfFromFloat64(acc)) + HalfToFloat64(a.halves[pos])
|
|
}
|
|
out.halves[pos] = HalfFromFloat64(acc)
|
|
}
|
|
}
|
|
}
|
|
})
|
|
case Float32:
|
|
// Compute in float64 for accuracy, round back. Each step
|
|
// narrows the carry to float32 before combining, exactly as the
|
|
// previous element's stored value fed the next one.
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
for blk := ls * perLine; blk < le*perLine; blk += perLine {
|
|
for s := range stride {
|
|
pos := blk + s
|
|
acc := float64(a.floats32[pos])
|
|
out.floats32[pos] = float32(acc)
|
|
for off := 1; off < line; off++ {
|
|
pos += stride
|
|
if mul {
|
|
acc = float64(float32(acc)) * float64(a.floats32[pos])
|
|
} else {
|
|
acc = float64(float32(acc)) + float64(a.floats32[pos])
|
|
}
|
|
out.floats32[pos] = float32(acc)
|
|
}
|
|
}
|
|
}
|
|
})
|
|
case Float:
|
|
if mul {
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
for blk := ls * perLine; blk < le*perLine; blk += perLine {
|
|
for s := range stride {
|
|
pos := blk + s
|
|
acc := a.floats[pos]
|
|
out.floats[pos] = acc
|
|
for off := 1; off < line; off++ {
|
|
pos += stride
|
|
acc *= a.floats[pos]
|
|
out.floats[pos] = acc
|
|
}
|
|
}
|
|
}
|
|
})
|
|
break
|
|
}
|
|
// The float64 sum carry is compensated: a running high part
|
|
// plus a correction, the output the corrected running value.
|
|
// The plain chain's accumulation error was measured against a
|
|
// 512-bit big.Float referent at 9.4e12 absolute on a 2^20
|
|
// random sample and a full loss of every small addend on the
|
|
// cancellation pattern [1, 1e100, 1, -1e100], where the
|
|
// compensated walk answers exactly; on benign data it answers
|
|
// to the last ulp. A non-finite partial freezes the
|
|
// correction, so an overflow sticks to infinity and a NaN
|
|
// poisons the tail exactly as the plain chain's would.
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
for blk := ls * perLine; blk < le*perLine; blk += perLine {
|
|
for s := range stride {
|
|
pos := blk + s
|
|
sum := a.floats[pos]
|
|
var comp float64
|
|
out.floats[pos] = sum
|
|
for off := 1; off < line; off++ {
|
|
pos += stride
|
|
x := a.floats[pos]
|
|
t := sum + x
|
|
if !math.IsInf(t, 0) && t == t {
|
|
if math.Abs(sum) >= math.Abs(x) {
|
|
comp += (sum - t) + x
|
|
} else {
|
|
comp += (x - t) + sum
|
|
}
|
|
}
|
|
sum = t
|
|
out.floats[pos] = sum + comp
|
|
}
|
|
}
|
|
}
|
|
})
|
|
default:
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
for blk := ls * perLine; blk < le*perLine; blk += perLine {
|
|
for s := range stride {
|
|
pos := blk + s
|
|
acc := a.complexes[pos]
|
|
out.complexes[pos] = acc
|
|
for off := 1; off < line; off++ {
|
|
pos += stride
|
|
if mul {
|
|
acc *= a.complexes[pos]
|
|
} else {
|
|
acc += a.complexes[pos]
|
|
}
|
|
out.complexes[pos] = acc
|
|
}
|
|
}
|
|
}
|
|
})
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
// reduceDimProd multiplies along dim, same pattern as the axis reductions
|
|
// but a product. Used by Prod.
|
|
func (a *Array) reduceDimProd(dim int, name string) (*Array, error) {
|
|
if dim < 0 || dim >= a.NDim() {
|
|
return nil, errf("%s: dimension %d out of range for shape %s", name, dim, shapeText(a.shape))
|
|
}
|
|
if narrowRefused(a.dt) {
|
|
// The narrow dtypes are refused for Prod; the refusal names
|
|
// the dtype and the conversion.
|
|
return nil, errf("%s: dtype %s is not supported; convert with Astype", name, a.dt)
|
|
}
|
|
if a.dt == Complex {
|
|
return nil, errf("%s: complex arrays have no real-valued product", name)
|
|
}
|
|
outShape := reduceShape(a.shape, dim)
|
|
out := &Array{shape: outShape, dt: a.dt}
|
|
total := 1
|
|
for _, d := range outShape {
|
|
total *= d
|
|
}
|
|
out.alloc(total)
|
|
stride := 1
|
|
for k := dim + 1; k < a.NDim(); k++ {
|
|
stride *= a.shape[k]
|
|
}
|
|
// Line by line like Norm: each line is gathered into the payload's
|
|
// own scratch and multiplied through the canonical partition the
|
|
// global product uses, one block partial per block and the balanced
|
|
// tree over them. The partition follows from the line length alone,
|
|
// so a single-line product answers prodGlobal1D's exact bits
|
|
// whatever the shape, the stride or the worker split, and the
|
|
// payload's own arithmetic (the half and float32 narrowings) is
|
|
// preserved.
|
|
line := a.shape[dim]
|
|
perLine := stride * line
|
|
lines := 0
|
|
if perLine > 0 {
|
|
lines = a.Len() / perLine
|
|
}
|
|
if line == 0 {
|
|
// An empty reduced dimension is the empty product: every line's
|
|
// answer is the multiplicative identity in the array's own dtype,
|
|
// exactly the value prodGlobal1D answers for an empty
|
|
// one-dimensional array. The zeroed allocation must not surface.
|
|
for i := range total {
|
|
switch a.dt {
|
|
case Int:
|
|
out.ints[i] = 1
|
|
case Float16:
|
|
out.halves[i] = halfOne
|
|
case Float32:
|
|
out.floats32[i] = 1
|
|
default:
|
|
out.floats[i] = 1
|
|
}
|
|
}
|
|
return out, nil
|
|
}
|
|
switch a.dt {
|
|
case Int:
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
var scratch []int64
|
|
for b := ls; b < le; b++ {
|
|
base := b * perLine
|
|
for s := range stride {
|
|
if stride == 1 {
|
|
// A contiguous line is a slice of the payload:
|
|
// the fold reads it where it lies, no gather.
|
|
out.ints[b*stride+s] = prodLine(a.ints[base+s:base+s+line], foldProdRangeI64)
|
|
continue
|
|
}
|
|
if scratch == nil {
|
|
scratch = make([]int64, line)
|
|
}
|
|
for off := range line {
|
|
scratch[off] = a.ints[base+off*stride+s]
|
|
}
|
|
out.ints[b*stride+s] = prodLine(scratch, foldProdRangeI64)
|
|
}
|
|
}
|
|
})
|
|
case Float16:
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
var scratch []uint16
|
|
for b := ls; b < le; b++ {
|
|
base := b * perLine
|
|
for s := range stride {
|
|
if stride == 1 {
|
|
out.halves[b*stride+s] = HalfFromFloat64(prodLineHalf(a.halves[base+s : base+s+line]))
|
|
continue
|
|
}
|
|
if scratch == nil {
|
|
scratch = make([]uint16, line)
|
|
}
|
|
for off := range line {
|
|
scratch[off] = a.halves[base+off*stride+s]
|
|
}
|
|
out.halves[b*stride+s] = HalfFromFloat64(prodLineHalf(scratch))
|
|
}
|
|
}
|
|
})
|
|
case Float32:
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
var scratch []float32
|
|
for b := ls; b < le; b++ {
|
|
base := b * perLine
|
|
for s := range stride {
|
|
if stride == 1 {
|
|
out.floats32[b*stride+s] = prodLine(a.floats32[base+s:base+s+line], foldProdRangeF32)
|
|
continue
|
|
}
|
|
if scratch == nil {
|
|
scratch = make([]float32, line)
|
|
}
|
|
for off := range line {
|
|
scratch[off] = a.floats32[base+off*stride+s]
|
|
}
|
|
out.floats32[b*stride+s] = prodLine(scratch, foldProdRangeF32)
|
|
}
|
|
}
|
|
})
|
|
default:
|
|
splitCapped(lines, a.Len(), reduceSplitFloor, func(ls, le int) {
|
|
var scratch []float64
|
|
for b := ls; b < le; b++ {
|
|
base := b * perLine
|
|
for s := range stride {
|
|
if stride == 1 {
|
|
out.floats[b*stride+s] = prodLine(a.floats[base+s:base+s+line], foldProdRange)
|
|
continue
|
|
}
|
|
if scratch == nil {
|
|
scratch = make([]float64, line)
|
|
}
|
|
for off := range line {
|
|
scratch[off] = a.floats[base+off*stride+s]
|
|
}
|
|
out.floats[b*stride+s] = prodLine(scratch, foldProdRange)
|
|
}
|
|
}
|
|
})
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
// prodLine multiplies one gathered line through the canonical partition:
|
|
// one block partial per block, the balanced tree over them, the shape
|
|
// prodGlobal1D multiplies the whole array with.
|
|
func prodLine[T int64 | float64 | float32](line []T, block func([]T) T) T {
|
|
parts := foldParts(len(line))
|
|
if parts == 1 {
|
|
return block(line)
|
|
}
|
|
partials := make([]T, parts)
|
|
for c := range parts {
|
|
partials[c] = block(line[c*len(line)/parts : (c+1)*len(line)/parts])
|
|
}
|
|
return treeProd(partials)
|
|
}
|
|
|
|
// prodLineHalf is prodLine for a half payload: the block partials carry
|
|
// exact half values in float64 and the tree narrows through half, the
|
|
// rounding foldProdRangeF16 and treeProdHalf keep.
|
|
func prodLineHalf(line []uint16) float64 {
|
|
parts := foldParts(len(line))
|
|
if parts == 1 {
|
|
return foldProdRangeF16(line)
|
|
}
|
|
partials := make([]float64, parts)
|
|
for c := range parts {
|
|
partials[c] = foldProdRangeF16(line[c*len(line)/parts : (c+1)*len(line)/parts])
|
|
}
|
|
return treeProdHalf(partials)
|
|
}
|
|
|
|
// normGlobal1D folds the whole one-dimensional array's power sums
|
|
// through the canonical partition: each block sums |v|^p with the
|
|
// norm fold's own per-element arithmetic and the partials combine
|
|
// through the balanced tree, the exact shape the spmd shards
|
|
// reproduce. The root closes through the same Sqrt and Pow the line
|
|
// norm keeps.
|
|
func normGlobal1D(a *Array, p float64) (*Array, error) {
|
|
n := a.Len()
|
|
parts := foldParts(n)
|
|
sums := make([]float64, parts)
|
|
switch a.dt {
|
|
case Int:
|
|
for c := range parts {
|
|
sums[c] = foldNormPowerRangeI64(a.ints[c*n/parts:(c+1)*n/parts], p)
|
|
}
|
|
case Float16:
|
|
for c := range parts {
|
|
sums[c] = foldNormPowerRangeF16(a.halves[c*n/parts:(c+1)*n/parts], p)
|
|
}
|
|
case Float32:
|
|
for c := range parts {
|
|
sums[c] = foldNormPowerRangeF32(a.floats32[c*n/parts:(c+1)*n/parts], p)
|
|
}
|
|
default:
|
|
for c := range parts {
|
|
sums[c] = foldNormPowerRange(a.floats[c*n/parts:(c+1)*n/parts], p)
|
|
}
|
|
}
|
|
out := &Array{shape: []int{1}, dt: Float}
|
|
out.alloc(1)
|
|
out.floats[0] = normRoot(treeSum(sums), p)
|
|
return out, nil
|
|
}
|
|
|
|
// reduceShape returns the shape with dim dropped. Used by the axis-style
|
|
// helpers so they all agree.
|
|
func reduceShape(shape []int, dim int) []int {
|
|
out := make([]int, 0, len(shape)-1)
|
|
out = append(out, shape[:dim]...)
|
|
out = append(out, shape[dim+1:]...)
|
|
if len(out) == 0 {
|
|
out = []int{1}
|
|
}
|
|
return out
|
|
}
|
|
|
|
// keepReducedDim returns the array with the reduced dimension reinserted
|
|
// as size 1, used by keepDim=true on Prod and Norm.
|
|
func keepReducedDim(a *Array, original []int, dim int) *Array {
|
|
sh := make([]int, 0, len(original))
|
|
sh = append(sh, original[:dim]...)
|
|
sh = append(sh, 1)
|
|
sh = append(sh, original[dim+1:]...)
|
|
// Every payload field the dtype may carry rides along: a narrow
|
|
// result must never lose its elements to a five-slice copy.
|
|
return &Array{shape: sh, dt: a.dt,
|
|
ints: a.ints, halves: a.halves, floats32: a.floats32,
|
|
floats: a.floats, complexes: a.complexes,
|
|
bools: a.bools, i8s: a.i8s, u8s: a.u8s,
|
|
i16s: a.i16s, u16s: a.u16s, i32s: a.i32s, u32s: a.u32s}
|
|
}
|