// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import ( "math" "testing" ) // TestEinsumGeneralReductions pins patterns only the general engine // reaches: reduction-only axes, implicit output, arbitrary order. func TestEinsumGeneralReductions(t *testing.T) { a, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) // "ij->i": row sums. rs, err := Einsum("ij->i", a) if err != nil { t.Fatalf("ij->i: %v", err) } if rs.FloatAt(0) != 6 || rs.FloatAt(1) != 15 { t.Fatalf("row sums = (%g, %g), want (6, 15)", rs.FloatAt(0), rs.FloatAt(1)) } // "ij->j": column sums. cs, err := Einsum("ij->j", a) if err != nil { t.Fatalf("ij->j: %v", err) } for j, want := range []float64{5, 7, 9} { if cs.FloatAt(j) != want { t.Fatalf("col sum %d = %g, want %g", j, cs.FloatAt(j), want) } } // Implicit output: "ij,jk" == "ij,jk->ik". b, _ := FromFloats([]float64{1, 0, 0, 1, 2, -1}, 3, 2) m1, err := Einsum("ij,jk", a, b) if err != nil { t.Fatalf("implicit: %v", err) } m2, _ := Einsum("ij,jk->ik", a, b) if !sameValues(m1, m2) { t.Fatal("implicit output differs from explicit") } // Output order "ij->ji" through the general engine as well. tr, err := Einsum("ij->ji", a) if err != nil { t.Fatalf("ij->ji: %v", err) } if tr.FloatAt(0) != 1 || tr.FloatAt(1) != 4 || tr.FloatAt(2) != 2 { t.Fatalf("transpose wrong: %v %v %v", tr.FloatAt(0), tr.FloatAt(1), tr.FloatAt(2)) } } func sameValues(a, b *Array) bool { if a.Len() != b.Len() { return false } for i := range a.Len() { if a.Dtype() == Complex { if a.ComplexAt(i) != b.ComplexAt(i) { return false } } else if a.FloatAt(i) != b.FloatAt(i) { return false } } return true } // TestEinsumEllipsis pins batched and broadcast patterns. func TestEinsumEllipsis(t *testing.T) { // Batched matmul "...ij,...jk->...ik" against the manual loop. av := make([]float64, 2*3*4) bv := make([]float64, 2*4*5) for i := range av { av[i] = float64(i%7) - 3 } for i := range bv { bv[i] = float64(i%5) - 2 } a, _ := FromFloats(av, 2, 3, 4) b, _ := FromFloats(bv, 2, 4, 5) got, err := Einsum("...ij,...jk->...ik", a, b) if err != nil { t.Fatalf("ellipsis batched: %v", err) } if got.NDim() != 3 || got.Shape()[0] != 2 || got.Shape()[1] != 3 || got.Shape()[2] != 5 { t.Fatalf("shape %v, want (2, 3, 5)", got.Shape()) } for bt := range 2 { for i := range 3 { for k := range 5 { want := 0.0 for j := range 4 { want += a.FloatAt((bt*3+i)*4+j) * b.FloatAt((bt*4+j)*5+k) } if math.Abs(got.FloatAt((bt*3+i)*5+k)-want) > 1e-12 { t.Fatalf("batched [%d,%d,%d] = %g, want %g", bt, i, k, got.FloatAt((bt*3+i)*5+k), want) } } } } // Broadcast: (1, 4) x (2, 4) -> (2, ...) inner product per batch. u, _ := FromFloats([]float64{1, 2, 3, 4}, 4) v, _ := FromFloats([]float64{1, 0, 0, 0, 0, 1, 0, 0}, 2, 4) dots, err := Einsum("i,...i->...", u, v) if err != nil { t.Fatalf("broadcast dots: %v", err) } if dots.Shape()[0] != 2 || dots.FloatAt(0) != 1 || dots.FloatAt(1) != 2 { t.Fatalf("broadcast dots = %v %v", dots.FloatAt(0), dots.FloatAt(1)) } } // TestEinsumDiagonalGeneral pins repeated labels through the engine. func TestEinsumDiagonalGeneral(t *testing.T) { a, _ := FromFloats([]float64{1, 2, 3, 4}, 2, 2) // "ii->i" still hits the fast path; "iij->ij" only the engine can. b, _ := FromFloats([]float64{ 1, 2, 3, 4, 5, 6, 7, 8, }, 2, 2, 2) out, err := Einsum("iij->ij", b) if err != nil { t.Fatalf("iij->ij: %v", err) } if out.Shape()[0] != 2 || out.Shape()[1] != 2 { t.Fatalf("shape %v, want (2, 2)", out.Shape()) } // element [i][j] = b[i][i][j]. for i := range 2 { for j := range 2 { want := b.FloatAt((i*2+i)*2 + j) if out.FloatAt(i*2+j) != want { t.Fatalf("out[%d][%d] = %g, want %g", i, j, out.FloatAt(i*2+j), want) } } } // Bad specs stay errors. if _, err := Einsum("ij->jj", a); err == nil { t.Fatal("repeated output label accepted") } if _, err := Einsum("ij->k", a); err == nil { t.Fatal("unknown output label accepted") } } // TestEinsumSlotSplitWideWorkload pins the parallel slot walk's chunk // cursor. The general engine hands every worker a contiguous chunk of // output slots and each worker rebuilds the cursor of its own first slot // from the slot index, so the walk must return the same bits however the // slots are split. Both workloads hold thousands of slots, and the split // is asserted before it is compared so a shrunken workload cannot // quietly fall back to the single-chunk walk. func TestEinsumSlotSplitWideWorkload(t *testing.T) { const workers = 4 cases := []struct { spec string shapes [][]int sumTotal int // product of the summed labels' sizes }{ // 64·8·8·4 = 16384 slots and nothing summed, so the cursor is // the only source of every operand offset. {"ij,kl->ijkl", [][]int{{64, 8}, {8, 4}}, 1}, // 64·16 = 1024 slots of 8·8·3 = 192 visits: the sum runs inside // the worker that owns the slot. {"ik,kj,jl->il", [][]int{{64, 8}, {8, 8}, {8, 16}}, 64}, } prev := NumWorkers() defer SetNumCPU(prev) for _, tc := range cases { for _, dt := range []Dtype{Int, Float} { operands := make([]*Array, len(tc.shapes)) for i, shape := range tc.shapes { operands[i] = einsumOperand(t, dt, shape...) } SetNumCPU(1) want, err := Einsum(tc.spec, operands...) if err != nil { t.Fatalf("%s/%s: %v", tc.spec, dt, err) } // The dispatch splits only while a worker's chunk reaches // the visit floor, so assert the workload still does. perSlot := tc.sumTotal * len(operands) minSlots := 1 if perSlot < einsumSlotFloor { minSlots = (einsumSlotFloor + perSlot - 1) / perSlot } slots := want.Len() if chunk := (slots + workers - 1) / workers; chunk < minSlots { t.Fatalf("%s/%s: %d slots no longer split at %d workers: chunk %d below the %d-slot floor", tc.spec, dt, slots, workers, chunk, minSlots) } // A pure outer product is the product of one element of each // operand, so its corners are checked against that definition: // the comparison below cannot pass on a walk that is wrong in // every chunk. if tc.sumTotal == 1 { for _, c := range [][4]int{{0, 0, 0, 0}, {7, 3, 5, 1}, {63, 7, 7, 3}} { var x, y, g float64 if dt == Int { xi, _ := IntAt(operands[0], c[0], c[1]) yi, _ := IntAt(operands[1], c[2], c[3]) gi, _ := IntAt(want, c[0], c[1], c[2], c[3]) x, y, g = float64(xi), float64(yi), float64(gi) } else { x, _ = FloatAt(operands[0], c[0], c[1]) y, _ = FloatAt(operands[1], c[2], c[3]) g, _ = FloatAt(want, c[0], c[1], c[2], c[3]) } if g != x*y { t.Fatalf("%s/%s: slot %v = %v, want %v·%v = %v", tc.spec, dt, c, g, x, y, x*y) } } } SetNumCPU(workers) got, err := Einsum(tc.spec, operands...) if err != nil { t.Fatalf("%s/%s: split walk: %v", tc.spec, dt, err) } if !einsumBitsEqual(got, want) { t.Fatalf("%s/%s: the split walk disagrees with the serial one at element %d", tc.spec, dt, firstDifferingSlot(got, want)) } } } } // firstDifferingSlot returns the flat index of the first element on // which two results disagree, or -1 when none does: a failure names one // element instead of printing two whole payloads. func firstDifferingSlot(a, b *Array) int { if a.Len() != b.Len() || a.Dtype() != b.Dtype() { return -1 } for i := range a.Len() { switch a.dt { case Int: if a.ints[i] != b.ints[i] { return i } case Float32: if a.floats32[i] != b.floats32[i] { return i } case Float: if a.floats[i] != b.floats[i] { return i } case Complex: if a.complexes[i] != b.complexes[i] { return i } } } return -1 }