Files

966 lines
28 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
package core
import (
"math"
"strings"
"testing"
)
// Reference walks: the bisections, the odometer packing and the
// trapezoid sums the fixed-iteration searches, the count-then-fill and
// the halved products replace. They pin the exact semantics the fast
// forms must reproduce bit for bit: bin clamping with NaN to the
// outermost bin, the rightmost insertion rule, the left-segment knot
// rule, the row-major coordinate order and the exact accumulation.
func refBinFloat(ev []float64, v float64) int {
if v < ev[0] {
return 0
}
lo, hi := 0, len(ev)-1
for lo < hi {
mid := int(uint(lo+hi) >> 1)
if ev[mid] <= v {
lo = mid + 1
} else {
hi = mid
}
}
if lo == 0 {
return len(ev) - 2
}
return lo - 1
}
func refBinInt(ev []int64, v int64) int {
if v < ev[0] {
return 0
}
lo, hi := 0, len(ev)-1
for lo < hi {
mid := int(uint(lo+hi) >> 1)
if ev[mid] <= v {
lo = mid + 1
} else {
hi = mid
}
}
return lo - 1
}
func refUpperFloat(h []float64, q float64) int {
lo, hi := 0, len(h)
for lo < hi {
mid := int(uint(lo+hi) >> 1)
if h[mid] <= q {
lo = mid + 1
} else {
hi = mid
}
}
return lo
}
func refUpperInt(h []int64, q int64) int {
lo, hi := 0, len(h)
for lo < hi {
mid := int(uint(lo+hi) >> 1)
if h[mid] <= q {
lo = mid + 1
} else {
hi = mid
}
}
return lo
}
func refInterpSegment(xv []float64, q float64) int {
n := len(xv)
if q <= xv[0] {
return 0
}
if q >= xv[n-1] {
return n - 2
}
lo, hi := 1, n
for lo < hi {
mid := int(uint(lo+hi) >> 1)
if xv[mid] < q {
lo = mid + 1
} else {
hi = mid
}
}
return lo - 1
}
// refArgwhere collects coordinates the way the single odometer walk did.
func refArgwhere(a *Array) []int64 {
var rows []int64
coord := make([]int, a.NDim())
for i := range a.Len() {
if !isZero(a, i) {
for d := range a.NDim() {
rows = append(rows, int64(coord[d]))
}
}
advanceOdometer(coord, a.shape)
}
return rows
}
// refGather is the per-element accessor walk Gather replaced.
func refGather(src *Array, dim int, index *Array) (*Array, error) {
out := &Array{shape: index.Shape(), dt: src.dt}
out.alloc(index.Len())
dst := make([]int, index.NDim())
for i := range index.Len() {
idx := int(index.ints[index.physIndex(i)])
if idx < 0 || idx >= src.shape[dim] {
return nil, errf("Gather: index %d out of range for dimension %d of size %d at position %d", idx, dim, src.shape[dim], i)
}
off := 0
for d := range dst {
c := dst[d]
if d == dim {
c = idx
}
off = off*src.shape[d] + c
}
out.setFrom(i, src, off)
advanceOdometer(dst, index.shape)
}
return out, nil
}
// refTake is the validated per-element copy Take replaced.
func refTake(a *Array, indices *Array) (*Array, error) {
out := &Array{shape: []int{indices.Len()}, dt: a.dt}
out.alloc(indices.Len())
src := a
if !src.isContiguous() {
src = src.materialise()
}
k := intPayload(indices)
for i, v := range k {
if v < 0 || v >= int64(a.Len()) {
return nil, errf("Take: index %d out of range for flat size %d at position %d", v, a.Len(), i)
}
}
switch out.dt {
case Int:
for i, v := range k {
out.ints[i] = src.ints[v]
}
case Float16:
for i, v := range k {
out.halves[i] = src.halves[v]
}
case Float32:
for i, v := range k {
out.floats32[i] = src.floats32[v]
}
case Float:
for i, v := range k {
out.floats[i] = src.floats[v]
}
default:
for i, v := range k {
out.complexes[i] = src.complexes[v]
}
}
return out, nil
}
// probeSearchValues returns the adversarial probe values every search
// equivalence test walks: exact edges, the two zeros, neighbours of the
// edges, out-of-range ends, both NaN signs, infinities, subnormals and
// a deterministic random spread.
func probeSearchValues(edges []float64) []float64 {
vals := make([]float64, 0, 64)
vals = append(vals, edges...)
for _, e := range edges {
vals = append(vals,
math.Nextafter(e, math.Inf(1)),
math.Nextafter(e, math.Inf(-1)),
e+0,
e-0)
}
vals = append(vals,
0, math.Copysign(0, -1),
math.Inf(1), math.Inf(-1),
math.NaN(), math.Copysign(math.NaN(), -1),
5e-324, -5e-324, 1e-320, -1e-320,
1e300, -1e300)
g := NewGenerator(3)
rnd, err := Floats(g, 40)
if err != nil {
panic(err)
}
vals = append(vals, rnd.RawFloats()[:40]...)
return vals
}
// TestBinSearchMatchesReference pins AssignBins' float bin selection
// against the bisection walk on every edge class: uniform edges, a
// single bin, repeated edges, negative ranges, infinite outer edges,
// subnormal edges and both large and small magnitudes. NaN edges are out
// of scope: an ascending edge set holds none.
func TestBinSearchMatchesReference(t *testing.T) {
edgeSets := [][]float64{
{0, 1},
{0, 0.25, 0.5, 0.75, 1},
{1, 2, 3},
{0, 0, 0, 1, 2, 2},
{-5, -3.5, -1, 0, 0, 2},
{math.Inf(-1), -1, 0, 1, math.Inf(1)},
{0, 5e-324, 1e-320, 1},
{1e300, 2e300, 3e300},
linspaceVals(0, 1, 256),
linspaceVals(-3, 7, 257),
}
for _, ev := range edgeSets {
vals := probeSearchValues(ev)
edges := mustFloats(t, ev)
a := mustFloats(t, vals)
out, err := AssignBins(a, edges)
if err != nil {
t.Fatalf("AssignBins m=%d: %v", len(ev), err)
}
got := out.RawInts()[:out.Len()]
for i, v := range vals {
want := refBinFloat(ev, v)
if got[i] != int64(want) {
t.Fatalf("edges %v…%v (m=%d), v=%v: bin %d, want %d",
ev[0], ev[len(ev)-1], len(ev), v, got[i], want)
}
}
}
}
// TestBinSearchMatchesReferenceInt pins the int bin selection, native
// int64 comparisons included: the extremes, the neighbours above 2^53
// that a float64 detour would fold together, and repeated edges.
func TestBinSearchMatchesReferenceInt(t *testing.T) {
edgeSets := [][]int64{
{0, 1},
{1, 2, 3},
{math.MinInt64, -1, 0, 1, math.MaxInt64},
{1 << 53, 1<<53 + 1, 1<<53 + 2},
{5, 5, 7, 9, 9},
{math.MinInt64, math.MaxInt64},
}
values := []int64{
0, 1, 2, 3, 5, 7, 9,
-1, 4, 6, 8, 10,
math.MinInt64, math.MaxInt64, math.MinInt64 + 1, math.MaxInt64 - 1,
1 << 53, 1<<53 + 1, 1<<53 + 2, 1<<53 + 3,
-(1 << 53), -(1<<53 + 1),
}
for _, ev := range edgeSets {
edges, err := FromInts(ev, len(ev))
if err != nil {
t.Fatal(err)
}
a, err := FromInts(values, len(values))
if err != nil {
t.Fatal(err)
}
out, err := AssignBins(a, edges)
if err != nil {
t.Fatalf("AssignBins int m=%d: %v", len(ev), err)
}
got := out.RawInts()[:out.Len()]
for i, v := range values {
want := refBinInt(ev, v)
if got[i] != int64(want) {
t.Fatalf("int edges %v, v=%d: bin %d, want %d", ev, v, got[i], want)
}
}
}
}
// TestAssignBinsPinsDocumentedSemantics fixes the observable edge rules
// with literals: below-range clamps to bin 0, at-or-above the last edge
// clamps to the outermost bin, an exact edge lands in the bin it opens,
// and a NaN keeps the outermost bin whatever its sign.
func TestAssignBinsPinsDocumentedSemantics(t *testing.T) {
edges := mustFloats(t, []float64{0, 1.0 / 3, 2.0 / 3, 1})
vals := []float64{
-1, 0, 1.0 / 3, 0.5, 2.0 / 3, 0.999, 1, 2,
math.NaN(), math.Copysign(math.NaN(), -1),
math.Inf(1), math.Inf(-1), math.Copysign(0, -1),
}
a := mustFloats(t, vals)
out, err := AssignBins(a, edges)
if err != nil {
t.Fatalf("AssignBins: %v", err)
}
want := []int64{0, 0, 1, 1, 2, 2, 2, 2, 2, 2, 2, 0, 0}
for i, w := range want {
if got := out.RawInts()[i]; got != w {
t.Fatalf("value %v (bits %#x): bin %d, want %d", vals[i], math.Float64bits(vals[i]), got, w)
}
}
}
// TestSearchSortedMatchesReference pins the rightmost insertion rule
// against the bisection for both operand kinds through the public entry
// point, including the int64 range above 2^53, NaN needles, infinite
// needles and infinite haystack ends, ties and the two zeros.
func TestSearchSortedMatchesReference(t *testing.T) {
haySetsF := [][]float64{
{1},
{0, 0},
{1, 2},
{-1, 0, 0, 1, 4, 4, 4, 9},
{math.Inf(-1), -3, 0, 3, math.Inf(1)},
{math.Copysign(0, -1), 0, 0, 1},
linspaceVals(0, 1, 33),
}
needleSetsF := [][]float64{
{0, 0.5, 1, 4, 9, 10},
{-1, math.Copysign(0, -1), 0, 3},
{math.NaN(), math.Copysign(math.NaN(), -1), math.Inf(1), math.Inf(-1)},
probeSearchValues([]float64{-1, 0, 1, 4}),
}
for _, h := range haySetsF {
for _, qs := range needleSetsF {
hay := mustFloats(t, h)
nd := mustFloats(t, qs)
out, err := SearchSorted(hay, nd)
if err != nil {
t.Fatalf("SearchSorted: %v", err)
}
got := out.RawInts()[:out.Len()]
for i, q := range qs {
want := refUpperFloat(h, q)
if got[i] != int64(want) {
t.Fatalf("haystack %v, needle %v: position %d, want %d", h, q, got[i], want)
}
}
}
}
haySetsI := [][]int64{
{1},
{0, 0, 1},
{1 << 53, 1<<53 + 1, 1<<53 + 1, 1 << 54},
{math.MinInt64, -1, 0, 1, math.MaxInt64},
}
needleValsI := []int64{
0, 1, 2, 1 << 53, 1<<53 + 1, 1<<53 + 2, 1 << 54,
math.MinInt64, math.MaxInt64, -1,
}
for _, h := range haySetsI {
hay, err := FromInts(h, len(h))
if err != nil {
t.Fatal(err)
}
nd, err := FromInts(needleValsI, len(needleValsI))
if err != nil {
t.Fatal(err)
}
out, err := SearchSorted(hay, nd)
if err != nil {
t.Fatalf("SearchSorted int: %v", err)
}
got := out.RawInts()[:out.Len()]
for i, q := range needleValsI {
want := refUpperInt(h, q)
if got[i] != int64(want) {
t.Fatalf("int haystack %v, needle %d: position %d, want %d", h, q, got[i], want)
}
}
}
// An int haystack above 2^53 must search natively, and an empty
// haystack reports position zero everywhere.
hi, err := FromInts([]int64{1 << 53, 1<<53 + 1, 1 << 54}, 3)
if err != nil {
t.Fatal(err)
}
ni, err := FromInts([]int64{1 << 53, 1<<53 + 1, 1<<53 + 2}, 3)
if err != nil {
t.Fatal(err)
}
out, err := SearchSorted(hi, ni)
if err != nil {
t.Fatalf("SearchSorted: %v", err)
}
for i, w := range []int64{1, 2, 2} {
if got := out.RawInts()[i]; got != w {
t.Fatalf("int needle %d: position %d, want %d", ni.RawInts()[i], got, w)
}
}
empty, err := FromInts(nil, 0)
if err != nil {
t.Fatal(err)
}
out, err = SearchSorted(empty, ni)
if err != nil {
t.Fatalf("SearchSorted empty: %v", err)
}
for i := range 3 {
if got := out.RawInts()[i]; got != 0 {
t.Fatalf("empty haystack position %d = %d, want 0", i, got)
}
}
}
// TestInterpolateSegmentMatchesReference pins the whole per-point
// pipeline against the bisection-derived reference across knot classes:
// strict ascent, repeated knots at the edges and in the interior, the
// minimal two-knot set, and queries on knots, between them, outside the
// range and at both infinities. The comparison is the result value's
// exact bits, which is the observable contract.
func TestInterpolateSegmentMatchesReference(t *testing.T) {
knotSets := [][]float64{
{0, 1},
{0, 1, 2, 3},
{0, 0.5, 0.5, 0.5, 1},
{0, 0, 1, 2},
{0, 1, 2, 2},
{-2, -1, -1, 0},
{0, 1e-300, 2e-300, 3e-300},
linspaceVals(0, 1, 33),
}
for _, xv := range knotSets {
ys := make([]float64, len(xv))
for i := range ys {
ys[i] = float64(i*i-3*i+1) / 7
}
queries := probeSearchValues([]float64{xv[0], xv[len(xv)/2], xv[len(xv)-1]})
qv := make([]float64, 0, len(queries))
for _, q := range queries {
if q != q {
continue // NaN queries are refused before the search
}
qv = append(qv, q)
}
xArr := mustFloats(t, xv)
yArr := mustFloats(t, ys)
qArr := mustFloats(t, qv)
out, err := Interpolate(xArr, yArr, qArr)
if err != nil {
t.Fatalf("Interpolate knots %v…%v: %v", xv[0], xv[len(xv)-1], err)
}
for i, q := range qv {
lo := refInterpSegment(xv, q)
x0, x1 := xv[lo], xv[lo+1]
y0, y1 := ys[lo], ys[lo+1]
tt := 0.0
if x1 > x0 {
tt = (q - x0) / (x1 - x0)
} else if q > x0 {
tt = 1
}
if tt < 0 {
tt = 0
}
if tt > 1 {
tt = 1
}
want := y0 + tt*(y1-y0)
got := out.RawFloats()[i]
if math.Float64bits(got) != math.Float64bits(want) {
t.Fatalf("knots %v…%v (n=%d), q=%v: value %#x, want %#x",
xv[0], xv[len(xv)-1], len(xv), q, math.Float64bits(got), math.Float64bits(want))
}
}
}
}
// TestInterpolateMatchesReferencePointwise pins the whole per-point
// pipeline: same segment, same t expression, same bits, for random
// queries against repeated-knot and strictly ascending knot sets.
func TestInterpolateMatchesReferencePointwise(t *testing.T) {
for _, knots := range [][]float64{
{0, 0.5, 0.5, 1, 2, 2, 3},
linspaceVals(-1, 2, 17),
} {
ys := make([]float64, len(knots))
for i := range ys {
ys[i] = float64(i*i-3*i+1) / 7
}
g := NewGenerator(9)
rnd, err := Floats(g, 200)
if err != nil {
t.Fatal(err)
}
queries := make([]float64, 0, 220)
queries = append(queries, rnd.RawFloats()[:200]...)
for _, k := range knots {
queries = append(queries, k)
}
queries = append(queries, -2, 5, math.Inf(1), math.Inf(-1), math.Copysign(0, -1))
xv := mustFloats(t, knots)
yv := mustFloats(t, ys)
qv := mustFloats(t, queries)
out, err := Interpolate(xv, yv, qv)
if err != nil {
t.Fatalf("Interpolate: %v", err)
}
n := len(knots)
for i, q := range queries {
lo := refInterpSegment(knots, q)
x0, x1 := knots[lo], knots[lo+1]
y0, y1 := ys[lo], ys[lo+1]
tt := 0.0
if x1 > x0 {
tt = (q - x0) / (x1 - x0)
} else if q > x0 {
tt = 1
}
if tt < 0 {
tt = 0
}
if tt > 1 {
tt = 1
}
want := y0 + tt*(y1-y0)
got := out.RawFloats()[i]
if math.Float64bits(got) != math.Float64bits(want) {
t.Fatalf("knots n=%d, query %v: value %#x, want %#x",
n, q, math.Float64bits(got), math.Float64bits(want))
}
}
}
}
// TestInterpolateNaNQueryReportsSmallestIndex pins the error contract
// under the parallel walk: the first NaN in order names the error.
func TestInterpolateNaNQueryReportsSmallestIndex(t *testing.T) {
xs := mustFloats(t, []float64{0, 1, 2})
ys := mustFloats(t, []float64{0, 1, 2})
q := mustFloats(t, []float64{0.5, 0.5, math.NaN(), 0.5, math.NaN(), 0.5})
if _, err := Interpolate(xs, ys, q); err == nil || !strings.Contains(err.Error(), "query 2 is NaN") {
t.Fatalf("Interpolate NaN error = %v, want the query 2 report", err)
}
}
// TestArgwhereMatchesReference pins the coordinate packing of the
// count-then-fill against the single odometer walk, across dtypes,
// ranks, densities and a chunk-boundary size; -0.0 counts as zero and
// NaN counts as non-zero, as the value tests have always had it.
func TestArgwhereMatchesReference(t *testing.T) {
shapes := [][]int{
{1}, {7}, {5, 4}, {3, 5, 4}, {2, 2, 2, 2}, {80, 80}, {1, 1, 5},
}
for _, shape := range shapes {
n := 1
for _, d := range shape {
n *= d
}
for _, density := range []int{0, 1, 2, 3, 64} {
// density: 0 all-zero, 1 all-non-zero, k every k-th non-zero.
valsF := make([]float64, n)
valsI := make([]int64, n)
for i := range n {
if density == 1 || (density > 1 && i%density == 0) {
valsF[i] = float64(i+1) / 3
valsI[i] = int64(i) + 1
}
}
if density > 0 && n > 3 {
valsF[2] = math.NaN()
valsF[3] = math.Copysign(0, -1) // zero
valsF[4] = math.Inf(1) // non-zero
}
af, _ := FromFloats(valsF, shape...)
ai, _ := FromInts(valsI, shape...)
for _, a := range []*Array{af, ai} {
want := refArgwhere(a)
out, err := Argwhere(a)
if err != nil {
t.Fatalf("Argwhere %v density %d: %v", shape, density, err)
}
got := out.RawInts()[:out.Len()]
if len(got) != len(want) {
t.Fatalf("Argwhere %v density %d: %d coordinates, want %d", shape, density, len(got), len(want))
}
for i := range want {
if got[i] != want[i] {
t.Fatalf("Argwhere %v density %d: coordinate %d = %d, want %d", shape, density, i, got[i], want[i])
}
}
if out.Shape()[0]*out.Shape()[1] != len(want) || out.Shape()[1] != a.NDim() {
t.Fatalf("Argwhere %v: shape %v does not pack %d coordinates of rank %d",
shape, out.Shape(), len(want), a.NDim())
}
}
}
}
// A rebased view: the walk is bounded by the view's own extent, and
// the coordinates are the view's.
base := mustFloats(t, []float64{0, 5, 0, 0, 7, 9, 0, 1, 0, 0, 0, 2})
view, err := Slice(base, 0, 3, 11)
if err != nil {
t.Fatal(err)
}
v2, err := Reshape(view, 2, 4)
if err != nil {
t.Fatal(err)
}
want := refArgwhere(v2)
out, err := Argwhere(v2)
if err != nil {
t.Fatalf("Argwhere view: %v", err)
}
got := out.RawInts()[:out.Len()]
if len(got) != len(want) {
t.Fatalf("Argwhere view: %d coordinates, want %d", len(got), len(want))
}
for i := range want {
if got[i] != want[i] {
t.Fatalf("Argwhere view: coordinate %d = %d, want %d", i, got[i], want[i])
}
}
}
// TestGatherTakeMatchReference pins the parallel payload walks against
// the per-element accessor walks, values and error reports alike, over
// every dtype, several ranks, repeated and out-of-range indices, view
// indices and view sources.
func TestGatherTakeMatchReference(t *testing.T) {
mk := func(vals []float64, shape ...int) *Array {
a, err := FromFloats(vals, shape...)
if err != nil {
t.Fatal(err)
}
return a
}
srcF := mk([]float64{
1, 2, 3, 4,
5, 6, 7, 8,
9, 10, 11, 12,
}, 3, 4)
srcI, _ := FromInts([]int64{
1, 2, 3, 4,
5, 6, 7, 8,
9, 10, 11, 12,
}, 3, 4)
srcC, _ := FromComplexes([]complex128{1, 2i, 3, 4i, 5, 6i}, 3, 2)
src3d := mk([]float64{
1, 2, 3, 4, 5, 6, 7, 8,
9, 10, 11, 12, 13, 14, 15, 16,
}, 2, 2, 4)
idxSets := []*Array{
mustInts(t, []int64{3, 0, 2, 2, 1, 3}, 3, 2), // for dim 1: (3, 2)
mustInts(t, []int64{2, 0, 2, 1}, 2, 2), // for dim 0: (2, 4) reshaped below
mustInts(t, []int64{0, 3, 2, 1, 1, 3, 0, 2}, 2, 4),
}
for _, src := range []*Array{srcF, srcI} {
for dim := range 2 {
for _, idx := range idxSets {
if !gatherCompatible(src.shape, idx.shape, dim) {
continue
}
want, wantErr := refGather(src, dim, idx)
got, gotErr := Gather(src, dim, idx)
if (wantErr == nil) != (gotErr == nil) || (wantErr != nil && wantErr.Error() != gotErr.Error()) {
t.Fatalf("Gather dim %d: error %v, want %v", dim, gotErr, wantErr)
}
if wantErr == nil && !Equal(got, want) {
t.Fatalf("Gather dim %d idx %v: %v, want %v", dim, idx.RawInts(), got, want)
}
}
}
}
// Complex source gathers unchanged; the index keeps the source's
// non-dim extent, as gatherCompatible requires.
idxC := mustInts(t, []int64{2, 0, 1, 1}, 2, 2)
wantC, _ := refGather(srcC, 0, idxC)
gotC, err := Gather(srcC, 0, idxC)
if err != nil || !Equal(gotC, wantC) {
t.Fatalf("Gather complex: %v vs %v (%v)", gotC, wantC, err)
}
// 3-D source, gather along the interior dimension: the index keeps
// the outer and inner extents and varies along dim 1.
idx3 := mustInts(t, []int64{1, 0, 1, 1, 0, 1, 0, 0}, 2, 1, 4)
want3, _ := refGather(src3d, 1, idx3)
got3, err := Gather(src3d, 1, idx3)
if err != nil || !Equal(got3, want3) {
t.Fatalf("Gather 3-D dim 1: %v vs %v (%v)", got3, want3, err)
}
// Out-of-range and negative indices: same error, same position.
bad := mustInts(t, []int64{0, 4, 2, 1, 0, 0, 0, 0}, 2, 4)
_, wantErr := refGather(srcF, 0, bad)
_, gotErr := Gather(srcF, 0, bad)
if wantErr == nil || gotErr == nil || wantErr.Error() != gotErr.Error() {
t.Fatalf("Gather range error: %v, want %v", gotErr, wantErr)
}
neg := mustInts(t, []int64{1, -1, 0, 0, 0, 0, 0, 0}, 2, 4)
_, wantErr = refGather(srcF, 0, neg)
_, gotErr = Gather(srcF, 0, neg)
if wantErr == nil || gotErr == nil || wantErr.Error() != gotErr.Error() {
t.Fatalf("Gather negative error: %v, want %v", gotErr, wantErr)
}
// A rebased view on the index side: the walk is bounded by the
// view's own extent, and the invisible payload tail never reads.
idxBase := mustInts(t, []int64{0, 1, 2, 0, 2, 1, 0, 2, 99, 99, 99, 99}, 3, 4)
idxView, err := Slice(idxBase, 0, 1, 2)
if err != nil {
t.Fatal(err)
}
wantV, _ := refGather(srcF, 0, idxView)
gotV, err := Gather(srcF, 0, idxView)
if err != nil || !Equal(gotV, wantV) {
t.Fatalf("Gather view index: %v vs %v (%v)", gotV, wantV, err)
}
// Take: every dtype, view source, view indices, range errors.
for _, src := range []*Array{srcF, srcI, srcC} {
tk := mustInts(t, []int64{int64(src.Len()) - 1, 0, 3, 3, 1}, 5)
wantT, wantErr := refTake(src, tk)
gotT, gotErr := Take(src, tk)
if (wantErr == nil) != (gotErr == nil) || (wantErr != nil && wantErr.Error() != gotErr.Error()) {
t.Fatalf("Take: error %v, want %v", gotErr, wantErr)
}
if wantErr == nil && !Equal(gotT, wantT) {
t.Fatalf("Take: %v, want %v", gotT, wantT)
}
}
// Take on a rebased view source: flat indices count from the view's
// origin, and the walk is bounded by the view's own extent.
longBase := mustFloats(t, []float64{99, 99, 99, 99, 10, 11, 12, 13, 14, 15, 16, 17})
viewSrc, err := Slice(longBase, 0, 4, 12)
if err != nil {
t.Fatal(err)
}
tk := mustInts(t, []int64{7, 0, 3}, 3)
wantT, _ := refTake(viewSrc, tk)
gotT, err := Take(viewSrc, tk)
if err != nil || !Equal(gotT, wantT) {
t.Fatalf("Take view source: %v vs %v (%v)", gotT, wantT, err)
}
badT := mustInts(t, []int64{1, 2, 99}, 3)
_, wantErr = refTake(srcF, badT)
_, gotErr = Take(srcF, badT)
if wantErr == nil || gotErr == nil || wantErr.Error() != gotErr.Error() {
t.Fatalf("Take range error: %v, want %v", gotErr, wantErr)
}
}
// TestIntegrateHalvingIsBitIdentical proves the multiplication by 0.5
// reproduces the division by 2 bit for bit on adversarial payloads:
// subnormals, the signed zeros, infinities, NaN payloads and random
// spreads, for a range of spacings. The accumulation order is untouched
// either way.
func TestIntegrateHalvingIsBitIdentical(t *testing.T) {
payloads := [][]float64{
{1, 2, 3, 4},
{5e-324, 1e-320, 2.5e-320, 1e-308},
{math.Copysign(0, -1), 0, 1, math.Copysign(0, -1)},
{math.Inf(1), 1, math.Inf(-1), 2},
{math.NaN(), 1, 2, math.Copysign(math.NaN(), -1)},
{1e308, 1.5e308, 2, 3},
{-1e308, 1e308, 1, -1},
}
g := NewGenerator(21)
rnd, err := Floats(g, 128)
if err != nil {
t.Fatal(err)
}
payloads = append(payloads, rnd.RawFloats()[:128])
for _, dx := range []float64{1, 0.5, 2, 1e-300, 1e300, -3, math.NaN(), math.Inf(1)} {
for _, y := range payloads {
yArr := mustFloats(t, y)
gotTotal, err := Integrate(yArr, dx)
if err != nil {
t.Fatalf("Integrate: %v", err)
}
var want float64
for i := 1; i < len(y); i++ {
want += (y[i-1] + y[i]) / 2
}
want *= dx
if math.Float64bits(gotTotal) != math.Float64bits(want) {
t.Fatalf("Integrate dx=%v payload %v…: %#x, want %#x",
dx, y[0], math.Float64bits(gotTotal), math.Float64bits(want))
}
gotCum, err := CumulativeIntegrate(yArr, dx)
if err != nil {
t.Fatalf("CumulativeIntegrate: %v", err)
}
ov := make([]float64, len(y))
for i := 1; i < len(y); i++ {
ov[i] = ov[i-1] + (y[i-1]+y[i])/2*dx
}
for i := range y {
if math.Float64bits(gotCum.RawFloats()[i]) != math.Float64bits(ov[i]) {
t.Fatalf("CumulativeIntegrate dx=%v at %d: %#x, want %#x",
dx, i, math.Float64bits(gotCum.RawFloats()[i]), math.Float64bits(ov[i]))
}
}
}
}
}
// TestArgwhereVariantsMatchReference pins every packing variant the
// probe walks, production included, against the single odometer walk:
// the coordinates, their order and the packed shape must match whatever
// the buffers do on the way.
func TestArgwhereVariantsMatchReference(t *testing.T) {
variants := append(argwhereStyles(),
struct {
name string
run func(*Array) *Array
}{"merge-capped-1024", func(a *Array) *Array { return probeArgwhereMerge(a, 1024) }},
)
shapes := [][]int{{1}, {7}, {5, 4}, {3, 5, 4}, {80, 80}, {2, 2, 2, 2}}
for _, shape := range shapes {
n := 1
for _, d := range shape {
n *= d
}
for _, density := range []int{0, 1, 2, 64} {
valsF := make([]float64, n)
for i := range n {
if density == 1 || (density > 1 && i%density == 0) {
valsF[i] = float64(i+1) / 3
}
}
if density > 0 && n > 4 {
valsF[2] = math.NaN()
valsF[3] = math.Copysign(0, -1)
valsF[4] = math.Inf(1)
}
a, _ := FromFloats(valsF, shape...)
want := refArgwhere(a)
for _, v := range variants {
out := v.run(a)
if out == nil {
t.Fatalf("%s %v density %d: nil result", v.name, shape, density)
}
got := out.RawInts()[:out.Len()]
if len(got) != len(want) {
t.Fatalf("%s %v density %d: %d coordinates, want %d", v.name, shape, density, len(got), len(want))
}
for i := range want {
if got[i] != want[i] {
t.Fatalf("%s %v density %d: coordinate %d = %d, want %d", v.name, shape, density, i, got[i], want[i])
}
}
}
}
}
// A rebased view: bounded by the view's own extent for every variant.
base := mustFloats(t, []float64{0, 5, 0, 0, 7, 9, 0, 1, 0, 0, 0, 2})
view, err := Slice(base, 0, 3, 11)
if err != nil {
t.Fatal(err)
}
v2, err := Reshape(view, 2, 4)
if err != nil {
t.Fatal(err)
}
want := refArgwhere(v2)
for _, v := range variants {
out := v.run(v2)
got := out.RawInts()[:out.Len()]
if len(got) != len(want) {
t.Fatalf("%s view: %d coordinates, want %d", v.name, len(got), len(want))
}
for i := range want {
if got[i] != want[i] {
t.Fatalf("%s view: coordinate %d = %d, want %d", v.name, i, got[i], want[i])
}
}
}
}
// TestSortRadixConfigsBitIdentical pins every probe radix configuration
// against the production Sort and ArgSort on adversarial fixtures: the
// sorted values bit for bit (NaN payloads included) and the
// permutations exactly, whatever crew cap and digit width carry them.
func TestSortRadixConfigsBitIdentical(t *testing.T) {
mkInts := func(vals []int64) *Array {
a, err := FromInts(vals, len(vals))
if err != nil {
t.Fatal(err)
}
return a
}
intRand := benchIntsShape(1000, 5000)
intWide := benchIntsShape(1<<30, 5000)
intParallel := benchIntsShape(1000, 20000)
floatRand := benchFloatsShape(5000)
floatParallel := benchFloatsShape(20000)
fixtures := []*Array{
intRand,
intWide,
intParallel,
mkInts([]int64{math.MinInt64, math.MaxInt64, 0, -1, 1, 1 << 53, 1<<53 + 1, math.MinInt64 + 1, math.MaxInt64 - 1, 5, 5, 5}),
mkInts([]int64{42, 42, 42, 42, 42}),
floatRand,
floatParallel,
mustFloats(t, []float64{math.NaN(), math.Copysign(math.NaN(), -1), math.Copysign(0, -1), 0, math.Inf(1), math.Inf(-1), 1, 1, 2.5, 2.5, -3, 0.1}),
mustFloats(t, []float64{7, 7, 7, 7, 7}),
}
cfgs := probeSortConfigs()
for _, a := range fixtures {
wantSort, err := Sort(a)
if err != nil {
t.Fatal(err)
}
wantIdx, err := ArgSort(a)
if err != nil {
t.Fatal(err)
}
for _, c := range cfgs {
var gotSort *Array
if a.dt == Int {
gotSort = probeSortInt(a, c.wcap, c.bits)
} else {
gotSort = probeSortFloat(a, c.wcap, c.bits)
}
if gotSort.Len() != wantSort.Len() {
t.Fatalf("%v %s: sorted length %d, want %d", a.Shape(), c.name, gotSort.Len(), wantSort.Len())
}
if a.dt == Int {
g, w := gotSort.RawInts()[:gotSort.Len()], wantSort.RawInts()[:wantSort.Len()]
for i := range w {
if g[i] != w[i] {
t.Fatalf("int fixture %v %s: sorted[%d] = %d, want %d", a.Shape(), c.name, i, g[i], w[i])
}
}
} else {
g, w := gotSort.RawFloats()[:gotSort.Len()], wantSort.RawFloats()[:wantSort.Len()]
for i := range w {
if math.Float64bits(g[i]) != math.Float64bits(w[i]) {
t.Fatalf("float fixture %v %s: sorted[%d] = %#x, want %#x", a.Shape(), c.name, i, math.Float64bits(g[i]), math.Float64bits(w[i]))
}
}
}
gotIdx := probeArgSort(a, c.wcap, c.bits)
g, w := gotIdx.RawInts()[:gotIdx.Len()], wantIdx.RawInts()[:wantIdx.Len()]
for i := range w {
if g[i] != w[i] {
t.Fatalf("fixture %v %s: argsort[%d] = %d, want %d", a.Shape(), c.name, i, g[i], w[i])
}
}
}
}
}
// linspaceVals builds n evenly spaced values, the deterministic edge
// ladder the bin equivalence tests walk.
func linspaceVals(start, stop float64, n int) []float64 {
a, err := Linspace(start, stop, n)
if err != nil {
panic(err)
}
return append([]float64(nil), a.RawFloats()[:n]...)
}
func mustInts(t *testing.T, vals []int64, shape ...int) *Array {
t.Helper()
a, err := FromInts(vals, shape...)
if err != nil {
t.Fatal(err)
}
return a
}