142 lines
3.8 KiB
Go
142 lines
3.8 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package stats
|
|
|
|
import (
|
|
"math"
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|
"strings"
|
|
"testing"
|
|
)
|
|
|
|
func TestVar(t *testing.T) {
|
|
// The classic set: mean 5, population variance 4.
|
|
a := mustFromInts(t, []int64{2, 4, 4, 4, 5, 5, 7, 9}, 8)
|
|
v, err := Var(a)
|
|
if err != nil {
|
|
t.Fatalf("Var: %v", err)
|
|
}
|
|
if math.Abs(v-4) > 1e-9 {
|
|
t.Fatalf("Var: %v", v)
|
|
}
|
|
|
|
// Sample variance of the same set: 32/7.
|
|
vs, err := VarSample(a)
|
|
if err != nil {
|
|
t.Fatalf("VarSample: %v", err)
|
|
}
|
|
if math.Abs(vs-32.0/7.0) > 1e-9 {
|
|
t.Fatalf("VarSample: %v", vs)
|
|
}
|
|
|
|
// Std squared is Var.
|
|
std, _ := Std(a)
|
|
if math.Abs(std*std-v) > 1e-12 {
|
|
t.Fatalf("Std² vs Var: %v %v", std*std, v)
|
|
}
|
|
|
|
one := mustFromFloats(t, []float64{1.5}, 1)
|
|
if _, err := VarSample(one); err == nil || !strings.Contains(err.Error(), "more than 1") {
|
|
t.Fatalf("VarSample single: %v", err)
|
|
}
|
|
empty := mustFromInts(t, nil, 0)
|
|
if _, err := Var(empty); err == nil || !strings.Contains(err.Error(), "more than 0") {
|
|
t.Fatalf("Var empty: %v", err)
|
|
}
|
|
c := mustFromComplexes(t, []complex128{1}, 1)
|
|
if _, err := Var(c); err == nil || !strings.Contains(err.Error(), "no float variance") {
|
|
t.Fatalf("Var complex: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestGeneratorNormal(t *testing.T) {
|
|
g := core.NewGenerator(11)
|
|
draws, err := core.Normal(g, 2000, 10, 2)
|
|
if err != nil {
|
|
t.Fatalf("Normal: %v", err)
|
|
}
|
|
if draws.Dtype() != core.Float || draws.Shape()[0] != 2000 {
|
|
t.Fatalf("Normal shape: %s", draws)
|
|
}
|
|
mean, _ := core.Mean(draws)
|
|
std, _ := Std(draws)
|
|
if mean < 9.8 || mean > 10.2 {
|
|
t.Fatalf("Normal mean: %v", mean)
|
|
}
|
|
if std < 1.8 || std > 2.2 {
|
|
t.Fatalf("Normal std: %v", std)
|
|
}
|
|
|
|
// Determinism: the same seed gives bit-identical draws.
|
|
n1, _ := core.Normal(core.NewGenerator(5), 16, 0, 1)
|
|
n2, _ := core.Normal(core.NewGenerator(5), 16, 0, 1)
|
|
if !core.Equal(n1, n2) {
|
|
t.Fatalf("Normal must be reproducible")
|
|
}
|
|
|
|
if _, err := core.Normal(g, 1, 0, -1); err == nil || !strings.Contains(err.Error(), "zero or greater") {
|
|
t.Fatalf("Normal negative std: %v", err)
|
|
}
|
|
if _, err := core.Normal(g, -1, 0, 1); err == nil || !strings.Contains(err.Error(), "zero or greater") {
|
|
t.Fatalf("Normal negative n: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestGeneratorPermutation(t *testing.T) {
|
|
g := core.NewGenerator(13)
|
|
p, err := core.Permutation(g, 50)
|
|
if err != nil {
|
|
t.Fatalf("Permutation: %v", err)
|
|
}
|
|
seen := make([]bool, 50)
|
|
for i := range 50 {
|
|
v, _ := core.IntAt(p, i)
|
|
if v < 0 || v >= 50 || seen[v] {
|
|
t.Fatalf("Permutation not a permutation at %d: %d", i, v)
|
|
}
|
|
seen[v] = true
|
|
}
|
|
|
|
// Determinism and difference between seeds.
|
|
p2, _ := core.Permutation(core.NewGenerator(13), 50)
|
|
if !core.Equal(p, p2) {
|
|
t.Fatalf("Permutation must be reproducible")
|
|
}
|
|
p3, _ := core.Permutation(core.NewGenerator(14), 50)
|
|
if core.Equal(p, p3) {
|
|
t.Fatalf("different seeds must give different permutations")
|
|
}
|
|
|
|
if _, err := core.Permutation(g, -1); err == nil || !strings.Contains(err.Error(), "zero or greater") {
|
|
t.Fatalf("Permutation negative: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestGeneratorShuffle(t *testing.T) {
|
|
g := core.NewGenerator(17)
|
|
a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6, 7, 8}, 8)
|
|
|
|
sh := core.Shuffle(g, a)
|
|
if sh.Len() != 8 || sh.Dtype() != core.Int {
|
|
t.Fatalf("Shuffle shape: %s", sh)
|
|
}
|
|
// A shuffle is a permutation of the same multiset.
|
|
sorted, _ := core.Sort(sh)
|
|
want, _ := core.Sort(a)
|
|
if !core.Equal(want, sorted) {
|
|
t.Fatalf("Shuffle is not a permutation: %s", sh)
|
|
}
|
|
// The receiver is untouched.
|
|
if v, _ := core.IntAt(a, 0); v != 1 {
|
|
t.Fatalf("Shuffle mutated the receiver: %d", v)
|
|
}
|
|
|
|
// core.Complex arrays shuffle too.
|
|
c := mustFromComplexes(t, []complex128{1, complex(2, 2)}, 2)
|
|
cs := core.Shuffle(g, c)
|
|
if cs.Dtype() != core.Complex || cs.Len() != 2 {
|
|
t.Fatalf("Shuffle complex: %s", cs)
|
|
}
|
|
}
|