Files

823 lines
25 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"
"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)
}
}