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