// Copyright (c) 2026 Petr Balvín (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) } }