Files
tensor/internal/core/guard_pins_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

823 lines
25 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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)
}
}