Files
tensor/stats/stats_random_test.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

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)
}
}