Files

493 lines
16 KiB
Go
Raw Permalink 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
// Regression pins for input validation in internal/core: one test per
// defect, named after what it pins; the radix chunk split also carries
// a deterministic invariant test.
package core
import (
"fmt"
"math"
"slices"
"strings"
"sync"
"testing"
"time"
"sourcedock.dev/petrbalvin/tensor/internal/engine"
)
// callNoPanic runs fn on the test goroutine and reports a panic as a
// test failure: a validation gap must surface as an error the caller can
// handle, never as a crash inside the library.
func callNoPanic(t *testing.T, what string, fn func() error) error {
t.Helper()
var err error
func() {
defer func() {
if r := recover(); r != nil {
t.Fatalf("%s panicked instead of returning an error: %v", what, r)
}
}()
err = fn()
}()
return err
}
// Pad used to accept negative pad values, drive the new shape
// negative and die in alloc with "makeslice: len out of range". Pad
// documents an error for a malformed pad argument, so the negatives have
// to be refused by name.
func TestPadRejectsNegativePadValues(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3}, 3)
err := callNoPanic(t, "Pad(-2, -2)", func() error {
_, err := Pad(a, []int{-2, -2}, "constant", 0)
return err
})
if err == nil {
t.Fatal("Pad accepted the negative pad pair (-2, -2)")
}
if !strings.Contains(err.Error(), "non-negative") {
t.Fatalf("Pad error does not name the negative pads: %v", err)
}
// One negative side of a pair is just as malformed.
err = callNoPanic(t, "Pad(1, -1)", func() error {
_, err := Pad(a, []int{1, -1}, "constant", 0)
return err
})
if err == nil {
t.Fatal("Pad accepted the mixed pad pair (1, -1)")
}
// The valid path is untouched.
ok, err := Pad(a, []int{1, 1}, "constant", 0)
if err != nil {
t.Fatalf("Pad valid pair: %v", err)
}
if !slices.Equal(ok.Shape(), []int{5}) {
t.Fatalf("Pad valid pair shape %v, want [5]", ok.Shape())
}
}
// TruncatedNormal with std = NaN used to pass the `std <= 0`
// guard, leave the rejection window NaN and spin in the retry loop
// forever. The draw must fall back to the degenerate all-zero result
// promptly, on this goroutine: the watchdog keeps a regression from
// hanging the suite.
func TestTruncatedNormalNaNStdReturnsPromptly(t *testing.T) {
done := make(chan *Array, 1)
go func() {
done <- TruncatedNormal(NewGenerator(1), []int{3}, 0, math.NaN())
}()
select {
case got := <-done:
if got == nil {
t.Fatal("TruncatedNormal(std=NaN) returned nil")
}
for i, v := range got.RawFloat32s() {
if v != 0 {
t.Fatalf("TruncatedNormal(std=NaN) value %d = %v, want the degenerate 0", i, v)
}
}
case <-time.After(5 * time.Second):
t.Fatal("TruncatedNormal(std=NaN) still running after 5s: the rejection loop cannot exit")
}
// The documented degenerate path (std <= 0) still returns zeros, and
// a valid std still draws inside the window.
zero := TruncatedNormal(NewGenerator(2), []int{4}, 0, 0)
if zero == nil {
t.Fatal("TruncatedNormal(std=0) returned nil")
}
for i, v := range zero.RawFloat32s() {
if v != 0 {
t.Fatalf("TruncatedNormal(std=0) value %d = %v, want 0", i, v)
}
}
drawn := TruncatedNormal(NewGenerator(3), []int{64}, 0, 1)
if drawn == nil {
t.Fatal("TruncatedNormal(std=1) returned nil")
}
for i, v := range drawn.RawFloat32s() {
if v < -2 || v > 2 {
t.Fatalf("TruncatedNormal(std=1) value %d = %v outside the +-2 sigma window", i, v)
}
}
}
// Normal let a NaN std through (`std < 0` is false for NaN) and
// returned an array of NaNs silently. This is the feeder of the TruncatedNormal case above, so it
// has to be a loud error.
func TestNormalRejectsNaNStd(t *testing.T) {
g := NewGenerator(1)
if _, err := Normal(g, 3, 0, math.NaN()); err == nil {
t.Fatal("Normal accepted a NaN std and drew silently wrong values")
}
if _, err := Normal(g, 3, 0, -1); err == nil {
t.Fatal("Normal accepted a negative std")
}
arr, err := Normal(g, 3, 0, 1)
if err != nil {
t.Fatalf("Normal std=1: %v", err)
}
for i, v := range arr.RawFloats() {
if math.IsNaN(v) {
t.Fatalf("Normal std=1 value %d is NaN", i)
}
}
}
// InterpolateGrid with a NaN query used to pass both clamps,
// convert to the platform's indefinite integer and panic inside
// FloatAt; with a stride above 1 it silently returned a value read from
// a wrapped index instead. The documented clamp must hold or the call
// must be refused. +Inf and -Inf keep clamping, as documented.
func TestInterpolateGridRejectsNaNQuery(t *testing.T) {
grid := mustFloats(t, []float64{0, 10, 20, 30}, 4)
origins := []float64{0}
steps := []float64{1}
nan := mustFloats(t, []float64{math.NaN()}, 1, 1)
err := callNoPanic(t, "InterpolateGrid(NaN)", func() error {
_, err := InterpolateGrid(grid, origins, steps, nan)
return err
})
if err == nil {
t.Fatal("InterpolateGrid accepted a NaN query")
}
if !strings.Contains(err.Error(), "NaN") {
t.Fatalf("InterpolateGrid error does not name the NaN position: %v", err)
}
// The infinite queries keep their clamp: -Inf to the first sample,
// +Inf to the last.
inf := mustFloats(t, []float64{math.Inf(-1), math.Inf(1)}, 2, 1)
out, err := InterpolateGrid(grid, origins, steps, inf)
if err != nil {
t.Fatalf("InterpolateGrid with infinite queries: %v", err)
}
if got := out.FloatAt(0); got != 0 {
t.Fatalf("-Inf query = %v, want the first sample 0", got)
}
if got := out.FloatAt(1); got != 30 {
t.Fatalf("+Inf query = %v, want the last sample 30", got)
}
}
// HaltonPoints had no skip + n bound, so an index near MaxInt
// wrapped the int arithmetic and every point collapsed to the origin
// with no error. The constructor now enforces the same 2^32 index
// budget SobolPoints does.
func TestHaltonPointsRejectsSkipOverflow(t *testing.T) {
if _, err := HaltonPoints(2, 2, math.MaxInt); err == nil {
t.Fatal("HaltonPoints accepted skip = MaxInt: the points collapse to the origin silently")
}
// The boundary the twin constructor enforces: n + skip must stay
// below 2^32, so the last acceptable index is 2^32 - 1.
if _, err := HaltonPoints(2, 2, 1<<32-2); err == nil {
t.Fatal("HaltonPoints accepted n + skip = 2^32")
}
if _, err := HaltonPoints(2, 2, 1<<32-3); err != nil {
t.Fatalf("HaltonPoints rejected n + skip = 2^32 - 1: %v", err)
}
if _, err := HaltonPoints(1<<32, 2, 0); err == nil {
t.Fatal("HaltonPoints accepted n = 2^32")
}
// The ordinary path is untouched.
if _, err := HaltonPoints(4, 2, 3); err != nil {
t.Fatalf("HaltonPoints valid skip: %v", err)
}
}
// Unsqueeze(-1) appends the new axis at the end (the convention
// the code has always followed and the frozen API keeps); the doc
// comment claimed it inserts before the last dimension. The behaviour
// is pinned here so the comment and the code cannot drift apart again
// in either direction.
func TestUnsqueezeNegativeDimAppendsAxis(t *testing.T) {
a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
got, err := Unsqueeze(a, -1)
if err != nil {
t.Fatalf("Unsqueeze(-1): %v", err)
}
if !slices.Equal(got.Shape(), []int{2, 3, 1}) {
t.Fatalf("Unsqueeze((2,3), -1) shape %v, want the appended [2 3 1]", got.Shape())
}
// Counting from the end of the result rank: -2 lands one axis earlier.
got, err = Unsqueeze(a, -2)
if err != nil {
t.Fatalf("Unsqueeze(-2): %v", err)
}
if !slices.Equal(got.Shape(), []int{2, 1, 3}) {
t.Fatalf("Unsqueeze((2,3), -2) shape %v, want [2 1 3]", got.Shape())
}
// One step past the rank is refused.
if _, err := Unsqueeze(a, -4); err == nil {
t.Fatal("Unsqueeze(-4) accepted for a rank-2 input")
}
}
// seekCoord was dead code with no caller in the module. Removing
// it must not move a single element, so BroadcastTo is pinned against
// the naive per-element odometer reference the run-filled fast path
// replaced.
func TestBroadcastToMatchesOdometerReference(t *testing.T) {
cases := []struct {
name string
vals []int64
srcSh []int
target []int
}{
{"column", []int64{1, 2, 3}, []int{3, 1}, []int{3, 2}},
{"prepend", []int64{1, 2, 3}, []int{3, 1}, []int{2, 3, 1}},
{"scalar", []int64{7}, []int{1}, []int{2, 3, 4}},
{"bias", []int64{1, 2, 3, 4, 5, 6}, []int{1, 3, 2}, []int{4, 3, 2}},
{"same-shape", []int64{9, 8, 7, 6}, []int{2, 2}, []int{2, 2}},
{"trailing-run", []int64{1, 2}, []int{2, 1}, []int{2, 3}},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
src := mustFromInts(t, tc.vals, tc.srcSh...)
got, err := BroadcastTo(src, tc.target...)
if err != nil {
t.Fatalf("BroadcastTo: %v", err)
}
want := broadcastOdometerReference(t, src, tc.target)
if !slices.Equal(want, got.RawInts()) {
t.Fatalf("BroadcastTo diverged from the odometer reference:\n got %v\nwant %v",
got.RawInts(), want)
}
})
}
}
// broadcastOdometerReference expands src to target one element at a
// time, coordinates recomputed from scratch: the slow path BroadcastTo's
// run fill replaced.
func broadcastOdometerReference(t *testing.T, src *Array, target []int) []int64 {
t.Helper()
total := 1
for _, d := range target {
total *= d
}
sh := src.Shape()
out := make([]int64, total)
coord := make([]int, len(target))
srcCoord := make([]int, len(sh))
off := len(target) - len(sh)
for i := range total {
for d := range sh {
c := 0
if sh[d] != 1 {
c = coord[off+d]
}
srcCoord[d] = c
}
v, err := IntAt(src, srcCoord...)
if err != nil {
t.Fatalf("IntAt(%v): %v", srcCoord, err)
}
out[i] = v
advanceOdometer(coord, target)
}
return out
}
// Radix chunk split, invariant half: the parallel radix derives its histogram rows
// from its own split of [0, n) instead of re-reading the global worker
// count the way engine.ParallelMin did, so a concurrent SetNumCPU can
// move work between goroutines but never two live chunks onto one row.
// This test reproduces the disagreement deterministically and with no
// reliance on scheduling: the stale-snapshot case raises the live worker
// count while the chunk width handed to the split stays the one a caller
// would have snapshotted earlier, which is exactly the window SetNumCPU
// opens (the raising read is what used to collapse two live chunks onto
// one row). Every row must be handed out once with the range it owns,
// start = row*chunk, and the ranges must partition [0, n); the remaining
// cases pin that contract either side of the radixParMin spawn floor.
func TestRadixChunksRowMatchesOwnedRange(t *testing.T) {
cases := []struct {
name string
n int
workers int
chunk int // the caller's snapshot width, stale where noted
}{
{"exact-fit", 40_000, 2, 20_000},
{"ragged", 40_001, 3, 13_334},
{"one-chunk", 40_000, 1, 40_000},
{"floor-exact", 40_000, 1, radixParMin},
{"stale-snapshot", 40_000, 8, 20_000}, // width from a 2-worker snapshot
{"below-floor", 40_000, 8, radixParMin - 1},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
prev := engine.SetNumWorkers(tc.workers)
defer engine.SetNumWorkers(prev)
parallel := tc.chunk >= radixParMin && tc.chunk < tc.n
nchunks := 1
if parallel {
nchunks = (tc.n + tc.chunk - 1) / tc.chunk
}
// The callback may run on a spawned goroutine, so failures
// are collected and reported on the test goroutine.
var mu sync.Mutex
var problems []string
touched := make([]int, tc.n)
rows := make([]bool, nchunks)
forEachRadixChunk(tc.n, tc.chunk, func(row, start, end int) {
var local []string
if row < 0 || row >= nchunks {
local = append(local, fmt.Sprintf("row %d out of range for %d chunks", row, nchunks))
} else if rows[row] {
local = append(local, fmt.Sprintf("row %d handed out twice", row))
} else {
rows[row] = true
}
// The row identity is the chunk position: the range a
// worker owns and the histogram row it writes are the
// same object, so they cannot disagree.
wantStart, wantEnd := 0, tc.n
if parallel {
wantStart, wantEnd = row*tc.chunk, min(row*tc.chunk+tc.chunk, tc.n)
}
if start != wantStart || end != wantEnd {
local = append(local, fmt.Sprintf("row %d owns [%d, %d), want [%d, %d)",
row, start, end, wantStart, wantEnd))
if start < 0 || end > tc.n || start >= end {
local = append(local, fmt.Sprintf("row %d owns the illegal range [%d, %d)", row, start, end))
}
}
if start >= 0 && end <= tc.n && start < end {
for i := start; i < end; i++ {
touched[i]++
}
}
mu.Lock()
problems = append(problems, local...)
mu.Unlock()
})
if len(problems) > 0 {
t.Fatalf("chunk split broken: %s", strings.Join(problems, "; "))
}
for i, c := range touched {
if c != 1 {
t.Fatalf("index %d visited %d times", i, c)
}
}
for row, seen := range rows {
if !seen {
t.Fatalf("row %d never ran", row)
}
}
})
}
}
// flipWorkers runs fn at least rounds times while a second goroutine
// flips the global worker count between low and high. SetNumCPU is
// documented as safe to call at any time, so a kernel that derives its
// chunking from the global count must still finish correctly; before
// Radix chunk split, growth inside the kernel's own window collapsed two live chunks
// onto one histogram row and corrupted the output.
func flipWorkers(t *testing.T, low, high, rounds int, fn func(round int)) {
t.Helper()
prev := engine.SetNumWorkers(low)
defer engine.SetNumWorkers(prev)
stop := make(chan struct{})
var flipper sync.WaitGroup
flipper.Go(func() {
w := low
for {
select {
case <-stop:
return
default:
}
if w == low {
w = high
} else {
w = low
}
engine.SetNumWorkers(w)
}
})
defer func() {
close(stop)
flipper.Wait()
}()
for round := range rounds {
fn(round)
}
}
// Radix chunk split, value-radix half: Sort of a large float64 payload must stay an
// ascending permutation while a concurrent SetNumCPU moves the worker
// count from 2 to 8 and back. The output is checked against the sorted
// order itself, so a lost histogram update (rows colliding), a
// mis-ordered scatter (bit-identity broken across chunks) or an
// out-of-range write all fail here.
func TestSortRadixStableUnderConcurrentSetNumCPU(t *testing.T) {
const n = 200_000
g := NewGenerator(7)
src, err := Floats(g, n)
if err != nil {
t.Fatalf("Floats: %v", err)
}
template := slices.Clone(src.RawFloats())
// Ties force the scatter to be exercised across chunk boundaries.
for i := range template {
if i%3 == 0 {
template[i] = float64(i % 17)
}
}
reference := sortReference(template)
flipWorkers(t, 2, 8, 8, func(round int) {
a := mustFromFloats(t, slices.Clone(template), n)
got, err := Sort(a)
if err != nil {
t.Fatalf("round %d: Sort: %v", round, err)
}
vals := got.RawFloats()
if !equalSortedFloats(reference, vals) {
t.Fatalf("round %d: Sort corrupted the parallel radix output", round)
}
})
}
// Radix chunk split, permutation-radix half: ArgSort must keep returning the same
// stable permutation the serial radix returns while the worker count
// changes concurrently. A row collision shows up either as a repeated
// or missing index or as a change of the permutation itself.
func TestArgSortRadixStableUnderConcurrentSetNumCPU(t *testing.T) {
const n = 200_000
g := NewGenerator(11)
src, err := Floats(g, n)
if err != nil {
t.Fatalf("Floats: %v", err)
}
vals := slices.Clone(src.RawFloats())
for i := range vals {
if i%3 == 0 {
vals[i] = float64(i % 17)
}
}
reference := argSortReference(vals)
flipWorkers(t, 2, 8, 8, func(round int) {
a := mustFromFloats(t, slices.Clone(vals), n)
got, err := ArgSort(a)
if err != nil {
t.Fatalf("round %d: ArgSort: %v", round, err)
}
if !slices.Equal(reference, got.RawInts()) {
t.Fatalf("round %d: ArgSort diverged from the stable serial permutation", round)
}
})
}