Files
tensor/internal/core/random_test.go
T

150 lines
3.6 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"math"
"strings"
"testing"
)
func TestGeneratorDeterminism(t *testing.T) {
g1 := NewGenerator(42)
g2 := NewGenerator(42)
f1, err := Floats(g1, 8)
if err != nil {
t.Fatalf("Floats: %v", err)
}
f2, err := Floats(g2, 8)
if err != nil {
t.Fatalf("Floats: %v", err)
}
if !Equal(f1, f2) {
t.Fatalf("the same seed must give identical arrays")
}
g3 := NewGenerator(43)
f3, _ := Floats(g3, 8)
if Equal(f1, f3) {
t.Fatalf("different seeds must give different arrays")
}
}
func TestGeneratorFloats(t *testing.T) {
g := NewGenerator(7)
f, err := Floats(g, 1000)
if err != nil {
t.Fatalf("Floats: %v", err)
}
if f.Dtype() != Float || f.Shape()[0] != 1000 {
t.Fatalf("Floats shape: %s", f)
}
for i := range f.Len() {
v, _ := FloatAt(f, i)
if v < 0 || v >= 1 {
t.Fatalf("Floats out of [0,1): %v", v)
}
}
// A thousand draws should cover both halves of the interval.
m, _ := Mean(f)
if m < 0.4 || m > 0.6 {
t.Fatalf("Floats mean suggests bias: %v", m)
}
if _, err := Floats(g, -1); err == nil || !strings.Contains(err.Error(), "zero or greater") {
t.Fatalf("Floats negative: %v", err)
}
}
func TestGeneratorInts(t *testing.T) {
g := NewGenerator(9)
v, err := Ints(g, 600, 1, 4)
if err != nil {
t.Fatalf("Ints: %v", err)
}
if v.Dtype() != Int {
t.Fatalf("Ints dtype: %s", v.Dtype())
}
seen := map[int64]bool{}
for i := range v.Len() {
x, _ := IntAt(v, i)
if x < 1 || x >= 4 {
t.Fatalf("Ints out of [1,4): %d", x)
}
seen[x] = true
}
// Six hundred draws over three values must hit all of them.
if len(seen) != 3 {
t.Fatalf("Ints coverage: %v", seen)
}
if _, err := Ints(g, 1, 5, 5); err == nil || !strings.Contains(err.Error(), "min must be less than max") {
t.Fatalf("Ints range: %v", err)
}
if _, err := Ints(g, -1, 0, 1); err == nil || !strings.Contains(err.Error(), "zero or greater") {
t.Fatalf("Ints negative: %v", err)
}
}
// TestGeneratorBoundedUniform pins the two defects of the old rejection
// rule: it rejected the whole [t, n) residue band, starving the upper
// buckets (for spans above 2^63 some values never appeared), and a full
// [MinInt64, MaxInt64) span accepted only lo == 2^64-1, hanging the draw.
func TestGeneratorBoundedUniform(t *testing.T) {
g := NewGenerator(11)
const draws = 30_000
counts := map[int64]int{}
for range draws {
v, err := Ints(g, 1, 0, 3)
if err != nil {
t.Fatalf("Ints: %v", err)
}
x, _ := IntAt(v, 0)
counts[x]++
}
// The counts are binomial around draws/3; the old biased rule kept
// bucket 2 a full draws/3 * (1/3) short, far outside this band.
want := draws / 3
tol := want / 10
for b := range int64(3) {
if diff := counts[b] - want; diff > tol || diff < -tol {
t.Fatalf("bucket %d drawn %d times, want %d +/- %d", b, counts[b], want, tol)
}
}
// Full-span draws must terminate promptly and spread over the range.
g = NewGenerator(12)
seen := map[int64]bool{}
for range 64 {
v, err := Ints(g, 1, math.MinInt64, math.MaxInt64)
if err != nil {
t.Fatalf("full-span Ints: %v", err)
}
x, _ := IntAt(v, 0)
seen[x] = true
}
if len(seen) < 2 {
t.Fatalf("full-span draws collapsed to %d distinct value(s)", len(seen))
}
}
func TestGeneratorZeroSeed(t *testing.T) {
// Seed zero must not degenerate into the all-zero state.
g := NewGenerator(0)
f, err := Floats(g, 4)
if err != nil {
t.Fatalf("Floats from seed 0: %v", err)
}
nonzero := false
for i := range f.Len() {
v, _ := FloatAt(f, i)
if v != 0 {
nonzero = true
}
}
if !nonzero {
t.Fatalf("seed 0 produced only zeros")
}
}