Files
tensor/stats/stats.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

544 lines
17 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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))
}