181 lines
5.6 KiB
Go
181 lines
5.6 KiB
Go
// 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)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|