Files

544 lines
17 KiB
Go
Raw Permalink 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 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))
}