Files
tensor/internal/core/reduction2.go
T

880 lines
26 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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}
}