257 lines
7.7 KiB
Go
257 lines
7.7 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
|
||
|
|
}
|