75 lines
2.1 KiB
Go
75 lines
2.1 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package stats
|
||
|
||
import (
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
"testing"
|
||
)
|
||
|
||
// TestPoissonDrawsZeroLambda pins the degenerate distribution: λ = 0
|
||
// must draw 0, never the −1 the multiplication loop used to hand back.
|
||
func TestPoissonDrawsZeroLambda(t *testing.T) {
|
||
g := core.NewGenerator(7)
|
||
out, err := PoissonDraws(g, 16, 0)
|
||
if err != nil {
|
||
t.Fatalf("PoissonDraws: %v", err)
|
||
}
|
||
for i := range 16 {
|
||
if out.FloatAt(i) != 0 {
|
||
t.Fatalf("draw %d = %g, want 0", i, out.FloatAt(i))
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestChiSquareComplexRejects pins the dtype guard.
|
||
func TestChiSquareComplexRejects(t *testing.T) {
|
||
obs, err := core.FromComplexes([]complex128{1, 2, 3}, 3)
|
||
if err != nil {
|
||
t.Fatalf("FromComplexes: %v", err)
|
||
}
|
||
exp, _ := core.FromFloats([]float64{1, 2, 3}, 3)
|
||
if _, _, _, err := ChiSquareGoodnessOfFit(obs, exp); err == nil {
|
||
t.Fatal("ChiSquareGoodnessOfFit accepted complex observed")
|
||
}
|
||
if _, _, _, err := ChiSquareGoodnessOfFit(exp, obs); err == nil {
|
||
t.Fatal("ChiSquareGoodnessOfFit accepted complex expected")
|
||
}
|
||
}
|
||
|
||
// TestBootstrapCIFreshArray pins the documented contract: a statistic
|
||
// that retains its argument must observe values frozen at call time,
|
||
// never the next resample's mutation through the shared buffer.
|
||
func TestBootstrapCIFreshArray(t *testing.T) {
|
||
data, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8}, 8)
|
||
if err != nil {
|
||
t.Fatalf("FromFloats: %v", err)
|
||
}
|
||
type snapshot struct {
|
||
arr *core.Array
|
||
mean float64
|
||
}
|
||
var kept []snapshot
|
||
stat := func(a *core.Array) (float64, error) {
|
||
m, err := core.Mean(a)
|
||
if err != nil {
|
||
return 0, err
|
||
}
|
||
kept = append(kept, snapshot{a, m})
|
||
return m, nil
|
||
}
|
||
if _, _, err := BootstrapCI(data, stat, 0.9, 8, 42); err != nil {
|
||
t.Fatalf("BootstrapCI: %v", err)
|
||
}
|
||
for i, s := range kept {
|
||
m, err := core.Mean(s.arr)
|
||
if err != nil {
|
||
t.Fatalf("kept[%d] mean: %v", i, err)
|
||
}
|
||
if m != s.mean {
|
||
t.Fatalf("kept[%d] mean drifted from %g to %g (buffer was mutated)", i, s.mean, m)
|
||
}
|
||
}
|
||
}
|