233 lines
7.0 KiB
Go
233 lines
7.0 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package stats
|
||
|
|
|
||
|
|
import (
|
||
|
|
"math"
|
||
|
|
"strings"
|
||
|
|
"testing"
|
||
|
|
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/engine"
|
||
|
|
)
|
||
|
|
|
||
|
|
func TestMedian(t *testing.T) {
|
||
|
|
odd := mustFromInts(t, []int64{3, 1, 2}, 3)
|
||
|
|
m, err := Median(odd)
|
||
|
|
if err != nil || m != 2 {
|
||
|
|
t.Fatalf("Median odd: %v %v", m, err)
|
||
|
|
}
|
||
|
|
|
||
|
|
even := mustFromInts(t, []int64{1, 2, 3, 4}, 4)
|
||
|
|
m, err = Median(even)
|
||
|
|
if err != nil || m != 2.5 {
|
||
|
|
t.Fatalf("Median even: %v %v", m, err)
|
||
|
|
}
|
||
|
|
|
||
|
|
f := mustFromFloats(t, []float64{7.5, 2.5, 5.0}, 3)
|
||
|
|
m, err = Median(f)
|
||
|
|
if err != nil || m != 5.0 {
|
||
|
|
t.Fatalf("Median float: %v %v", m, err)
|
||
|
|
}
|
||
|
|
|
||
|
|
empty := mustFromInts(t, nil, 0)
|
||
|
|
if _, err := Median(empty); err == nil || !strings.Contains(err.Error(), "empty array") {
|
||
|
|
t.Fatalf("Median empty: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestStd(t *testing.T) {
|
||
|
|
// The classic set: mean 5, population variance 4, std 2.
|
||
|
|
a := mustFromInts(t, []int64{2, 4, 4, 4, 5, 5, 7, 9}, 8)
|
||
|
|
s, err := Std(a)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Std: %v", err)
|
||
|
|
}
|
||
|
|
if math.Abs(s-2) > 1e-9 {
|
||
|
|
t.Fatalf("Std: %v", s)
|
||
|
|
}
|
||
|
|
|
||
|
|
single := mustFromFloats(t, []float64{3.5}, 1)
|
||
|
|
s, err = Std(single)
|
||
|
|
if err != nil || s != 0 {
|
||
|
|
t.Fatalf("Std single: %v %v", s, err)
|
||
|
|
}
|
||
|
|
|
||
|
|
empty := mustFromInts(t, nil, 0)
|
||
|
|
if _, err := Std(empty); err == nil || !strings.Contains(err.Error(), "empty array") {
|
||
|
|
t.Fatalf("Std empty: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestHistogram(t *testing.T) {
|
||
|
|
a := mustFromFloats(t, []float64{1, 1.5, 2, 3, 3.5, 4}, 6)
|
||
|
|
counts, edges, err := Histogram(a, 3)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Histogram: %v", err)
|
||
|
|
}
|
||
|
|
// Bins over [1, 4]: [1,2), [2,3), [3,4]; 4 lands in the top bin.
|
||
|
|
wantCounts := mustFromInts(t, []int64{2, 1, 3}, 3)
|
||
|
|
if !core.Equal(wantCounts, counts) {
|
||
|
|
t.Fatalf("Histogram counts: %s", counts)
|
||
|
|
}
|
||
|
|
if edges.Len() != 4 {
|
||
|
|
t.Fatalf("Histogram edges len: %d", edges.Len())
|
||
|
|
}
|
||
|
|
e0, _ := core.FloatAt(edges, 0)
|
||
|
|
e3, _ := core.FloatAt(edges, 3)
|
||
|
|
if e0 != 1 || e3 != 4 {
|
||
|
|
t.Fatalf("Histogram edges: %s", edges)
|
||
|
|
}
|
||
|
|
|
||
|
|
// An all-equal sample widens to [v-0.5, v+0.5].
|
||
|
|
same := mustFromInts(t, []int64{5, 5}, 2)
|
||
|
|
_, edges, err = Histogram(same, 2)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Histogram same: %v", err)
|
||
|
|
}
|
||
|
|
e0, _ = core.FloatAt(edges, 0)
|
||
|
|
e2, _ := core.FloatAt(edges, 2)
|
||
|
|
if e0 != 4.5 || e2 != 5.5 {
|
||
|
|
t.Fatalf("Histogram same edges: %s", edges)
|
||
|
|
}
|
||
|
|
|
||
|
|
empty := mustFromInts(t, nil, 0)
|
||
|
|
if _, _, err := Histogram(empty, 3); err == nil || !strings.Contains(err.Error(), "empty array") {
|
||
|
|
t.Fatalf("Histogram empty: %v", err)
|
||
|
|
}
|
||
|
|
if _, _, err := Histogram(a, 0); err == nil || !strings.Contains(err.Error(), "at least one bin") {
|
||
|
|
t.Fatalf("Histogram bins: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestHistogramNonFinite pins the finite-sample contract: a NaN or
|
||
|
|
// infinite sample makes Histogram error instead of producing a
|
||
|
|
// corrupted or empty binning.
|
||
|
|
func TestHistogramNonFinite(t *testing.T) {
|
||
|
|
withNaN := mustFromFloats(t, []float64{1, 2, math.NaN()}, 3)
|
||
|
|
if _, _, err := Histogram(withNaN, 2); err == nil {
|
||
|
|
t.Error("expected an error for a NaN sample")
|
||
|
|
}
|
||
|
|
withInf := mustFromFloats(t, []float64{1, 2, math.Inf(1)}, 3)
|
||
|
|
if _, _, err := Histogram(withInf, 2); err == nil {
|
||
|
|
t.Error("expected an error for an infinite sample")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestHistogramViewCountsVisibleElements(t *testing.T) {
|
||
|
|
// A rebased view shares the parent's payload, which runs past the
|
||
|
|
// view's own count: the sweep must cover the visible elements only,
|
||
|
|
// exactly as the range reductions and the edges already do.
|
||
|
|
dense := mustFromFloats(t, []float64{100, 200, 1, 2, 3, 4, 500, 600}, 8)
|
||
|
|
view, err := core.Slice(dense, 0, 2, 6)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Slice: %v", err)
|
||
|
|
}
|
||
|
|
counts, edges, err := Histogram(view, 4)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Histogram: %v", err)
|
||
|
|
}
|
||
|
|
got := counts.RawInts()[:counts.Len()]
|
||
|
|
wantCounts := []int64{1, 1, 1, 1}
|
||
|
|
for i, w := range wantCounts {
|
||
|
|
if got[i] != w {
|
||
|
|
t.Errorf("counts[%d] = %d, want %d (all %v)", i, got[i], w, got)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
gotEdges := edges.RawFloats()[:edges.Len()]
|
||
|
|
wantEdges := []float64{1, 1.75, 2.5, 3.25, 4}
|
||
|
|
for i, w := range wantEdges {
|
||
|
|
if gotEdges[i] != w {
|
||
|
|
t.Errorf("edges[%d] = %g, want %g (all %v)", i, gotEdges[i], w, gotEdges)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestViewNonFiniteTailIsNotScanned(t *testing.T) {
|
||
|
|
// The tail of the parent's payload is invisible to the view, so a
|
||
|
|
// non-finite value there must not refuse an estimation on the view.
|
||
|
|
dense := mustFromFloats(t, []float64{1, 2, 3, 4, math.NaN(), math.Inf(1)}, 6)
|
||
|
|
view, err := core.Slice(dense, 0, 0, 4)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Slice: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := Median(view); err != nil {
|
||
|
|
t.Errorf("Median over a view with a non-finite tail: %v", err)
|
||
|
|
}
|
||
|
|
if _, _, err := Histogram(view, 2); err != nil {
|
||
|
|
t.Errorf("Histogram over a view with a non-finite tail: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestHistogramParallelCounting pins the split counting sweep: past the
|
||
|
|
// per-worker floor the payload is cut into chunks, each worker counts
|
||
|
|
// into a private array of one cell per bin and the merge adds every
|
||
|
|
// chunk into the total. The counts are exact integers and every sample
|
||
|
|
// carries exactly one increment, so the split must produce the very
|
||
|
|
// counts the single-worker sweep produces, bin for bin.
|
||
|
|
func TestHistogramParallelCounting(t *testing.T) {
|
||
|
|
const n, bins = 1 << 15, 256
|
||
|
|
vals := make([]float64, n)
|
||
|
|
for i := range vals {
|
||
|
|
// Repeated values over a wide range, with the sample's own
|
||
|
|
// extremes at both ends so the first and the top bin are hit.
|
||
|
|
vals[i] = float64((i*7919)%1000)/1000*4 - 2
|
||
|
|
}
|
||
|
|
vals[0], vals[1] = -2, 2
|
||
|
|
a := mustFloats(t, vals, n)
|
||
|
|
|
||
|
|
prev := engine.SetNumWorkers(4)
|
||
|
|
defer engine.SetNumWorkers(prev)
|
||
|
|
if w := engine.WorkersFor(n); w < 2 {
|
||
|
|
t.Fatalf("the sweep did not split across workers: %d", w)
|
||
|
|
}
|
||
|
|
split, splitEdges, err := Histogram(a, bins)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Histogram across workers: %v", err)
|
||
|
|
}
|
||
|
|
engine.SetNumWorkers(1)
|
||
|
|
serial, serialEdges, err := Histogram(a, bins)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Histogram on one worker: %v", err)
|
||
|
|
}
|
||
|
|
if split.Len() != bins || splitEdges.Len() != bins+1 {
|
||
|
|
t.Fatalf("lengths %d and %d, want %d and %d", split.Len(), splitEdges.Len(), bins, bins+1)
|
||
|
|
}
|
||
|
|
got, want := split.RawInts(), serial.RawInts()
|
||
|
|
total := int64(0)
|
||
|
|
for i := range bins {
|
||
|
|
if got[i] != want[i] {
|
||
|
|
t.Fatalf("count[%d] = %d across workers, %d on one, want the same", i, got[i], want[i])
|
||
|
|
}
|
||
|
|
if splitEdges.RawFloats()[i] != serialEdges.RawFloats()[i] {
|
||
|
|
t.Fatalf("edge[%d] = %g across workers, %g on one", i,
|
||
|
|
splitEdges.RawFloats()[i], serialEdges.RawFloats()[i])
|
||
|
|
}
|
||
|
|
total += got[i]
|
||
|
|
}
|
||
|
|
if total != n {
|
||
|
|
t.Fatalf("the counts total %d, want %d", total, n)
|
||
|
|
}
|
||
|
|
// The binning arithmetic, recomputed from the returned edges: a
|
||
|
|
// chunk boundary must not move a sample between bins.
|
||
|
|
lo := splitEdges.RawFloats()[0]
|
||
|
|
width := splitEdges.RawFloats()[1] - lo
|
||
|
|
ref := make([]int64, bins)
|
||
|
|
for _, v := range vals {
|
||
|
|
b := int((v - lo) / width)
|
||
|
|
if b >= bins {
|
||
|
|
b = bins - 1
|
||
|
|
}
|
||
|
|
if b < 0 {
|
||
|
|
b = 0
|
||
|
|
}
|
||
|
|
ref[b]++
|
||
|
|
}
|
||
|
|
for i := range bins {
|
||
|
|
if got[i] != ref[i] {
|
||
|
|
t.Fatalf("count[%d] = %d, the binning arithmetic gives %d", i, got[i], ref[i])
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|