// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package stats import ( "sourcedock.dev/petrbalvin/tensor/internal/base" "sourcedock.dev/petrbalvin/tensor/internal/core" ) import ( "math" "math/bits" "slices" "sync" "sourcedock.dev/petrbalvin/tensor/internal/engine" ) // The descriptive summaries of the package. Both entry points follow // the standard conventions: Median always returns float (averaging the // two middle values on even length) and Std is the population standard // deviation. // rawFloats returns the array's float64 payload when a is a dense // float64 array and nil otherwise: hot loops branch once on the result // and sweep the payload directly, falling back to the widening // accessor for views and other dtypes. The elements are identical // either way, so every raw sweep computes the same bits as the // accessor walk it replaces. func rawFloats(a *core.Array) []float64 { if !a.Strided() && a.Dtype() == core.Float { return a.RawFloats() } return nil } // Median returns the median as a float64, averaging the two middle values // when the length is even; an empty, non-finite or complex array is an // error. func Median(a *core.Array) (float64, error) { if a.Dtype() == core.Complex { return 0, base.Errf("Median: complex arrays have no median") } if a.Len() == 0 { return 0, base.Errf("Median: an empty array has no median") } if err := checkFinite("Median", "the sample", a); err != nil { return 0, err } vals := make([]float64, a.Len()) if fs := rawFloats(a); fs != nil { copy(vals, fs) } else if a.Dtype() == core.Int { // An int sample sorts and averages in its own type: widening // first rounds every value above 2^53, and the median of // {2^53, 2^53+1} came back as 2^53 instead of 2^53+0.5. iv := make([]int64, a.Len()) for i := range iv { v, ierr := core.IntAt(a, i) if ierr != nil { return 0, base.Errf("Median: %w", ierr) } iv[i] = v } slices.Sort(iv) n := len(iv) if n%2 == 1 { return float64(iv[n/2]), nil } // (x+y)/2 as halved magnitudes plus the carried halves: exact // whenever the average is representable, correctly rounded // beyond, and immune to the int64 sum overflow. x, y := iv[n/2-1], iv[n/2] return float64((x>>1)+(y>>1)) + float64((x&1)+(y&1))/2, nil } else { for i := range vals { vals[i] = a.FloatAt(i) } } slices.Sort(vals) n := len(vals) if n%2 == 1 { return vals[n/2], nil } // (lo+hi)/2 as halved magnitudes summed: each division by two is // exact, so the sum cannot overflow where the average itself is // representable, Median([MaxFloat64, MaxFloat64]) being MaxFloat64 // rather than the +Inf the literal sum produces. The int64 branch // above is the same averaging shape in integer arithmetic. lo, hi := vals[n/2-1], vals[n/2] return lo/2 + hi/2, nil } // sqDeviationBlock folds one block's (v − mean)² over vals[lo:hi] in // the shape the core block fold keeps: four interleaved chains, so the // adds of a long block overlap instead of queueing on one adder, // combined as ((s0+s1)+(s2+s3)). func sqDeviationBlock(vals []float64, mean float64, lo, hi int) float64 { var s0, s1, s2, s3 float64 i := lo for ; i+4 <= hi; i += 4 { d0 := vals[i] - mean d1 := vals[i+1] - mean d2 := vals[i+2] - mean d3 := vals[i+3] - mean s0 += d0 * d0 s1 += d1 * d1 s2 += d2 * d2 s3 += d3 * d3 } for ; i < hi; i++ { d := vals[i] - mean s0 += d * d } return (s0 + s1) + (s2 + s3) } // sqDeviationBlockAt is sqDeviationBlock over an accessor walk: the // elements are the ones FloatAt returns, so the block answers the same // bits the payload walk answers. func sqDeviationBlockAt(a *core.Array, mean float64, lo, hi int) float64 { var s0, s1, s2, s3 float64 i := lo for ; i+4 <= hi; i += 4 { d0 := a.FloatAt(i) - mean d1 := a.FloatAt(i+1) - mean d2 := a.FloatAt(i+2) - mean d3 := a.FloatAt(i+3) - mean s0 += d0 * d0 s1 += d1 * d1 s2 += d2 * d2 s3 += d3 * d3 } for ; i < hi; i++ { d := a.FloatAt(i) - mean s0 += d * d } return (s0 + s1) + (s2 + s3) } // sqDeviations sums (v − mean)² over vals through the canonical // partition the core reductions keep: fixed blocks of the length alone, // one partial per block, the partials combined through the balanced // tree. The squared deviations are all non-negative, but a single chain // a million long still sheds the rounding of every add against an // accumulator already near the total, and the partition shortens each // chain by the block count: measured against an exact referent at // n = 2^20 the error falls by roughly an order of magnitude on // adversarial magnitude orders, and the fold answers the same bits // whatever the worker count, the partition being a function of the // length alone. func sqDeviations(vals []float64, mean float64) float64 { n := len(vals) parts := core.FoldParts(n) if parts == 1 { return sqDeviationBlock(vals, mean, 0, n) } partials := make([]float64, parts) for c := range parts { partials[c] = sqDeviationBlock(vals, mean, core.FoldBoundary(n, c), core.FoldBoundary(n, c+1)) } return core.TreeSum(partials) } // sqDeviationsAt is sqDeviations over an accessor walk, the same // partition and the same block shape, so both routes answer identical // bits for identical elements. func sqDeviationsAt(a *core.Array, mean float64) float64 { n := a.Len() parts := core.FoldParts(n) if parts == 1 { return sqDeviationBlockAt(a, mean, 0, n) } partials := make([]float64, parts) for c := range parts { partials[c] = sqDeviationBlockAt(a, mean, core.FoldBoundary(n, c), core.FoldBoundary(n, c+1)) } return core.TreeSum(partials) } // Std returns the population standard deviation (ddof = 0); an empty or // complex array is an error. func Std(a *core.Array) (float64, error) { if a.Dtype() == core.Complex { return 0, base.Errf("Std: complex arrays have no float standard deviation") } if a.Len() == 0 { return 0, base.Errf("Std: an empty array has no standard deviation") } mean, _ := core.Mean(a) var sum float64 if fs := rawFloats(a); fs != nil { sum = sqDeviations(fs[:a.Len()], mean) } else { sum = sqDeviationsAt(a, mean) } return math.Sqrt(sum / float64(a.Len())), nil } // Var returns the population variance (ddof = 0, Std squared); an empty // or complex array is an error. func Var(a *core.Array) (float64, error) { return variance(a, 0) } // VarSample returns the unbiased sample variance (ddof = 1); fewer than // two elements, or a complex array, is an error. func VarSample(a *core.Array) (float64, error) { return variance(a, 1) } // variance computes the squared deviation from the mean with the given // ddof. func variance(a *core.Array, ddof int) (float64, error) { name := "Var" if ddof == 1 { name = "VarSample" } if a.Dtype() == core.Complex { return 0, base.Errf("%s: complex arrays have no float variance", name) } if a.Len() <= ddof { return 0, base.Errf("%s: needs more than %d element(s)", name, ddof) } mean, _ := core.Mean(a) var sum float64 if fs := rawFloats(a); fs != nil { sum = sqDeviations(fs[:a.Len()], mean) } else { sum = sqDeviationsAt(a, mean) } return sum / float64(a.Len()-ddof), nil } // maxHistBins bounds the bin count a histogram may request. The edges // and the counts together cost sixteen bytes per bin, so a million bins // is already a sixteen-megabyte answer to a question no histogram plot // asks; a larger request is refused instead of handed to the allocator, // which also keeps every xBins*yBins product far inside an int. const maxHistBins = 1 << 20 // histSerialSamples is the sample count one counting worker must carry // before the sweep splits across goroutines, and histBinsPerWorker the // number of samples it must carry per bin on top of that: a worker // counts into a private array of one cell per bin, and the split only // pays where the samples counted outweigh the cells that array costs. // Together they bound the total private scratch by the sample's own // size. const ( histSerialSamples = 1 << 10 histBinsPerWorker = 8 ) // histCountBins adds one slice of the sample to counts: the bin is the // sample's position on [lo, lo + bins·width], the maximum folds into the // top bin and anything below the range into the first. func histCountBins(vals []float64, lo, width float64, bins int, counts []int64) { for _, v := range vals { bin := int((v - lo) / width) if bin >= bins { bin = bins - 1 // the maximum lands in the top bin } if bin < 0 { bin = 0 } counts[bin]++ } } // Histogram bins the values over [min, max] into bins equal-width bins. // It returns int counts of length bins and float edges of length bins+1. // The top bin includes the maximum; an all-equal sample widens to // [v-0.5, v+0.5]; bins < 1, more than maxHistBins bins, an empty array, // or a non-finite sample is an error. func Histogram(a *core.Array, bins int) (*core.Array, *core.Array, error) { if a.Dtype() == core.Complex { return nil, nil, base.Errf("Histogram: complex arrays have no histogram") } if a.Len() == 0 { return nil, nil, base.Errf("Histogram: an empty array has no histogram") } if bins < 1 { return nil, nil, base.Errf("Histogram: needs at least one bin, got %d", bins) } if bins > maxHistBins { return nil, nil, base.Errf("Histogram: %d bins exceed the %d-bin limit", bins, maxHistBins) } // Dense float64 payloads are swept through the raw slice: the // elements are the ones FloatAt returns, without the per-element // dtype dispatch. There the finiteness scan and the range come out of // one pass; every other layout keeps the accessor scan and the two // reduction passes. The int branch of the counting below is reached // only when fs is nil, the branch that fills these two scalars. fs := rawFloats(a) var minS, maxS core.Scalar var lo, hi float64 if fs != nil { // Element i of a dense array sits at payload index i, and a // rebased view's payload may run past its own count: the walk is // bounded by Len, the elements a caller can see. fs = fs[:a.Len()] lo, hi = fs[0], fs[0] for i, v := range fs { if math.IsNaN(v) || math.IsInf(v, 0) { return nil, nil, base.Errf("Histogram: sample %d is not finite (%g)", i, v) } if v < lo { lo = v } if v > hi { hi = v } } } else { for i := range a.Len() { if v := a.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { return nil, nil, base.Errf("Histogram: sample %d is not finite (%g)", i, v) } } var err error if minS, err = core.Min(a); err != nil { return nil, nil, err } if maxS, err = core.Max(a); err != nil { return nil, nil, err } lo, hi = minS.Float(), maxS.Float() } if lo == hi { lo -= 0.5 hi += 0.5 } width := (hi - lo) / float64(bins) if math.IsInf(width, 0) || math.IsNaN(width) { // A sample holding both float extremes spans more than the // float64 range: no finite edges exist, and the int((v−lo)/width) // binning below would clamp every sample into bin 0 over ±Inf // edges. Refuse rather than publish that. return nil, nil, base.Errf("Histogram: the sample spans more than the float64 range (%g to %g)", lo, hi) } edges := make([]float64, bins+1) for i := range bins + 1 { edges[i] = lo + float64(i)*width } counts := make([]int64, bins) if fs != nil { // Counting a bin is exact integer arithmetic and every sample // carries exactly one increment, so the sweep splits over disjoint // slices of the payload and the private counters merge in any // order into the totals a single pass produces. minPerWorker := max(histSerialSamples, histBinsPerWorker*bins) var mu sync.Mutex engine.ParallelMin(len(fs), minPerWorker, func(start, end int) { if start == 0 && end == len(fs) { // The whole payload runs inline, the worker policy having // found the split uneconomical: count straight into the // result. histCountBins(fs[start:end], lo, width, bins, counts) return } local := make([]int64, bins) histCountBins(fs[start:end], lo, width, bins, local) mu.Lock() for i, c := range local { counts[i] += c } mu.Unlock() }) } else if a.Dtype() == core.Int && maxS.Int() > minS.Int() { // An int sample is binned on exact integer arithmetic: // widening v to float64 rounds above 2^53 and has misbinned // legal samples (nanosecond timestamps live there). The bin is // floor((v−lo)·bins/(hi−lo)) over the uint64 modular distance, // the quotient taken from float64 and corrected against exact // 128-bit products, which matches the rational equal-width // edges the float path approximates. An all-equal sample took // the widened float range above and stays on the float path. loI, hiI := minS.Int(), maxS.Int() span := uint64(hiI) - uint64(loI) for i := range a.Len() { v, verr := core.IntAt(a, i) if verr != nil { return nil, nil, base.Errf("Histogram: %w", verr) } d := uint64(v) - uint64(loI) bin := int(float64(d) * float64(bins) / float64(span)) // One exact correction step each way: the float quotient is // within a few 2^-32 of the true index, so at most one // boundary is crossed. dh, dl := bits.Mul64(d, uint64(bins)) for { qh, ql := bits.Mul64(uint64(bin+1), span) if qh < dh || (qh == dh && ql <= dl) { bin++ continue } qh, ql = bits.Mul64(uint64(bin), span) if bin > 0 && (qh > dh || (qh == dh && ql > dl)) { bin-- continue } break } if bin >= bins { bin = bins - 1 // the maximum lands in the top bin } if bin < 0 { bin = 0 } counts[bin]++ } } else { for i := range a.Len() { v := a.FloatAt(i) bin := int((v - lo) / width) if bin >= bins { bin = bins - 1 // the maximum lands in the top bin } if bin < 0 { bin = 0 } counts[bin]++ } } countsArr, err := core.FromInts(counts, bins) if err != nil { return nil, nil, err } edgesArr, err := core.FromFloats(edges, bins+1) if err != nil { return nil, nil, err } return countsArr, edgesArr, nil } // BinCounts counts values falling into uniformly spaced bins over // [min, max]. func BinCounts(a *core.Array, bins int) (*core.Array, error) { countsArr, _, err := Histogram(a, bins) return countsArr, err } // floatsToArray copies data into a new float array of the given shape. func floatsToArray(data []float64, shape []int) *core.Array { out := core.New(core.Float, shape...) copy(out.RawFloats(), data) return out } // checkFinite scans a real-valued array for a non-finite entry and // names it, the refusal every estimation entry point of the package // shares: a single NaN or ±Inf would otherwise spread silently through // the whole result. The label describes the array in the caller's own // vocabulary ("the sample", "the covariance"). Callers must have // refused complex input first, because the scan reads through FloatAt. func checkFinite(name, label string, a *core.Array) error { if fs := rawFloats(a); fs != nil { // A rebased view's payload runs past its own count: only the // visible elements are scanned. for _, v := range fs[:a.Len()] { if math.IsNaN(v) || math.IsInf(v, 0) { return base.Errf("%s: %s holds the non-finite value %g", name, label, v) } } return nil } for i := range a.Len() { if v := a.FloatAt(i); math.IsNaN(v) || math.IsInf(v, 0) { return base.Errf("%s: %s holds the non-finite value %g", name, label, v) } } return nil } // Quantile computes q quantiles (each in [0, 1]) of the sample using // linear interpolation on the sorted values; an empty, non-finite or // complex array is an error. func Quantile(a *core.Array, qs []float64) (*core.Array, error) { if a.Dtype() == core.Complex { return nil, base.Errf("Quantile: complex arrays have no ordering") } n := a.Len() if n == 0 { return nil, base.Errf("Quantile: empty array has no quantiles") } // A NaN sorts into an arbitrary position and an infinity drags the // interpolation across it, so the refusal comes before the sort, // as in every other estimation entry point. if err := checkFinite("Quantile", "the sample", a); err != nil { return nil, err } sortedVals, err := core.Sort(a) if err != nil { return nil, err } intSample := a.Dtype() == core.Int vals := make([]float64, len(qs)) for qi, q := range qs { // NaN-rejecting on purpose: NaN compares false against both // bounds, so `q < 0 || q > 1` would let it through to the index // conversion, where int(NaN) is a meaningless index. if !(q >= 0 && q <= 1) { return nil, base.Errf("Quantile: q must be in [0, 1], got %v", q) } pos := q * float64(n-1) lo := int(math.Floor(pos)) frac := pos - float64(lo) if intSample { // Interpolate on the exact integer difference: widening the // endpoints first rounds both above 2^53, where a rounded // pair can even collapse to equal floats and flatten the // interpolation entirely. x, xerr := core.IntAt(sortedVals, lo) if xerr != nil { return nil, base.Errf("Quantile: %w", xerr) } if lo+1 >= n || frac == 0 { vals[qi] = float64(x) continue } y, yerr := core.IntAt(sortedVals, lo+1) if yerr != nil { return nil, base.Errf("Quantile: %w", yerr) } d := y - x // exact unless the sorted pair spans more than 2^63 if (x < 0) != (y < 0) && d < 0 { vals[qi] = float64(x) + frac*(float64(y)-float64(x)) continue } vals[qi] = float64(x) + frac*float64(d) continue } v := sortedVals.FloatAt(lo) if lo+1 < n { v += frac * (sortedVals.FloatAt(lo+1) - v) } vals[qi] = v } return core.FromFloats(vals, len(qs)) }