Files
tensor/internal/core/walks_pin_test.go
T

181 lines
5.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"
"slices"
"strings"
"testing"
)
// Pins for the parallel payload walks: strided operands must be read in
// logical order through the accessors, Take must report the smallest
// offending position, and the radix ArgSort permutation must equal the
// stable comparator permutation on tie-heavy inputs.
func pinFloats(t *testing.T, vals []float64, shape ...int) *Array {
t.Helper()
a, err := FromFloats(vals, shape...)
if err != nil {
t.Fatal(err)
}
return a
}
func pinInts(t *testing.T, vals []int64, shape ...int) *Array {
t.Helper()
a := New(Int, shape...)
copy(a.RawInts(), vals)
return a
}
// pinStridedFloat builds a test-mechanics strided array: logical (r, c)
// reads payload[r*rowStride + c]. The payload carries an invisible tail
// element at an unaddressed slot, the state a raw payload walk gets
// wrong.
func pinStridedFloat(t *testing.T, payload []float64, shape []int, strides []int) *Array {
t.Helper()
return &Array{shape: shape, dt: Float, floats: payload, strides: strides}
}
func pinStridedInt(t *testing.T, payload []int64, shape []int, strides []int) *Array {
t.Helper()
return &Array{shape: shape, dt: Int, ints: payload, strides: strides}
}
func TestArgwhereNonzeroStridedUseLogicalElements(t *testing.T) {
// Logical window [[1, 0], [0, 5]] at strides [3, 1]: the elements
// sit at payload slots 0, 1, 3, 4, and slot 2 holds an invisible
// 99. A payload walk at the logical position reads slot 3 for the
// bottom-right element and calls the 5 a zero.
stridedF := pinStridedFloat(t, []float64{1, 0, 99, 0, 5, 7, 7, 7}, []int{2, 2}, []int{3, 1})
twinF := pinFloats(t, []float64{1, 0, 0, 5}, 2, 2)
gotF, err := Argwhere(stridedF)
if err != nil {
t.Fatalf("Argwhere(strided float): %v", err)
}
wantF, err := Argwhere(twinF)
if err != nil {
t.Fatalf("Argwhere(twin float): %v", err)
}
if !Equal(gotF, wantF) {
t.Fatalf("Argwhere(strided float) = %v, want %v", gotF, wantF)
}
nzS, err := Nonzero(stridedF)
if err != nil {
t.Fatalf("Nonzero(strided float): %v", err)
}
nzT, err := Nonzero(twinF)
if err != nil {
t.Fatalf("Nonzero(twin float): %v", err)
}
if !slices.Equal(nzS[0], nzT[0]) || !slices.Equal(nzS[1], nzT[1]) {
t.Fatalf("Nonzero(strided float) = %v, want %v", nzS, nzT)
}
stridedI := pinStridedInt(t, []int64{1, 0, 99, 0, 5, 7, 7, 7}, []int{2, 2}, []int{3, 1})
twinI := pinInts(t, []int64{1, 0, 0, 5}, 2, 2)
gotI, err := Argwhere(stridedI)
if err != nil {
t.Fatalf("Argwhere(strided int): %v", err)
}
wantI, err := Argwhere(twinI)
if err != nil {
t.Fatalf("Argwhere(twin int): %v", err)
}
if !Equal(gotI, wantI) {
t.Fatalf("Argwhere(strided int) = %v, want %v", gotI, wantI)
}
nzSI, err := Nonzero(stridedI)
if err != nil {
t.Fatalf("Nonzero(strided int): %v", err)
}
nzTI, err := Nonzero(twinI)
if err != nil {
t.Fatalf("Nonzero(twin int): %v", err)
}
if !slices.Equal(nzSI[0], nzTI[0]) || !slices.Equal(nzSI[1], nzTI[1]) {
t.Fatalf("Nonzero(strided int) = %v, want %v", nzSI, nzTI)
}
}
func TestTakeReportsSmallestOffender(t *testing.T) {
src := pinFloats(t, []float64{10, 11, 12, 13, 14}, 5)
idx := pinInts(t, []int64{99, 88, 0}, 3)
_, err := Take(src, idx)
if err == nil {
t.Fatal("Take accepted out-of-range indices")
}
if !strings.Contains(err.Error(), "position 0") || strings.Contains(err.Error(), "position 1") {
t.Fatalf("Take offender report = %q, want the smallest offending position 0", err.Error())
}
}
func TestArgSortTiePermutationsMatchStableReference(t *testing.T) {
for _, n := range []int{64, 1024, 4096} {
vals := make([]float64, n)
ivals := make([]int64, n)
state := uint64(0x9E3779B97F4A7C15 + uint64(n))
for i := range vals {
state = state*6364136223846793005 + 1442695040888963407
// A small level set, so ties dominate the permutation.
level := int64(state>>60) % 5
ivals[i] = level - 2
vals[i] = float64(level - 2)
}
ref := make([]int, n)
for i := range ref {
ref[i] = i
}
// The stable comparator permutation: ties keep index order, the
// contract a stable digit scatter must reproduce.
slices.SortStableFunc(ref, func(x, y int) int {
switch {
case vals[x] < vals[y]:
return -1
case vals[x] > vals[y]:
return 1
}
return 0
})
gotF, err := ArgSort(pinFloats(t, vals, n))
if err != nil {
t.Fatalf("n=%d ArgSort(float): %v", n, err)
}
gotI, err := ArgSort(pinInts(t, ivals, n))
if err != nil {
t.Fatalf("n=%d ArgSort(int): %v", n, err)
}
// ArgSort answers the permutation as an int array whatever the
// sorted dtype was.
rf, ri := gotF.RawInts(), gotI.RawInts()
for i := range ref {
if int(rf[i]) != ref[i] {
t.Fatalf("n=%d float permutation at %d = %d, want %d (values %v)", n, i, int(rf[i]), ref[i], vals)
}
if int(ri[i]) != ref[i] {
t.Fatalf("n=%d int permutation at %d = %d, want %d (values %v)", n, i, int(ri[i]), ref[i], ivals)
}
}
}
// The NaN and signed-zero contract on the same walk: NaN sorts
// last, -0 folds onto +0 in the value order but keeps its index
// order among the zeros.
withSpecial := pinFloats(t, []float64{2, math.NaN(), math.Copysign(0, -1), -1, 0, math.NaN(), 3}, 7)
got, err := ArgSort(withSpecial)
if err != nil {
t.Fatalf("ArgSort(special): %v", err)
}
perm := got.RawInts()
// Sorted values: -1, then the zeros at indices 2 and 4 in index
// order, then 2, 3, then the NaNs at indices 1 and 5 in index order.
want := []int{3, 2, 4, 0, 6, 1, 5}
for i := range want {
if int(perm[i]) != want[i] {
t.Fatalf("special permutation at %d = %d, want %d (full %v)", i, int(perm[i]), want[i], perm)
}
}
}