493 lines
16 KiB
Go
493 lines
16 KiB
Go
// 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)
|
|
}
|
|
})
|
|
}
|