Files
tensor/internal/core/einsum_test.go
T

257 lines
7.7 KiB
Go
Raw 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"
"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
}