// Copyright (c) 2026 Petr Balvín (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) } }