823 lines
25 KiB
Go
823 lines
25 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package core
|
||
|
||
import (
|
||
"math"
|
||
"math/rand/v2"
|
||
"strings"
|
||
"testing"
|
||
"time"
|
||
)
|
||
|
||
// Guard pins: each test guards a repair. An invalid
|
||
// einsum label is an error instead of a silent contraction, the pInf
|
||
// norm reports NaN for an all-NaN line, circular padding folds in
|
||
// constant time, BesselKn owns its domain, the float32 sparse matmul
|
||
// accumulates in float64, a NaN interpolation query is an
|
||
// error, strided views read their own elements everywhere, unknown
|
||
// dtypes are refused at the constructors, Scalar.String prints one
|
||
// sign, and Mean survives a sum that would overflow. The TestPin
|
||
// guards at the end cover paths that already agreed with their naive
|
||
// references and only needed the coverage pinned.
|
||
|
||
// ---------- Einsum rejects labels outside a-z and A-Z ----------
|
||
|
||
func TestEinsumInvalidLabels(t *testing.T) {
|
||
m := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2)
|
||
v := mustFloats(t, []float64{7, 8}, 2)
|
||
// Three specs the old label reader silently computed: the pound
|
||
// sign decoded as a label and returned a copy, the euro sign did
|
||
// the same on a vector, and the digit is refused for consistency
|
||
// with the general engine's parser.
|
||
for _, c := range []struct {
|
||
spec string
|
||
op *Array
|
||
bad string
|
||
}{
|
||
{"£->£", m, "£"},
|
||
{"€->€", v, "€"},
|
||
{"i0,i0->i0", m, "0"},
|
||
} {
|
||
out, err := Einsum(c.spec, c.op)
|
||
if err == nil {
|
||
t.Errorf("Einsum(%q) = %v, want an error", c.spec, out.Shape())
|
||
continue
|
||
}
|
||
if !strings.Contains(err.Error(), c.bad) {
|
||
t.Errorf("Einsum(%q) error %q does not name the offending character %q", c.spec, err, c.bad)
|
||
}
|
||
}
|
||
// The valid surface is unchanged: the table path still answers the
|
||
// same results bit for bit, uppercase letters included.
|
||
a := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2)
|
||
b := mustFloats(t, []float64{5, 6, 7, 8}, 2, 2)
|
||
table, err := Einsum("ij,jk->ik", a, b)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
direct, derr := MatMul2D(a, b)
|
||
if derr != nil {
|
||
t.Fatal(derr)
|
||
}
|
||
if !samePayload(t, table, direct) {
|
||
t.Error("Einsum ij,jk->ik no longer matches MatMul2D bit for bit")
|
||
}
|
||
general, err := Einsum("ij,jk", a, b)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if !samePayload(t, table, general) {
|
||
t.Error("Einsum table and general paths disagree")
|
||
}
|
||
upA := mustFloats(t, []float64{1, 2, 3, 4}, 2, 2)
|
||
upB := mustFloats(t, []float64{5, 6, 7, 8}, 2, 2)
|
||
up, err := Einsum("AB,BC->AC", upA, upB)
|
||
if err != nil {
|
||
t.Fatalf("uppercase labels rejected: %v", err)
|
||
}
|
||
if !samePayload(t, up, table) {
|
||
t.Error("uppercase labels answered a different product")
|
||
}
|
||
x := mustFromInts(t, []int64{1, 2}, 2)
|
||
y := mustFromInts(t, []int64{3, 4}, 2)
|
||
dot, err := Einsum("i,i->", x, y)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if dot.Dtype() != Int || dot.ints[0] != 11 {
|
||
t.Errorf("Einsum i,i-> = %s %v, want int [11]", dot.Dtype(), dot.ints)
|
||
}
|
||
}
|
||
|
||
// samePayload reports whether two arrays hold identical values, dtype
|
||
// included, comparing int payloads as integers so the check stays bit
|
||
// exact for every element type.
|
||
func samePayload(t *testing.T, a, b *Array) bool {
|
||
t.Helper()
|
||
if a.Dtype() != b.Dtype() || !sameShape(a.Shape(), b.Shape()) {
|
||
return false
|
||
}
|
||
for i := range a.Len() {
|
||
switch a.Dtype() {
|
||
case Int:
|
||
if a.ints[i] != b.ints[i] {
|
||
return false
|
||
}
|
||
case Float32:
|
||
if math.Float32bits(a.floats32[i]) != math.Float32bits(b.floats32[i]) {
|
||
return false
|
||
}
|
||
case Float:
|
||
if math.Float64bits(a.floats[i]) != math.Float64bits(b.floats[i]) {
|
||
return false
|
||
}
|
||
default:
|
||
if a.complexes[i] != b.complexes[i] {
|
||
return false
|
||
}
|
||
}
|
||
}
|
||
return true
|
||
}
|
||
|
||
// ---------- Norm with p = +Inf propagates NaN ----------
|
||
|
||
func TestNormInfNaN(t *testing.T) {
|
||
allNaN := mustFloats(t, []float64{math.NaN(), math.NaN()}, 2)
|
||
out, err := Norm(allNaN, math.Inf(1), 0, false)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if v := out.FloatAt(0); !math.IsNaN(v) {
|
||
t.Errorf("Norm(all-NaN, Inf) = %v, want NaN like p = 1 and p = 2", v)
|
||
}
|
||
// NaN between finite values is skipped, as it always was: the
|
||
// magnitudes decide.
|
||
mixed := mustFloats(t, []float64{math.NaN(), 3, -9}, 3)
|
||
mixOut, err := Norm(mixed, math.Inf(1), 0, false)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if v := mixOut.FloatAt(0); v != 9 {
|
||
t.Errorf("Norm([NaN, 3, -9], Inf) = %v, want 9", v)
|
||
}
|
||
// Lines without NaN keep their exact answers.
|
||
clean := mustFloats(t, []float64{1, -4, 2}, 3)
|
||
cleanOut, err := Norm(clean, math.Inf(1), 0, false)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if v := cleanOut.FloatAt(0); v != 4 {
|
||
t.Errorf("Norm([1, -4, 2], Inf) = %v, want 4", v)
|
||
}
|
||
// The float32 path seeds the same way.
|
||
f32NaN := &Array{shape: []int{2}, dt: Float32, floats32: []float32{float32(math.NaN()), 5}}
|
||
f32Out, err := Norm(f32NaN, math.Inf(1), 0, false)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if v := f32Out.FloatAt(0); v != 5 {
|
||
t.Errorf("Norm(float32 [NaN, 5], Inf) = %v, want 5", v)
|
||
}
|
||
f32All := &Array{shape: []int{2}, dt: Float32, floats32: []float32{float32(math.NaN()), float32(math.NaN())}}
|
||
f32AllOut, err := Norm(f32All, math.Inf(1), 0, false)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if v := f32AllOut.FloatAt(0); !math.IsNaN(v) {
|
||
t.Errorf("Norm(float32 all-NaN, Inf) = %v, want NaN", v)
|
||
}
|
||
// Per-line behaviour on a 2-D array along both orientations.
|
||
g := mustFloats(t, []float64{math.NaN(), 2, 3, math.NaN()}, 2, 2)
|
||
byDim1, err := Norm(g, math.Inf(1), 1, false)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if byDim1.FloatAt(0) != 2 || byDim1.FloatAt(1) != 3 {
|
||
t.Errorf("Norm dim 1 = [%v, %v], want [2, 3]", byDim1.FloatAt(0), byDim1.FloatAt(1))
|
||
}
|
||
byDim0, err := Norm(g, math.Inf(1), 0, false)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if byDim0.FloatAt(0) != 3 || byDim0.FloatAt(1) != 2 {
|
||
t.Errorf("Norm dim 0 = [%v, %v], want [3, 2]", byDim0.FloatAt(0), byDim0.FloatAt(1))
|
||
}
|
||
}
|
||
|
||
// ---------- Circular padding folds in constant time ----------
|
||
|
||
func TestPadCircularConstantTime(t *testing.T) {
|
||
// Small cases agree with a naive modulo fold, both directions.
|
||
src := mustFromInts(t, []int64{1, 2, 3}, 3)
|
||
got, err := Pad(src, []int{4, 5}, "circular", 0)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
for i := range 12 {
|
||
want := src.ints[((i-4)%3+3)%3]
|
||
if got.ints[i] != want {
|
||
t.Fatalf("circular pad [%d] = %d, want %d", i, got.ints[i], want)
|
||
}
|
||
}
|
||
// A pad far longer than the axis used to fold one step at a time,
|
||
// a quadratic walk: a pre-pad of two million on a one-element axis
|
||
// must land in constants, not hours.
|
||
one := mustFloats(t, []float64{7}, 1)
|
||
start := time.Now()
|
||
big, err := Pad(one, []int{2_000_000, 0}, "circular", 0)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if elapsed := time.Since(start); elapsed > 30*time.Second {
|
||
t.Fatalf("circular pad of 2000000 on a length-1 axis took %s", elapsed)
|
||
}
|
||
if big.Len() != 2_000_001 {
|
||
t.Fatalf("padded length %d, want 2000001", big.Len())
|
||
}
|
||
for _, i := range []int{0, 1, 999_999, 2_000_000} {
|
||
if v := big.FloatAt(i); v != 7 {
|
||
t.Fatalf("circular pad [%d] = %v, want 7", i, v)
|
||
}
|
||
}
|
||
}
|
||
|
||
// ---------- BesselKn refuses every non-positive argument ----------
|
||
|
||
func TestBesselKnDomain(t *testing.T) {
|
||
neg := mustFloats(t, []float64{-1}, 1)
|
||
for _, n := range []int{0, 2} {
|
||
out, err := BesselKn(n, neg)
|
||
if err == nil {
|
||
t.Errorf("BesselKn(%d, -1) = %v, want an error", n, out.FloatAt(0))
|
||
} else if !strings.Contains(err.Error(), "the argument must be positive") {
|
||
t.Errorf("BesselKn(%d, -1) error %q does not state the domain", n, err)
|
||
} else if !strings.HasPrefix(err.Error(), "tensor: BesselKn") {
|
||
t.Errorf("BesselKn(%d, -1) error %q lacks the prefixed name", n, err)
|
||
}
|
||
}
|
||
zero := mustFloats(t, []float64{0}, 1)
|
||
if _, err := BesselKn(0, zero); err == nil {
|
||
t.Error("BesselKn(0, 0) accepted the origin")
|
||
}
|
||
nan := mustFloats(t, []float64{math.NaN()}, 1)
|
||
if _, err := BesselKn(1, nan); err == nil {
|
||
t.Error("BesselKn(1, NaN) accepted NaN")
|
||
}
|
||
// The error names the offending element in a mixed array.
|
||
mixed := mustFloats(t, []float64{1, -2}, 2)
|
||
if _, err := BesselKn(3, mixed); err == nil || !strings.Contains(err.Error(), "element 1") {
|
||
t.Errorf("BesselKn(3, [1, -2]) error %v does not name element 1", err)
|
||
}
|
||
// Positive arguments are untouched: the recurrence still holds.
|
||
pos := mustFloats(t, []float64{1, 4}, 2)
|
||
for _, n := range []int{0, 1, 4} {
|
||
out, err := BesselKn(n, pos)
|
||
if err != nil {
|
||
t.Fatalf("BesselKn(%d, [1, 4]): %v", n, err)
|
||
}
|
||
for i, x := range []float64{1, 4} {
|
||
if v := out.FloatAt(i); !(v > 0) || math.IsInf(v, 1) {
|
||
t.Errorf("BesselKn(%d, %g) = %v, want a finite positive value", n, x, v)
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
// ---------- SpMatMul accumulates float32 in float64 ----------
|
||
|
||
func TestSpMatMulFloat32Accumulation(t *testing.T) {
|
||
// One output cell fed by seven stored coordinates: the found case
|
||
// where narrowing every product before the float32 addition drifts
|
||
// ten ulps from the float64 reference, while widening once per
|
||
// coordinate stays within one.
|
||
vals := []float32{-0.72071517, 0.0010362695, 0.9238949, 0.0019417476, 205664.48, 0.8203619, 0.0019267823}
|
||
dens := []float32{-0.066897266, 0, 0.7552673, 0, 0, -0.8196033, 0}
|
||
indices := make([]int64, 0, 14)
|
||
for j := range vals {
|
||
indices = append(indices, 0, int64(j))
|
||
}
|
||
idx, err := FromInts(indices, 7, 2)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
v32, err := FromFloat32s(vals, 7)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
s, err := NewSparseCOO(idx, v32, []int{1, 7})
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
d32, err := FromFloat32s(dens, 7, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
got, err := SpMatMul(s, d32)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if got.Dtype() != Float32 {
|
||
t.Fatalf("dtype %s, want float32", got.Dtype())
|
||
}
|
||
// The widening walk: every coordinate widens, accumulates in float64
|
||
// and narrows the cell on arrival.
|
||
exp := float32(0)
|
||
for j := range vals {
|
||
exp = float32(float64(exp) + float64(vals[j])*float64(dens[j]))
|
||
}
|
||
if bits := math.Float32bits(got.floats32[0]); bits != math.Float32bits(exp) {
|
||
t.Errorf("SpMatMul float32 = %v (%08x), want %v (%08x)",
|
||
got.floats32[0], bits, exp, math.Float32bits(exp))
|
||
}
|
||
// Within two ulps of the float64 accumulation narrowed once.
|
||
acc := 0.0
|
||
for j := range vals {
|
||
acc += float64(vals[j]) * float64(dens[j])
|
||
}
|
||
ref := float32(acc)
|
||
if d := ulpDiff32(got.floats32[0], ref); d > 2 {
|
||
t.Errorf("SpMatMul float32 = %v sits %d ulps from the float64 reference %v, want at most 2",
|
||
got.floats32[0], d, ref)
|
||
}
|
||
// An int sparse times a float32 dense promotes and stays sane.
|
||
vi, err := FromInts([]int64{3, -1, 2}, 3)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
idxI, err := FromInts([]int64{0, 0, 0, 1, 0, 2}, 3, 2)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
si, err := NewSparseCOO(idxI, vi, []int{1, 3})
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
di, err := FromFloat32s([]float32{0.1, 0.2, 0.3}, 3, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
gotI, err := SpMatMul(si, di)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if gotI.Dtype() != Float32 {
|
||
t.Fatalf("mixed dtype %s, want float32", gotI.Dtype())
|
||
}
|
||
if want := float32(3*0.1 - 0.2 + 2*0.3); math.Abs(float64(gotI.floats32[0]-want)) > 1e-6 {
|
||
t.Errorf("int×float32 sparse product = %v, want about %v", gotI.floats32[0], want)
|
||
}
|
||
}
|
||
|
||
// ulpDiff32 counts the representable float32 steps between a and b.
|
||
func ulpDiff32(a, b float32) int {
|
||
return int(int64(math.Float32bits(a)) - int64(math.Float32bits(b)))
|
||
}
|
||
|
||
// ---------- InterpolateMonotone refuses a NaN query ----------
|
||
|
||
func TestPCHIPNaNQuery(t *testing.T) {
|
||
xs := mustFloats(t, []float64{0, 1, 2}, 3)
|
||
ys := mustFloats(t, []float64{0, 1, 4}, 3)
|
||
q := mustFloats(t, []float64{0.5, math.NaN()}, 2)
|
||
out, err := InterpolateMonotone(xs, ys, q)
|
||
if err == nil {
|
||
t.Errorf("InterpolateMonotone with a NaN query = %v, want an error", out.Shape())
|
||
} else if !strings.Contains(err.Error(), "NaN") {
|
||
t.Errorf("InterpolateMonotone NaN error %q does not name NaN", err)
|
||
}
|
||
fine, err := InterpolateMonotone(xs, ys, mustFloats(t, []float64{0.5, 1.9}, 2))
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if v := fine.FloatAt(0); !(v > 0 && v < 1) {
|
||
t.Errorf("InterpolateMonotone(0.5) = %v, want a value inside (0, 1)", v)
|
||
}
|
||
}
|
||
|
||
// ---------- Strided views read their own elements ----------
|
||
|
||
func TestStridedRawReads(t *testing.T) {
|
||
// Equal walks both arrays' own elements, not their payloads: two
|
||
// views carrying {4, 1} over different padding are equal, and one
|
||
// of them equals the dense array of the same elements.
|
||
v1 := &Array{shape: []int{2}, dt: Int, ints: []int64{4, 9, 1, 9}, strides: []int{2}}
|
||
v2 := &Array{shape: []int{2}, dt: Int, ints: []int64{4, 7, 1, 7}, strides: []int{2}}
|
||
dense, err := FromInts([]int64{4, 1}, 2)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if !Equal(v1, v2) {
|
||
t.Error("Equal of two identical strided int views = false, want true")
|
||
}
|
||
if !Equal(v1, dense) {
|
||
t.Error("Equal of a strided int view and its dense twin = false, want true")
|
||
}
|
||
v3 := &Array{shape: []int{2}, dt: Int, ints: []int64{4, 9, 2, 9}, strides: []int{2}}
|
||
if Equal(v1, v3) {
|
||
t.Error("Equal of views holding {4, 1} and {4, 2} = true, want false")
|
||
}
|
||
// ArgMax/ArgMin compare the view's elements as ints.
|
||
got, err := ArgMax(v1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if got != 0 {
|
||
t.Errorf("ArgMax(strided int view of [4, 1]) = %d, want 0", got)
|
||
}
|
||
gotMin, err := ArgMin(v1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if gotMin != 1 {
|
||
t.Errorf("ArgMin(strided int view of [4, 1]) = %d, want 1", gotMin)
|
||
}
|
||
// Sort gathers the view's own float32 elements.
|
||
fv := &Array{shape: []int{3}, dt: Float32, floats32: []float32{9, 7, 1, 7, 2, 7}, strides: []int{2}}
|
||
sorted, err := Sort(fv)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
wantSorted, err := FromFloat32s([]float32{1, 2, 9}, 3)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if !samePayload(t, sorted, wantSorted) {
|
||
t.Errorf("Sort(strided float32 view of [9, 1, 2]) = %v, want [1, 2, 9]", sorted.floats32)
|
||
}
|
||
}
|
||
|
||
// ---------- The constructors refuse unknown dtypes ----------
|
||
|
||
func TestUnknownDtypeRejected(t *testing.T) {
|
||
if _, err := Zeros(Dtype(99), 2); err == nil {
|
||
t.Error("Zeros(Dtype(99), 2) created a payload")
|
||
}
|
||
if _, err := Ones(Dtype(200), 2); err == nil {
|
||
t.Error("Ones(Dtype(200), 2) created a payload")
|
||
}
|
||
if New(Dtype(99), 2) != nil {
|
||
t.Error("New(Dtype(99), 2) returned an array, want nil")
|
||
}
|
||
for _, dt := range []Dtype{Int, Float32, Float, Complex} {
|
||
if _, err := Zeros(dt, 2); err != nil {
|
||
t.Errorf("Zeros(%s, 2): %v", dt, err)
|
||
}
|
||
}
|
||
}
|
||
|
||
// ---------- Scalar.String prints one sign ----------
|
||
|
||
func TestScalarString(t *testing.T) {
|
||
c := Scalar{isComplex: true, c: complex(4, -2)}
|
||
if got := c.String(); got != "complex (4-2i)" {
|
||
t.Errorf("Scalar.String() = %q, want %q", got, "complex (4-2i)")
|
||
}
|
||
pos := Scalar{isComplex: true, c: complex(0, 2.5)}
|
||
if got := pos.String(); got != "complex (0+2.5i)" {
|
||
t.Errorf("Scalar.String() = %q, want %q", got, "complex (0+2.5i)")
|
||
}
|
||
if strings.Contains(Scalar{isComplex: true, c: complex(1, -1)}.String(), "+-") {
|
||
t.Error("Scalar.String() still prints +-")
|
||
}
|
||
if got := (Scalar{i: 7}).String(); got != "int 7" {
|
||
t.Errorf("Scalar.String() = %q, want %q", got, "int 7")
|
||
}
|
||
if got := (Scalar{isFloat: true, f: 2.5}).String(); got != "float 2.5" {
|
||
t.Errorf("Scalar.String() = %q, want %q", got, "float 2.5")
|
||
}
|
||
}
|
||
|
||
// ---------- Mean survives a sum that would overflow ----------
|
||
|
||
func TestMeanOverflow(t *testing.T) {
|
||
maxs := mustFloats(t, []float64{math.MaxFloat64, math.MaxFloat64}, 2)
|
||
m, err := Mean(maxs)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if math.Float64bits(m) != math.Float64bits(math.MaxFloat64) {
|
||
t.Errorf("Mean of two MaxFloat64s = %v, want %v", m, math.MaxFloat64)
|
||
}
|
||
threes := mustFloats(t, []float64{math.MaxFloat64, math.MaxFloat64, math.MaxFloat64}, 3)
|
||
m3, err := Mean(threes)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if math.Float64bits(m3) != math.Float64bits(math.MaxFloat64) {
|
||
t.Errorf("Mean of three MaxFloat64s = %v, want %v", m3, math.MaxFloat64)
|
||
}
|
||
neg := mustFloats(t, []float64{-math.MaxFloat64, -math.MaxFloat64}, 2)
|
||
mn, err := Mean(neg)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if math.Float64bits(mn) != math.Float64bits(-math.MaxFloat64) {
|
||
t.Errorf("Mean of two −MaxFloat64s = %v, want %v", mn, -math.MaxFloat64)
|
||
}
|
||
// Ordinary inputs keep the plain Sum rounding bit for bit.
|
||
plain := mustFloats(t, []float64{0.1, 0.2, 0.3, 700, -3.25}, 5)
|
||
mp, err := Mean(plain)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if want := Sum(plain).Float() / 5; math.Float64bits(mp) != math.Float64bits(want) {
|
||
t.Errorf("Mean of ordinary input = %v, want the Sum rounding %v", mp, want)
|
||
}
|
||
// Int arrays keep their true-division mean.
|
||
mi, err := Mean(mustFromInts(t, []int64{1, 2}, 2))
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if mi != 1.5 {
|
||
t.Errorf("Mean of int [1, 2] = %v, want 1.5", mi)
|
||
}
|
||
// A NaN payload still reaches the caller as NaN, and an infinite
|
||
// element keeps the plain path's infinite mean.
|
||
if mNaN, err := Mean(mustFloats(t, []float64{math.NaN(), 1}, 2)); err != nil || !math.IsNaN(mNaN) {
|
||
t.Errorf("Mean of [NaN, 1] = %v, %v, want NaN", mNaN, err)
|
||
}
|
||
if mInf, err := Mean(mustFloats(t, []float64{math.Inf(1), 1}, 2)); err != nil || !math.IsInf(mInf, 1) {
|
||
t.Errorf("Mean of [+Inf, 1] = %v, %v, want +Inf", mInf, err)
|
||
}
|
||
}
|
||
|
||
// ---------- Coverage pins (no fix behind them) ----------
|
||
|
||
// TestPinEinsumBatchedVecMat walks the two batched
|
||
// matrix-vector patterns against a naive loop, for int and float32,
|
||
// with one batch and with several.
|
||
func TestPinEinsumBatchedVecMat(t *testing.T) {
|
||
rng := rand.New(rand.NewPCG(606, 6))
|
||
for _, batches := range []int{1, 3} {
|
||
n, k, m := 2, 3, 4
|
||
for _, dt := range []Dtype{Int, Float32} {
|
||
a := randWhole(t, rng, dt, batches, n, k) // (b, n, k)
|
||
bm := randWhole(t, rng, dt, batches, k, m) // (b, k, m)
|
||
v := randWhole(t, rng, dt, batches, k) // (b, k)
|
||
mv, err := Einsum("bij,bj->bi", a, v)
|
||
if err != nil {
|
||
t.Fatalf("bij,bj->bi (%d batches, %s): %v", batches, dt, err)
|
||
}
|
||
vm, err := Einsum("bi,bij->bj", v, bm)
|
||
if err != nil {
|
||
t.Fatalf("bi,bij->bj (%d batches, %s): %v", batches, dt, err)
|
||
}
|
||
for bt := range batches {
|
||
for i := range n {
|
||
var wantF float64
|
||
var wantI int64
|
||
for p := range k {
|
||
av := int64At(a, (bt*n+i)*k+p)
|
||
bv := int64At(v, bt*k+p)
|
||
if dt == Int {
|
||
wantI += av * bv
|
||
} else {
|
||
wantF += float64(av) * float64(bv)
|
||
}
|
||
}
|
||
if dt == Int {
|
||
if mv.ints[bt*n+i] != wantI {
|
||
t.Errorf("bij,bj->bi int batch %d row %d = %d, want %d", bt, i, mv.ints[bt*n+i], wantI)
|
||
}
|
||
} else if d := math.Abs(float64(mv.floats32[bt*n+i]) - wantF); d > 1e-5*judgingScale(wantF) {
|
||
t.Errorf("bij,bj->bi float32 batch %d row %d = %v, want %v", bt, i, mv.floats32[bt*n+i], wantF)
|
||
}
|
||
}
|
||
for j := range m {
|
||
var wantF float64
|
||
var wantI int64
|
||
for p := range k {
|
||
bv := int64At(v, bt*k+p)
|
||
mav := int64At(bm, (bt*k+p)*m+j)
|
||
if dt == Int {
|
||
wantI += bv * mav
|
||
} else {
|
||
wantF += float64(bv) * float64(mav)
|
||
}
|
||
}
|
||
if dt == Int {
|
||
if vm.ints[bt*m+j] != wantI {
|
||
t.Errorf("bi,bij->bj int batch %d col %d = %d, want %d", bt, j, vm.ints[bt*m+j], wantI)
|
||
}
|
||
} else if d := math.Abs(float64(vm.floats32[bt*m+j]) - wantF); d > 1e-5*judgingScale(wantF) {
|
||
t.Errorf("bi,bij->bj float32 batch %d col %d = %v, want %v", bt, j, vm.floats32[bt*m+j], wantF)
|
||
}
|
||
}
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
// judgingScale grows the tolerance with the magnitude it guards.
|
||
func judgingScale(want float64) float64 {
|
||
s := math.Abs(want)
|
||
if s < 1 {
|
||
return 1
|
||
}
|
||
return s
|
||
}
|
||
|
||
// TestPinScanDimTypes walks CumSum and CumProd for int,
|
||
// float32 and complex against step-by-step references.
|
||
func TestPinScanDimTypes(t *testing.T) {
|
||
iv := []int64{3, -1, 2, 5, 0, 7}
|
||
ia := mustFromInts(t, iv, 2, 3)
|
||
cs, err := CumSum(ia, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// Row-major (2, 3): row 0 scans 3, -1, 2 and row 1 scans 5, 0, 7.
|
||
wantRow := [][]int64{{3, 2, 4}, {5, 5, 12}}
|
||
for i := range 2 {
|
||
for j := range 3 {
|
||
if cs.ints[i*3+j] != wantRow[i][j] {
|
||
t.Errorf("CumSum int [%d][%d] = %d, want %d", i, j, cs.ints[i*3+j], wantRow[i][j])
|
||
}
|
||
}
|
||
}
|
||
cp, err := CumProd(ia, 0)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
for j := range 3 {
|
||
if cp.ints[j] != iv[j] {
|
||
t.Errorf("CumProd int first row [%d] = %d, want %d", j, cp.ints[j], iv[j])
|
||
}
|
||
if cp.ints[3+j] != iv[j]*iv[3+j] {
|
||
t.Errorf("CumProd int second row [%d] = %d, want %d", j, cp.ints[3+j], iv[j]*iv[3+j])
|
||
}
|
||
}
|
||
// Float32 narrows the carry once per step, exactly as the stored
|
||
// value feeds the next one.
|
||
fv := []float32{0.5, -1.25, 3.5, 2, -0.125, 4}
|
||
fa, err := FromFloat32s(fv, 2, 3)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
fcs, err := CumSum(fa, 0)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
for j := range 3 {
|
||
carry := float64(fv[j])
|
||
if math.Float32bits(fcs.floats32[j]) != math.Float32bits(float32(carry)) {
|
||
t.Errorf("CumSum float32 [%d] = %v, want %v", j, fcs.floats32[j], float32(carry))
|
||
}
|
||
for i := 1; i < 2; i++ {
|
||
carry = float64(float32(carry)) + float64(fv[i*3+j])
|
||
if math.Float32bits(fcs.floats32[i*3+j]) != math.Float32bits(float32(carry)) {
|
||
t.Errorf("CumSum float32 [%d] = %v, want %v", i*3+j, fcs.floats32[i*3+j], float32(carry))
|
||
}
|
||
}
|
||
}
|
||
fcp, err := CumProd(fa, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
for i := range 2 {
|
||
carry := float32(1)
|
||
for j := range 3 {
|
||
carry = float32(float64(carry) * float64(fv[i*3+j]))
|
||
if math.Float32bits(fcp.floats32[i*3+j]) != math.Float32bits(carry) {
|
||
t.Errorf("CumProd float32 [%d] = %v, want %v", i*3+j, fcp.floats32[i*3+j], carry)
|
||
}
|
||
}
|
||
}
|
||
// Complex adds exactly.
|
||
cv := []complex128{1 + 1i, 2, -1i, 1}
|
||
ca := mustFromComplexes(t, cv, 2, 2)
|
||
ccs, err := CumSum(ca, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if ccs.complexes[0] != 1+1i || ccs.complexes[1] != 3+1i ||
|
||
ccs.complexes[2] != -1i || ccs.complexes[3] != 1-1i {
|
||
t.Errorf("CumSum complex = %v", ccs.complexes)
|
||
}
|
||
// CumProd along dim 1 accumulates within each row.
|
||
ccp, err := CumProd(ca, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// Row 0: (1+1i)·2 = 2+2i; row 1: (−1i)·1 = −1i.
|
||
if ccp.complexes[0] != 1+1i || ccp.complexes[1] != 2+2i ||
|
||
ccp.complexes[2] != -1i || ccp.complexes[3] != -1i {
|
||
t.Errorf("CumProd complex = %v", ccp.complexes)
|
||
}
|
||
}
|
||
|
||
// TestPinMeanAxisEmptyDim pins the NaN fill a mean over an
|
||
// empty dimension reports.
|
||
func TestPinMeanAxisEmptyDim(t *testing.T) {
|
||
a, err := Zeros(Float, 3, 0, 4)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
m, err := MeanAxis(a, 1)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if !sameShape(m.Shape(), []int{3, 4}) {
|
||
t.Fatalf("MeanAxis over an empty dim has shape %s, want (3, 4)", shapeText(m.Shape()))
|
||
}
|
||
for i := range m.Len() {
|
||
if !math.IsNaN(m.floats[i]) {
|
||
t.Errorf("MeanAxis over an empty dim [%d] = %v, want NaN", i, m.floats[i])
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestPinEinsumSlotSumMultiOperand walks the three- and
|
||
// four-operand slot sums against naive loops, float64 and float32.
|
||
func TestPinEinsumSlotSumMultiOperand(t *testing.T) {
|
||
rng := rand.New(rand.NewPCG(607, 7))
|
||
for _, dt := range []Dtype{Float, Float32} {
|
||
a := randWhole(t, rng, dt, 3, 4)
|
||
b := randWhole(t, rng, dt, 4, 5)
|
||
c := randWhole(t, rng, dt, 5, 6)
|
||
got3, err := Einsum("ik,kj,jl->il", a, b, c)
|
||
if err != nil {
|
||
t.Fatalf("3 operands %s: %v", dt, err)
|
||
}
|
||
for i := range 3 {
|
||
for l := range 6 {
|
||
var want float64
|
||
for k := range 4 {
|
||
for j := range 5 {
|
||
want += float64At(a, i*4+k) * float64At(b, k*5+j) * float64At(c, j*6+l)
|
||
}
|
||
}
|
||
g := float64At(got3, i*6+l)
|
||
if math.Abs(g-want) > 1e-9*judgingScale(want) {
|
||
t.Errorf("3 operands %s [%d,%d] = %v, want %v", dt, i, l, g, want)
|
||
}
|
||
}
|
||
}
|
||
d := randWhole(t, rng, dt, 6, 2)
|
||
got4, err := Einsum("ik,kj,jl,lm->im", a, b, c, d)
|
||
if err != nil {
|
||
t.Fatalf("4 operands %s: %v", dt, err)
|
||
}
|
||
for i := range 3 {
|
||
for m := range 2 {
|
||
var want float64
|
||
for k := range 4 {
|
||
for j := range 5 {
|
||
for l := range 6 {
|
||
want += float64At(a, i*4+k) * float64At(b, k*5+j) *
|
||
float64At(c, j*6+l) * float64At(d, l*2+m)
|
||
}
|
||
}
|
||
}
|
||
g := float64At(got4, i*2+m)
|
||
tol := 1e-9
|
||
if dt == Float32 {
|
||
tol = 1e-5
|
||
}
|
||
if math.Abs(g-want) > tol*judgingScale(want) {
|
||
t.Errorf("4 operands %s [%d,%d] = %v, want %v", dt, i, m, g, want)
|
||
}
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
// randWhole builds a small random array of the given dtype, holding
|
||
// only whole numbers so the naive references stay exact.
|
||
func randWhole(t *testing.T, rng *rand.Rand, dt Dtype, shape ...int) *Array {
|
||
t.Helper()
|
||
n := 1
|
||
for _, d := range shape {
|
||
n *= d
|
||
}
|
||
switch dt {
|
||
case Int:
|
||
vals := make([]int64, n)
|
||
for i := range vals {
|
||
vals[i] = int64(rng.IntN(7) - 3)
|
||
}
|
||
return mustFromInts(t, vals, shape...)
|
||
case Float32:
|
||
vals := make([]float32, n)
|
||
for i := range vals {
|
||
vals[i] = float32(rng.IntN(7) - 3)
|
||
}
|
||
a, err := FromFloat32s(vals, shape...)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
return a
|
||
default:
|
||
vals := make([]float64, n)
|
||
for i := range vals {
|
||
vals[i] = float64(rng.IntN(7) - 3)
|
||
}
|
||
return mustFloats(t, vals, shape...)
|
||
}
|
||
}
|
||
|
||
// int64At reads element i as an int64 magnitudes-only value.
|
||
func int64At(a *Array, i int) int64 {
|
||
switch a.Dtype() {
|
||
case Int:
|
||
return a.ints[i]
|
||
case Float32:
|
||
return int64(a.floats32[i])
|
||
default:
|
||
return int64(a.floats[i])
|
||
}
|
||
}
|
||
|
||
// float64At reads element i as a float64 value.
|
||
func float64At(a *Array, i int) float64 {
|
||
switch a.Dtype() {
|
||
case Float32:
|
||
return float64(a.floats32[i])
|
||
default:
|
||
return a.floatAt(i)
|
||
}
|
||
}
|