feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,256 @@
|
||||
// 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
|
||||
}
|
||||
Reference in New Issue
Block a user