Files

1654 lines
49 KiB
Go
Raw Permalink 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 (
"slices"
"strings"
"sourcedock.dev/petrbalvin/tensor/internal/engine"
)
// Einsum implements Einstein summation for a useful subset of
// operations. The spec is a string of the form
// "lhs[,lhs,...]->rhs". Each label (an ASCII letter a-z or A-Z) names
// a dimension; the same letter across operands must agree. Any other
// character is an error, so an invalid spec can never silently compute.
//
// Supported cases:
//
// "ij,jk->ik" matrix multiply
// "ij,ji->" Frobenius inner product of two matrices
// "ij,ij->" same: sum of element-wise products
// "ij,ij->ij" element-wise product
// "ij->ji" transpose
// "ii->i" diagonal
// "ii->" trace
// "ij->" full sum
// "i->i" identity (returns the operand)
//
// Other patterns, among them broadcasting over the ellipsis,
// reduction-only axes (e.g. "ij->i") and batched matmul
// ("bij,bjk->bik"), are handled by the general engine behind the same
// entry point; an ellipsis must lead the subscripts it appears in.
func Einsum(spec string, operands ...*Array) (*Array, error) {
for _, op := range operands {
// The einsum engines (table, batched and general) are built on
// the matmul kernels and are not offered for the half dtype;
// the refusal is loud and named rather than a silent payload
// misread, and the Astype conversion is cheap. The narrow
// element types ride the same refusal until their kernels
// arrive.
if op.Dtype() == Float16 {
return nil, errf("Einsum: float16 operands are not supported; convert with Astype")
}
if narrowRefused(op.Dtype()) {
return nil, errf("Einsum: dtype %s is not supported; convert with Astype", op.Dtype())
}
}
lhsPart, rhsPart, hasArrow := strings.Cut(spec, "->")
lhsStrs := strings.Split(lhsPart, ",")
if len(lhsStrs) != len(operands) {
return nil, errf("Einsum: spec %q has %d operands, got %d", spec, len(lhsStrs), len(operands))
}
// The general engine owns the patterns the table below does not
// table: ellipsis broadcasting, implicit output, arbitrary sums.
if !hasArrow || strings.Contains(spec, ".") {
return einsumGeneral(lhsStrs, rhsPart, hasArrow, operands)
}
lhs := make([][]int, len(operands))
for i, s := range lhsStrs {
l, lerr := einsumLabels(s)
if lerr != nil {
return nil, lerr
}
lhs[i] = l
}
rhs, rerr := einsumLabels(rhsPart)
if rerr != nil {
return nil, rerr
}
out, err := einsumEval(lhs, rhs, operands)
if err == nil {
return out, nil
}
// The table declined; the general engine gets the last word so a
// valid pattern never errors twice. Its error is the one reported:
// it saw the whole pattern, while the table can only say that no
// kernel matched. A shape that would previously print a label table
// now names the label that disagrees.
if g, gerr := einsumGeneral(lhsStrs, rhsPart, true, operands); gerr == nil {
return g, nil
} else {
return nil, gerr
}
}
// einsumLabels converts a label string like "ij" into a slice of int
// labels (a=0, b=1, …, z=25, A=26, …, Z=51). Empty string returns nil.
// Every character must be an ASCII letter, matching einsumParse in the
// general engine: digits and anything else are a named error quoting
// the offending character, never a silent misread.
func einsumLabels(s string) ([]int, error) {
if s == "" {
return nil, nil
}
out := make([]int, 0, len(s))
for _, c := range s {
switch {
case c >= 'a' && c <= 'z':
out = append(out, int(c-'a'))
case c >= 'A' && c <= 'Z':
out = append(out, int(c-'A')+26)
default:
return nil, errf("Einsum: bad label %q in %q", string(c), s)
}
}
return out, nil
}
// einsumEval is the dispatcher. It inspects the spec and routes to
// one of the concrete implementations.
func einsumEval(lhs [][]int, rhs []int, operands []*Array) (*Array, error) {
// An output label may appear once: "ii->ii" is a typo for the
// diagonal "ii->i", and the identity fast path below would silently
// return a copy of the whole operand instead. The general engine
// refuses a repeated output label as well.
for i, l := range rhs {
if slices.Contains(rhs[:i], l) {
return nil, errf("Einsum: output label %q repeats", einsumLabelName(l))
}
}
// Validate operand shapes against labels.
for i, op := range operands {
if op.NDim() != len(lhs[i]) {
return nil, errf("Einsum: operand %d has shape %s but spec has %d labels", i, shapeText(op.shape), len(lhs[i]))
}
}
// Trace "ii->": sum of diagonal. Routed through Diagonal + Sum so the
// scalar keeps the operand's dtype (int stays int, exact above 2^53)
// and complex operands work, exactly as the general engine would
// answer them.
if len(lhs) == 1 && len(lhs[0]) == 2 && lhs[0][0] == lhs[0][1] && len(rhs) == 0 {
op := operands[0]
if op.shape[0] != op.shape[1] {
return nil, errf("Einsum: label %q has sizes %d and %d in operand 0",
einsumLabelName(lhs[0][0]), op.shape[0], op.shape[1])
}
diag, derr := Diagonal(op, 0)
if derr != nil {
return nil, derr
}
return scalarArray(Sum(diag))
}
// Diagonal "ii->i": extract main diagonal as 1-D.
if len(lhs) == 1 && len(lhs[0]) == 2 && lhs[0][0] == lhs[0][1] &&
len(rhs) == 1 && rhs[0] == lhs[0][0] {
return Diagonal(operands[0], 0)
}
// Identity "xn->xn": a copy, not the operand; results never alias
// their inputs.
if len(lhs) == 1 && einsumEqLabels(lhs[0], rhs) {
return Copy(operands[0]), nil
}
// Full sum "...": "xN,..->"
if len(rhs) == 0 {
return einsumReduceAll(lhs, operands)
}
// Transpose "ij->ji" or other permutation.
if len(lhs) == 1 && len(rhs) == len(lhs[0]) && einsumIsPermutation(lhs[0], rhs) {
return einsumTranspose(operands[0], lhs[0], rhs)
}
// Reduction-only axes "ij->i", "ijk->ik": the axes the output drops
// are summed away, the surviving ones keep their order.
if len(lhs) == 1 && len(rhs) > 0 && len(rhs) < len(lhs[0]) {
if res, ok, err := einsumSumAxes(lhs[0], rhs, operands[0]); ok {
return res, err
}
}
// Matrix multiply "ij,jk->ik" (or equivalent with letters).
if len(lhs) == 2 && len(rhs) == 2 &&
len(lhs[0]) == 2 && len(lhs[1]) == 2 &&
lhs[0][1] == lhs[1][0] && lhs[0][0] == rhs[0] && lhs[1][1] == rhs[1] {
return MatMul2D(operands[0], operands[1])
}
// Matrix-vector "ij,j->i".
if len(lhs) == 2 && len(rhs) == 1 &&
len(lhs[0]) == 2 && len(lhs[1]) == 1 &&
lhs[0][1] == lhs[1][0] && lhs[0][0] == rhs[0] {
return MatMul2D(operands[0], operands[1])
}
// Vector-matrix "i,ij->j".
if len(lhs) == 2 && len(rhs) == 1 &&
len(lhs[0]) == 1 && len(lhs[1]) == 2 &&
lhs[0][0] == lhs[1][0] && lhs[1][1] == rhs[0] {
return MatMul2D(operands[0], operands[1])
}
// Batched products "bij,bjk->bik", "bij,bj->bi" and "bi,bij->bj":
// the shared leading label multiplies matching contiguous batch
// slices into one shared payload. For a fixed output slot the
// walk's summed label visits ascending, which is MatMul2D's
// ascending-p accumulation per cell, so each batch is the walk's
// slot sum (the Float32 kernel keeps its documented narrow-once
// semantics, as on the 2-D path above). Only patterns with
// pairwise-distinct labels and exactly agreeing shared dimensions
// are taken; broadcast batches and repeated labels stay with the
// engine.
if res, ok, err := einsumBatched(lhs, rhs, operands); ok {
return res, err
}
// Dot product "i,i->".
if len(lhs) == 2 && len(rhs) == 0 &&
len(lhs[0]) == 1 && len(lhs[1]) == 1 && lhs[0][0] == lhs[1][0] {
s, err := Dot(operands[0], operands[1])
if err != nil {
return nil, err
}
return scalarArray(s)
}
// Outer product "i,j->ij".
if len(lhs) == 2 && len(rhs) == 2 &&
len(lhs[0]) == 1 && len(lhs[1]) == 1 && lhs[0][0] != lhs[1][0] {
return einsumOuter(operands[0], operands[1], rhs, lhs)
}
// Element-wise product with output "ij,ij->ij".
if len(lhs) == 2 && len(rhs) == 2 &&
einsumEqLabels(lhs[0], lhs[1]) && einsumEqLabels(lhs[0], rhs) {
return Mul(operands[0], operands[1])
}
// Element-wise multiply, reduce-all "ij,ij->".
if len(lhs) == 2 && len(rhs) == 0 &&
einsumEqLabels(lhs[0], lhs[1]) {
prod, err := Mul(operands[0], operands[1])
if err != nil {
return nil, err
}
return scalarArray(Sum(prod))
}
return nil, errf("Einsum: unsupported pattern lhs=%v rhs=%v", lhs, rhs)
}
// einsumSumAxes folds away the axes an explicit output leaves out of a
// single operand, the patterns "ij->i" and "kji->k" reduce to. It
// declines, reporting ok=false, unless the axis fold reproduces the
// engine's walk bit for bit:
//
// - the output must read as the operand's labels with some of them
// dropped, in order: a reordered output or a label the operand does
// not have is not this pattern, and a repeated label is a diagonal
// the fold would not take;
// - the operand must be contiguous, as every kernel reading a payload
// window demands;
// - a float32 operand stays with the engine, which narrows into the
// float32 slot once per addend while the fold widens, sums in
// float64 and narrows once at the end;
// - a complex operand stays with the engine, whose product seed
// multiplies the first operand by 1+0i: that is the identity for
// finite values but turns an infinite component into a NaN.
//
// The fold runs the dropped axes from the largest label id down, which
// is what makes its nesting the engine's: the engine walks the sorted
// label union with the smallest id outermost, so the innermost loop
// carries the largest id and has to be closed first.
func einsumSumAxes(labels, rhs []int, a *Array) (*Array, bool, error) {
if a.strides != nil || a.dt == Float32 || a.dt == Complex {
return nil, false, nil
}
for i, l := range labels {
if slices.Contains(labels[:i], l) {
return nil, false, nil
}
}
// Which axes the output skips: every output label must be met in
// the operand's own order.
keep := make([]bool, len(labels))
j := 0
for _, r := range rhs {
for j < len(labels) && labels[j] != r {
j++
}
if j == len(labels) {
return nil, false, nil
}
keep[j] = true
j++
}
// The dropped axes as (position, label); a fold removes one axis and
// shifts the positions after it down by one.
rest := make([]einsumAxis, 0, len(labels)-len(rhs))
for axis, l := range labels {
if !keep[axis] {
rest = append(rest, einsumAxis{pos: axis, label: l})
}
}
cur := a
for len(rest) > 0 {
last := 0
for i, r := range rest {
if r.label > rest[last].label {
last = i
}
}
next, err := SumAxis(cur, rest[last].pos)
if err != nil {
return nil, true, err
}
cur = next
removed := rest[last].pos
rest = slices.Delete(rest, last, last+1)
for i := range rest {
if rest[i].pos > removed {
rest[i].pos--
}
}
}
return cur, true, nil
}
// einsumAxis is one axis of an operand on its way to being folded: its
// current position and the label it carries.
type einsumAxis struct {
pos int
label int
}
// einsumBatched recognises the batched product patterns whose shared
// leading label multiplies matching batch slices into one payload:
// "bij,bjk->bik", "bij,bj->bi" and "bi,bij->bj". It declines,
// reporting ok=false, unless every label in the pattern is pairwise
// distinct and the shared dimensions agree exactly; broadcast batches
// and repeated labels keep the general engine's semantics. ok=true
// means the returned result, or error, is final.
func einsumBatched(lhs [][]int, rhs []int, operands []*Array) (*Array, bool, error) {
if len(lhs) != 2 {
return nil, false, nil
}
x, y := lhs[0], lhs[1]
sa, sb := operands[0].Shape(), operands[1].Shape()
switch {
// Matrix times matrix: a (b,i,j), b (b,j,k), out (b,i,k).
case len(x) == 3 && len(y) == 3 && len(rhs) == 3 &&
x[0] == y[0] && x[0] == rhs[0] && x[1] == rhs[1] && y[2] == rhs[2] && x[2] == y[1] &&
x[0] != x[1] && x[0] != x[2] && x[0] != y[2] &&
x[1] != x[2] && x[1] != y[2] && x[2] != y[2] &&
sa[0] == sb[0] && sa[2] == sb[1]:
// Matrix times vector: a (b,i,j), b (b,j), out (b,i).
case len(x) == 3 && len(y) == 2 && len(rhs) == 2 &&
x[0] == y[0] && x[0] == rhs[0] && x[1] == rhs[1] && x[2] == y[1] &&
x[0] != x[1] && x[0] != x[2] && x[1] != x[2] &&
sa[0] == sb[0] && sa[2] == sb[1]:
// Vector times matrix: a (b,i), b (b,i,j), out (b,j).
case len(x) == 2 && len(y) == 3 && len(rhs) == 2 &&
x[0] == y[0] && x[0] == rhs[0] && x[1] == y[1] && y[2] == rhs[1] &&
x[0] != x[1] && x[0] != y[2] && x[1] != y[2] &&
sa[0] == sb[0] && sa[1] == sb[1]:
default:
return nil, false, nil
}
res, err := einsumBatchedProduct(operands[0], operands[1])
if err != nil {
return nil, true, err
}
return res, true, nil
}
// einsumBatchedProduct multiplies matching batches of a and b straight
// into one shared output payload. Both operands are materialised
// first: the per-batch slices must be contiguous payload windows. Each
// batch runs the panel kernels MatMul2D runs for its shape over the
// same ascending-p walk, so the payload holds exactly the stacked
// MatMul2D results; assembling through the kernels rather than through
// MatMul2D only removes the per-batch output array and the per-batch
// goroutine fan-out. Mixed real operands convert once for the whole
// tensor, which reads back the values MatMul2D's per-batch conversions
// produced.
func einsumBatchedProduct(a, b *Array) (*Array, error) {
a = a.materialise()
b = b.materialise()
sa, sb := a.shape, b.shape
batches := sa[0]
// Per-batch result shape: a's batch dims minus the summed axis,
// then b's trailing dims.
outTail := append(append([]int{}, sa[1:len(sa)-1]...), sb[2:]...)
tailTotal := 1
for _, d := range outTail {
tailTotal *= d
}
out := &Array{shape: append([]int{batches}, outTail...), dt: promote(a.dt, b.dt)}
out.alloc(batches * tailTotal)
if batches == 0 {
return out, nil
}
// One batch's geometry, named as MatMul2D names its dimensions:
// "bij,bjk->bik" multiplies an n×k batch by a k×m batch,
// "bij,bj->bi" an n×k batch by a k batch, and "bi,bij->bj" a k
// batch by a k×m batch. nOut is the batch's output row count.
bmm, bmv, vbm := false, false, false
var n, k, m int
switch {
case len(sa) == 3 && len(sb) == 3:
bmm = true
n, k, m = sa[1], sa[2], sb[2]
case len(sa) == 3:
bmv = true
n, k = sa[1], sa[2]
default:
vbm = true
k, m = sb[1], sb[2]
}
nOut := n
if vbm {
nOut = m
}
perA := a.Len() / batches
perB := b.Len() / batches
// Workers own disjoint output row blocks, so the kernels need no
// locks. Two layouts: a worker takes each batch whole, running the
// batch's kernel alone where the wide float64 panels win, or the
// batches' output rows share one global split over the full pool
// with narrow panels. Measured on the batched benchmarks, whole
// batches win from two batches up: each kernel then streams its b
// once per wide panel on a private core, instead of narrow panels
// sharing L1 under SMT and paying a goroutine fan-out per split.
// A lone batch cannot feed the pool with whole batches, so it
// keeps MatMul2D's own row split.
whole := batches > 1
drive := func(body func(walk func(visit func(bt, r0, r1 int)))) {
if whole {
engine.Parallel(batches, func(ks, ke int) {
body(func(visit func(bt, r0, r1 int)) {
for bt := ks; bt < ke; bt++ {
visit(bt, 0, nOut)
}
})
})
return
}
engine.Parallel(batches*nOut, func(rs, re int) {
body(func(visit func(bt, r0, r1 int)) {
for bt := rs / nOut; bt*nOut < re; bt++ {
r0 := max(rs, bt*nOut) - bt*nOut
r1 := min(re, (bt+1)*nOut) - bt*nOut
visit(bt, r0, r1)
}
})
})
}
switch out.dt {
case Int:
ai, bi, oi := a.ints, b.ints, out.ints
drive(func(walk func(visit func(bt, r0, r1 int))) {
walk(func(bt, r0, r1 int) {
aK := ai[bt*perA : bt*perA+perA]
bK := bi[bt*perB : bt*perB+perB]
oK := oi[bt*tailTotal : bt*tailTotal+tailTotal]
switch {
case bmm:
matMulIntRows(aK, bK, oK, r0, r1, k, m)
case bmv:
matVecRows(aK, bK, oK, r0, r1, k)
default:
vecMatCols(aK, bK, oK[r0:r1], m, r0, r1)
}
})
})
case Float32:
// Accumulate in float64 and round once, exactly as
// MatMul2D's Float32 path: the payloads read directly, every
// widening is exact, and each finished row narrows once. The
// pooled scratch serves the matmul panels (four m-wide rows)
// or the vecMat column walk (one accumulator row); each
// kernel clears the scratch it uses, so one buffer outlasts
// the worker's whole chunk.
aIs32, bIs32 := a.dt == Float32, b.dt == Float32
drive(func(walk func(visit func(bt, r0, r1 int))) {
var scratch []float64
switch {
case bmm:
scratch = engine.GetFloat64Buf(4 * m)
case vbm:
scratch = engine.GetFloat64Buf(m)
}
if scratch != nil {
defer engine.PutFloat64Buf(scratch)
}
walk(func(bt, r0, r1 int) {
oK := out.floats32[bt*tailTotal : bt*tailTotal+tailTotal]
var aK32, bK32 []float32
var aKi, bKi []int64
if aIs32 {
aK32 = a.floats32[bt*perA : bt*perA+perA]
} else {
aKi = a.ints[bt*perA : bt*perA+perA]
}
if bIs32 {
bK32 = b.floats32[bt*perB : bt*perB+perB]
} else {
bKi = b.ints[bt*perB : bt*perB+perB]
}
switch {
case bmm:
switch {
case aIs32 && bIs32:
matMulF32Rows(aK32, bK32, oK, scratch, r0, r1, k, m)
case aIs32:
matMulF32Rows(aK32, bKi, oK, scratch, r0, r1, k, m)
case bIs32:
matMulF32Rows(aKi, bK32, oK, scratch, r0, r1, k, m)
default:
matMulF32Rows(aKi, bKi, oK, scratch, r0, r1, k, m)
}
case bmv:
switch {
case aIs32 && bIs32:
matVecF32Rows(aK32, bK32, oK, r0, r1, k)
case aIs32:
matVecF32Rows(aK32, bKi, oK, r0, r1, k)
case bIs32:
matVecF32Rows(aKi, bK32, oK, r0, r1, k)
default:
matVecF32Rows(aKi, bKi, oK, r0, r1, k)
}
default:
switch {
// The column walk narrows over the accumulator's
// own length, so the scratch must be exactly the
// window's width, as MatMul2D sizes it.
case aIs32 && bIs32:
vecMatF32Cols(aK32, bK32, oK[r0:r1], scratch[:r1-r0], m, r0, r1)
case aIs32:
vecMatF32Cols(aK32, bKi, oK[r0:r1], scratch[:r1-r0], m, r0, r1)
case bIs32:
vecMatF32Cols(aKi, bK32, oK[r0:r1], scratch[:r1-r0], m, r0, r1)
default:
vecMatF32Cols(aKi, bKi, oK[r0:r1], scratch[:r1-r0], m, r0, r1)
}
}
})
})
case Float:
// Mixed real operands convert once at entry, in the same
// element order MatMul2D's per-batch conversions read, so the
// kernel streams plain rows; a float64 operand's payload
// already is that stream and skips the copy.
af := a.floats
if a.dt != Float {
af = denseFloatsLocal(a, a.Len(), 1)
}
bf := b.floats
if b.dt != Float {
bf = denseFloatsLocal(b, b.Len(), 1)
}
of := out.floats
// A lone worker runs the wide 8-row panel; a split range runs
// the narrow one (see matMulF64Rows). The whole-batch layout
// runs each kernel alone; the row split mirrors MatMul2D's
// worker test for its rows.
wide := whole || workersFor(nOut) == 1
drive(func(walk func(visit func(bt, r0, r1 int))) {
walk(func(bt, r0, r1 int) {
aK := af[bt*perA : bt*perA+perA]
bK := bf[bt*perB : bt*perB+perB]
oK := of[bt*tailTotal : bt*tailTotal+tailTotal]
switch {
case bmm:
matMulF64Rows(aK, bK, oK, r0, r1, k, m, wide)
case bmv:
matVecF64Rows(aK, bK, oK, r0, r1, k)
default:
vecMatF64Cols(aK, bK, oK[r0:r1], m, r0, r1)
}
})
})
default:
ac := complexPayload(a)
bc := complexPayload(b)
oc := out.complexes
drive(func(walk func(visit func(bt, r0, r1 int))) {
walk(func(bt, r0, r1 int) {
aK := ac[bt*perA : bt*perA+perA]
bK := bc[bt*perB : bt*perB+perB]
oK := oc[bt*tailTotal : bt*tailTotal+tailTotal]
switch {
case bmm:
// The plain i-p-j walk MatMul2D runs, bounded to
// this batch's rows.
for i := r0; i < r1; i++ {
orow := oK[i*m : (i+1)*m]
for p := range k {
av := aK[i*k+p]
for j := range m {
orow[j] += av * bK[p*m+j]
}
}
}
case bmv:
matVecRows(aK, bK, oK, r0, r1, k)
default:
vecMatCols(aK, bK, oK[r0:r1], m, r0, r1)
}
})
})
}
return out, nil
}
// scalarArray wraps a Scalar as a length-1 array of its own dtype, so
// reduced einsum results keep the complex or int payload instead of
// collapsing to the real part.
func scalarArray(s Scalar) (*Array, error) {
switch {
case s.IsComplex():
return FromComplexes([]complex128{s.Complex()}, 1)
case s.IsFloat():
return FromFloats([]float64{s.Float()}, 1)
default:
return FromInts([]int64{s.Int()}, 1)
}
}
func einsumEqLabels(a, b []int) bool {
return slices.Equal(a, b)
}
// einsumIsPermutation reports whether rhs is a permutation of lhs. A
// label repeated on either side stays a permutation of itself, exactly
// as the membership test it replaces answered.
func einsumIsPermutation(lhs, rhs []int) bool {
if len(lhs) != len(rhs) {
return false
}
for _, r := range rhs {
if !slices.Contains(lhs, r) {
return false
}
}
return true
}
// einsumTranspose reorders the axes of a according to the mapping
// lhs to rhs. lhs is the current order; rhs is the desired order.
func einsumTranspose(a *Array, lhs, rhs []int) (*Array, error) {
dims := make([]int, len(rhs))
for i, r := range rhs {
// Find r in lhs.
j := -1
for k, l := range lhs {
if l == r {
j = k
break
}
}
if j < 0 {
return nil, errf("Einsum: rhs label %d not in lhs", r)
}
dims[i] = j
}
return TransposeAxes(a, dims...)
}
// einsumReduceAll handles patterns that reduce every axis (no
// "->rhs"). For 1 operand it's a full sum; for 2 operands it's a
// dot/inner product over matching axes: the second operand is first
// permuted into the first's label order, so "ij,ji->" pairs a's
// columns with b's rows instead of multiplying positionally.
func einsumReduceAll(lhs [][]int, operands []*Array) (*Array, error) {
if len(operands) == 1 {
return scalarArray(Sum(operands[0]))
}
if len(operands) == 2 {
b, err := einsumAlign(operands[1], lhs[1], lhs[0])
if err != nil {
return nil, err
}
prod, err := Mul(operands[0], b)
if err != nil {
return nil, err
}
return scalarArray(Sum(prod))
}
return nil, errf("Einsum: '->' with %d operands not supported", len(operands))
}
// einsumAlign permutes a so its labels read in the target order. Labels
// must appear exactly once in the operand; a already-aligned operand is
// returned unchanged, never a copy.
func einsumAlign(a *Array, labels, target []int) (*Array, error) {
for i, l := range labels {
if slices.Contains(labels[:i], l) {
return nil, errf("Einsum: repeated label in operand")
}
}
dims := make([]int, len(target))
for i, want := range target {
found := false
for k, l := range labels {
if l == want {
dims[i] = k
found = true
break
}
}
if !found {
return nil, errf("Einsum: label missing from operand, cannot align")
}
}
return TransposeAxes(a, dims...)
}
// einsumOuter computes the outer product of two 1-D vectors, producing
// a 2-D matrix indexed by the rhs labels. The result dtype follows the
// promotion ladder, so int vectors stay int and complex vectors stay
// complex.
func einsumOuter(a, b *Array, rhs []int, lhs [][]int) (*Array, error) {
if a.NDim() != 1 || b.NDim() != 1 {
return nil, errf("Einsum: outer product needs two 1-D operands, got %d and %d dims", a.NDim(), b.NDim())
}
// Output shape: in rhs order, each label picks the matching input
// size.
outShape := make([]int, len(rhs))
for i, label := range rhs {
switch label {
case lhs[0][0]:
outShape[i] = a.Len()
case lhs[1][0]:
outShape[i] = b.Len()
default:
return nil, errf("Einsum: outer label %d not in operands", label)
}
}
out := &Array{shape: outShape, dt: promote(a.dt, b.dt)}
total := 1
for _, d := range outShape {
total *= d
}
out.alloc(total)
// The first rhs slot is the row (slow) axis. When the rhs order
// swaps the operand labels ("i,j->ji"), the row index selects b's
// elements and the column index a's, so the flat write must follow
// the rhs order, not the operand order.
aIsRow := rhs[0] == lhs[0][0]
for i := range a.Len() {
for j := range b.Len() {
var flat int
if aIsRow {
flat = i*outShape[1] + j
} else {
flat = j*outShape[1] + i
}
switch out.dt {
case Int:
out.ints[flat] = a.ints[i] * b.ints[j]
case Complex:
out.complexes[flat] = a.complexAt(i) * b.complexAt(j)
default:
out.setFromValue(flat, a.floatAt(i)*b.floatAt(j))
}
}
}
return out, nil
}
// einsumLabelName spells a label the way the spec writes it: the lower
// case letters first, then the upper case ones.
func einsumLabelName(l int) string {
if l >= 26 {
return string(rune('A' + l - 26))
}
return string(rune('a' + l))
}
// einsumParse splits one operand's label string into explicit labels
// and an ellipsis flag. The ellipsis may appear once and must lead the
// subscripts, because the engine lays it onto the leading axes. Label
// ids run 0-25 for a-z and 26-51 for A-Z.
func einsumParse(s string) ([]int, bool, error) {
var labels []int
ellipsis := false
afterLabel := false
for _, c := range s {
switch {
case c == '.':
// Ellipsis dots arrive as three consecutive runes.
return nil, false, errf("Einsum: stray '.' in %q", s)
case c == '…':
if ellipsis {
return nil, false, errf("Einsum: two ellipses in %q", s)
}
if afterLabel {
// The engine lays the ellipsis axes onto the leading
// dimensions, so subscripts written after a label would
// be read as if they came first. Refuse rather than
// compute a different contraction than the one written.
return nil, false, errf("Einsum: the ellipsis must lead the subscripts, %q has a label before it", s)
}
ellipsis = true
case c >= 'a' && c <= 'z':
labels = append(labels, int(c-'a'))
afterLabel = true
case c >= 'A' && c <= 'Z':
labels = append(labels, int(c-'A')+26)
afterLabel = true
default:
return nil, false, errf("Einsum: bad label %q in %q", string(c), s)
}
}
return labels, ellipsis, nil
}
// einsumGeneral is the full Einstein-summation engine the pattern
// table above does not reach: ellipsis broadcasting, repeated labels
// (diagonals), reduction-only axes, arbitrary output order, implicit
// output, any operand count.
//
// Evaluation runs over label-position tables built once per call. The
// label union is sorted once; each operand carries, per sorted label
// position, the payload stride its dimensions contribute when that
// label's coordinate advances (repeated labels add, broadcast axes
// contribute zero). Output slots are written in row-major order by an
// odometer over the output labels, and within one slot the summed
// labels run as an inner odometer in ascending sorted order, so every
// slot receives exactly the products, in exactly the summation order,
// a full walk over the sorted label union would produce.
func einsumGeneral(specs []string, rhsSpec string, rhsExplicit bool, operands []*Array) (*Array, error) {
// The three-dot ellipsis becomes its single-rune form for parsing.
for i := range specs {
specs[i] = strings.ReplaceAll(specs[i], "...", "\u2026")
}
rhsSpec = strings.ReplaceAll(rhsSpec, "...", "\u2026")
type opInfo struct {
explicit []int
ellipsis bool
// label of each dimension after ellipsis expansion
dimLabels []int
// stride per dimension, zero on broadcast (size-1) axes
dimStrides []int
}
ops := make([]opInfo, len(operands))
// Synthetic labels for the ellipsis axes start past the alphabet.
nextSynthetic := 100
var allLabels []int
labelSeen := make(map[int]bool)
addLabel := func(l int) {
if !labelSeen[l] {
labelSeen[l] = true
allLabels = append(allLabels, l)
}
}
for p, s := range specs {
labels, ell, err := einsumParse(s)
if err != nil {
return nil, err
}
ops[p].explicit = labels
ops[p].ellipsis = ell
nd := operands[p].NDim()
ellDims := 0
if ell {
ellDims = nd - len(labels)
if ellDims < 0 {
return nil, errf("Einsum: operand %d has %d dims but spec %q wants more", p, nd, s)
}
} else if nd != len(labels) {
return nil, errf("Einsum: operand %d has shape %s but spec %q has %d labels",
p, shapeText(operands[p].shape), s, len(labels))
}
if ell {
for range ellDims {
addLabel(nextSynthetic)
nextSynthetic++
}
}
for _, l := range labels {
addLabel(l)
}
}
// Label sizes: explicit labels must agree exactly; an ellipsis axis
// broadcasts the same way, where only a size-1 axis stretches to the
// common size. The slot walk below expresses replication through a
// zero stride, and a zero stride is exactly what a size-1 axis
// contributes: a longer axis that disagrees with the broadcast
// extent has no stride that could read it, so it is refused rather
// than walked past the operand's elements.
sizes := make(map[int]int)
for p := range ops {
nd := operands[p].NDim()
ellDims := nd - len(ops[p].explicit)
dl := make([]int, 0, nd)
if ops[p].ellipsis {
for k := range ellDims {
dl = append(dl, 100+(nextSynthetic-100-ellDims+k))
}
}
dl = append(dl, ops[p].explicit...)
ops[p].dimLabels = dl
}
// The synthetic ids above were allocated per operand in order; two
// operands' ellipses must line up right-aligned, so renumber by
// distance from the right edge instead.
for p := range ops {
nd := operands[p].NDim()
ellDims := nd - len(ops[p].explicit)
if !ops[p].ellipsis {
continue
}
base := 100
for k := range ellDims {
// axis offset from the right among the ellipsis dims
ops[p].dimLabels[k] = base + (ellDims - 1 - k)
}
}
// Recompute the label union with the renumbered synthetic ids.
labelSeen = make(map[int]bool)
allLabels = allLabels[:0]
for p := range ops {
for _, l := range ops[p].dimLabels {
if !labelSeen[l] {
labelSeen[l] = true
allLabels = append(allLabels, l)
}
}
}
for l := range labelSeen {
sizes[l] = -1
}
for p := range ops {
for d, l := range ops[p].dimLabels {
dim := operands[p].shape[d]
cur := sizes[l]
if l >= 100 {
// A synthetic label is an ellipsis axis, and those
// broadcast: the widest operand sets the extent.
if dim > cur {
sizes[l] = dim
}
continue
}
// A written label must have the same length everywhere it
// appears, size 1 included: a lone 1 stretched against a
// longer labelled axis is refused, because a 1 in the wrong
// place is a shape typo far more often than an intended
// broadcast, and the package's own contract is that a
// shared letter means a shared length. Only ellipsis axes
// broadcast.
switch {
case cur == -1:
sizes[l] = dim
case cur != dim:
return nil, errf("Einsum: label %q is %d in one operand and %d in another",
einsumLabelName(l), cur, dim)
}
}
}
for _, s := range sizes {
if s == -1 {
return nil, errf("Einsum: unsized label")
}
}
// Every ellipsis axis must be 1 or the broadcast extent: the widest
// operand sets the size, and a size-1 axis beside it replicates. The
// check runs before any stride or slot walk, so a mismatched axis
// can never reach the readers, not even through a Slice view whose
// payload runs past Len().
for p := range ops {
for d, l := range ops[p].dimLabels {
if l < 100 {
continue
}
dim := operands[p].shape[d]
if dim == 1 || dim == sizes[l] {
continue
}
return nil, errf("Einsum: operand %d has shape %s, whose ellipsis axis of %d cannot broadcast against the extent %d; only a size-1 axis broadcasts",
p, shapeText(operands[p].shape), dim, sizes[l])
}
}
// Output labels.
var outLabels []int
if !rhsExplicit {
// Implicit output: labels appearing exactly once, sorted,
// ellipsis axes prepended.
counts := make(map[int]int)
for p := range ops {
for _, l := range ops[p].dimLabels {
counts[l]++
}
}
var once []int
for l, c := range counts {
if c == 1 && l < 100 {
once = append(once, l)
}
}
slices.Sort(once)
outLabels = append(outLabels, once...)
// Ellipsis dims go first in their broadcast order, leftmost
// operand axis first. The synthetic ids count from the right
// edge, so the operand's left-to-right order is descending id.
var ellAxes []int
for l := range sizes {
if l >= 100 {
ellAxes = append(ellAxes, l)
}
}
slices.Sort(ellAxes)
slices.Reverse(ellAxes)
outLabels = append(ellAxes, outLabels...)
} else {
rhsLabels, rhsEll, err := einsumParse(rhsSpec)
if err != nil {
return nil, err
}
seen := make(map[int]bool)
for _, l := range rhsLabels {
if seen[l] {
return nil, errf("Einsum: output label %q repeats", einsumLabelName(l))
}
seen[l] = true
if !labelSeen[l] {
return nil, errf("Einsum: output label %q not in any operand", einsumLabelName(l))
}
outLabels = append(outLabels, l)
}
if rhsEll {
var ellAxes []int
for l := range sizes {
if l >= 100 {
ellAxes = append(ellAxes, l)
}
}
slices.Sort(ellAxes)
slices.Reverse(ellAxes)
outLabels = slices.Concat(ellAxes, outLabels)
}
}
// Strides per operand dimension (0 on broadcast axes) and the
// output strides.
for p := range ops {
a := operands[p]
st := make([]int, a.NDim())
run := 1
for d := a.NDim() - 1; d >= 0; d-- {
st[d] = run
run *= a.shape[d]
}
for d, l := range ops[p].dimLabels {
if a.shape[d] == 1 && sizes[l] != 1 {
st[d] = 0
}
}
ops[p].dimStrides = st
}
outShape := make([]int, len(outLabels))
for i, l := range outLabels {
outShape[i] = sizes[l]
}
out := &Array{shape: outShape, dt: operands[0].Dtype()}
for _, o := range operands[1:] {
out.dt = promote(out.dt, o.Dtype())
}
total := 1
for _, d := range outShape {
total *= d
}
out.alloc(total)
outStrides := make([]int, len(outLabels))
run := 1
for i := len(outLabels) - 1; i >= 0; i-- {
outStrides[i] = run
run *= outShape[i]
}
// Label-position tables. sorted lists the label union ascending;
// depthOf maps a label to its position there and outAxisOf a label
// to its output axis, or -1 when the label is summed away. All hot
// lookups below are slice reads.
sorted := append([]int(nil), allLabels...)
slices.Sort(sorted)
nAll := len(sorted)
nOps := len(operands)
nOut := len(outLabels)
maxLabel := -1
for _, l := range allLabels {
if l > maxLabel {
maxLabel = l
}
}
depthOf := make([]int, maxLabel+1)
outAxisOf := make([]int, maxLabel+1)
for i := range outAxisOf {
outAxisOf[i] = -1
}
for i, l := range sorted {
depthOf[l] = i
}
for i, l := range outLabels {
outAxisOf[l] = i
}
// opDepth[p*nAll+d] is the payload offset operand p gains when the
// coordinate of sorted[d] advances by one: the sum of its
// dimensions' strides on that label, zero for broadcast axes and
// for labels the operand does not have. Repeated labels sum, which
// reproduces the diagonal extraction of the per-element walk.
opDepth := make([]int, nOps*nAll)
for p := range ops {
for d, l := range ops[p].dimLabels {
opDepth[p*nAll+depthOf[l]] += ops[p].dimStrides[d]
}
}
// The summed labels keep their ascending sorted order: index 0 is
// the outermost inner-loop axis, index nSum-1 the innermost, so the
// per-slot visit order matches the full walk with the output
// coordinates held fixed.
sumDepths := make([]int, 0, nAll)
for d, l := range sorted {
if outAxisOf[l] < 0 {
sumDepths = append(sumDepths, d)
}
}
nSum := len(sumDepths)
sumSize := make([]int, nSum)
sumTotal := 1
for t, d := range sumDepths {
sumSize[t] = sizes[sorted[d]]
sumTotal *= sumSize[t]
}
innerDelta := make([]int, nOps*nSum)
for p := range nOps {
for t, d := range sumDepths {
innerDelta[p*nSum+t] = opDepth[p*nAll+d]
}
}
outDelta := make([]int, nOps*nOut)
for p := range nOps {
for i, l := range outLabels {
outDelta[p*nOut+i] = opDepth[p*nAll+depthOf[l]]
}
}
// Readers resolve each operand to a dense payload slice indexed by
// the same logical offsets the per-element walk computed, widened
// exactly where its scalar read widened. Every conversion is
// deterministic, so widening once per call reads back the values
// the per-element widening produced.
// Each worker builds its own visitor, so the cursor and the odometer
// scratch below are per-worker: both are rewritten in full at every
// visit, which is what keeps the slots independent.
var newVisit func() func(slot int, base []int)
switch out.dt {
case Int:
// An Int result forces every operand to Int, read raw; a
// strided one is gathered through physIndex first, mirroring
// the complex branch below.
rv := make([][]int64, nOps)
for p, a := range operands {
if a.strides == nil {
rv[p] = a.ints
continue
}
w := make([]int64, a.Len())
einsumGather(w, func(i int) int64 { return a.ints[a.physIndex(i)] })
rv[p] = w
}
newVisit = func() func(slot int, base []int) {
off := make([]int, nOps)
coord := make([]int, nSum)
return func(slot int, base []int) {
einsumSlotSum(rv, out.ints, slot, base, off, coord, innerDelta, sumSize, sumTotal)
}
}
case Complex:
rv := make([][]complex128, nOps)
for p, a := range operands {
switch {
case a.dt == Complex && a.strides == nil:
rv[p] = a.complexes
case a.strides != nil:
// A strided view is gathered through the accessor: its
// payload is not the view's own element order.
w := make([]complex128, a.Len())
einsumGather(w, func(i int) complex128 { return a.complexAt(i) })
rv[p] = w
default:
// A real operand widens to complex(v, 0) element by
// element, exactly as the scalar read did.
rv[p] = complexPayload(a)
}
}
newVisit = func() func(slot int, base []int) {
off := make([]int, nOps)
coord := make([]int, nSum)
return func(slot int, base []int) {
einsumSlotSum(rv, out.complexes, slot, base, off, coord, innerDelta, sumSize, sumTotal)
}
}
default:
rv := make([][]float64, nOps)
for p, a := range operands {
rv[p] = einsumRealReader(a)
}
if out.dt == Float32 {
// Per-addend narrowing: each product lands through
// float32(float64(out[slot]) + acc), the exact write the
// walk performed per visit.
newVisit = func() func(slot int, base []int) {
off := make([]int, nOps)
coord := make([]int, nSum)
return func(slot int, base []int) {
einsumSlotSumF32(rv, out.floats32, slot, base, off, coord, innerDelta, sumSize, sumTotal)
}
}
} else {
newVisit = func() func(slot int, base []int) {
off := make([]int, nOps)
coord := make([]int, nSum)
return func(slot int, base []int) {
einsumSlotSum(rv, out.floats, slot, base, off, coord, innerDelta, sumSize, sumTotal)
}
}
}
}
einsumDriveSlots(total, sumTotal, nOut, outShape, outDelta, nOps, newVisit)
return out, nil
}
// einsumRealReader returns the operand's elements as float64, indexed
// by logical flat offset, widened exactly where the engine's scalar
// read widened: float32 and int from their raw payloads, a float64
// payload as is, a strided view gathered through the accessor. The
// widening splits across workers from widenParallelMin up: every
// element widens on its own into its own slot, so the buffer holds the
// serial walk's values.
func einsumRealReader(a *Array) []float64 {
if a.strides != nil {
n := a.Len()
w := make([]float64, n)
if n < widenParallelMin {
for i := range w {
w[i] = a.floatAt(i)
}
return w
}
engine.Parallel(n, func(ws, we int) {
for i := ws; i < we; i++ {
w[i] = a.floatAt(i)
}
})
return w
}
switch a.dt {
case Float32:
return einsumWiden(a.floats32)
case Int:
return einsumWiden(a.ints)
default:
return a.floats
}
}
// einsumWiden widens a raw payload to float64, exactly per element and
// across workers from widenParallelMin up, so the result reads back the
// serial walk's values.
func einsumWiden[T int64 | float32](src []T) []float64 {
w := make([]float64, len(src))
if len(src) < widenParallelMin {
for i, v := range src {
w[i] = float64(v)
}
return w
}
engine.Parallel(len(src), func(ws, we int) {
for i := ws; i < we; i++ {
w[i] = float64(src[i])
}
})
return w
}
// einsumGather fills dst by mapping every logical element index through
// read, exactly per element and across workers from widenParallelMin
// up, so the buffer reads back the serial walk's values.
func einsumGather[T any](dst []T, read func(int) T) {
if len(dst) < widenParallelMin {
for i := range dst {
dst[i] = read(i)
}
return
}
engine.Parallel(len(dst), func(ws, we int) {
for i := ws; i < we; i++ {
dst[i] = read(i)
}
})
}
// einsumSlotFloor is the number of per-slot visits a worker must carry
// before the slot walk is split across goroutines. A visit is a handful
// of multiply-adds and an odometer step, so a chunk below this costs
// more to spawn and synchronise than it runs: measured on a contraction
// of eight visits per slot, a worker carrying a thousand of them ran
// slower split than whole, one carrying two thousand broke even, and
// one carrying four thousand was ahead.
const einsumSlotFloor = 4096
// einsumDriveSlots walks the output slots in row-major order,
// maintaining per-operand base offsets through the output labels'
// strides, and hands each slot to the visitor built for the worker that
// owns it. The slots are split across workers as contiguous chunks,
// each walking its own cursor from its own first slot; which slot is
// written when is free to choose, and visit owns the summation order
// within the slot, so the chunking cannot move a bit.
func einsumDriveSlots(outTotal, sumTotal, nOut int, outShape, outDelta []int, nOps int, newVisit func() func(slot int, base []int)) {
if sumTotal == 0 || outTotal == 0 {
// A zero-sized summed label starves every slot, and a
// zero-sized output axis has no slot at all: the output stays
// zeroed, as the walk left it.
return
}
// The floor is per worker, so it converts into slots through the
// per-slot visit count; a chunk below it never splits.
perSlot := sumTotal * nOps
minSlots := 1
if perSlot < einsumSlotFloor {
minSlots = (einsumSlotFloor + perSlot - 1) / perSlot
}
engine.ParallelMin(outTotal, minSlots, func(ss, se int) {
visit := newVisit()
base := make([]int, nOps)
coord := make([]int, nOut)
// The chunk's own starting cursor: the walk below only advances
// offsets by strides, so it has to begin from the offset the
// first slot's coordinate already carries.
rem := ss
for i := nOut - 1; i >= 0; i-- {
coord[i] = rem % outShape[i]
rem /= outShape[i]
for p := range nOps {
base[p] += coord[i] * outDelta[p*nOut+i]
}
}
for slot := ss; slot < se; slot++ {
visit(slot, base)
for i := nOut - 1; i >= 0; i-- {
coord[i]++
for p := range nOps {
base[p] += outDelta[p*nOut+i]
}
if coord[i] < outShape[i] {
break
}
coord[i] = 0
for p := range nOps {
base[p] -= outDelta[p*nOut+i] * outShape[i]
}
}
}
})
}
// einsumSlotSum accumulates one output slot. It visits the summed
// labels' combinations in ascending-odometer order (index 0 outermost,
// the last innermost), multiplies the operands in order from the unit
// seed and adds each product into out[slot]. The innermost axis
// advances at every visit, so it is peeled into its own counted run:
// there only the operand cursors move, and the odometer above them
// steps once per run, which takes the digit bookkeeping off the hot
// path without moving a single addend. The run loop is written out per
// operand count up to three: the hot reads and cursor updates then hold
// their indices in registers, while the outer walk stays generic in
// nSum. The arithmetic of every branch is the walk's: seed, operand
// order, coordinate order and wrap bookkeeping included.
func einsumSlotSum[T int64 | float64 | complex128](rv [][]T, out []T, slot int, base, off, coord, delta, sizes []int, total int) {
nSum := len(sizes)
copy(off, base)
clear(coord[:nSum])
if nSum == 0 {
// Nothing is summed: the slot takes the product of the
// operands at the cursor, and a repeat of the walk would read
// the same elements again.
for range total {
acc := T(1)
for p := range rv {
acc *= rv[p][off[p]]
}
out[slot] += acc
}
return
}
// inner is the innermost axis' extent, runs the number of odometer
// combinations above it.
inner := sizes[nSum-1]
runs := 1
for t := range nSum - 1 {
runs *= sizes[t]
}
switch len(rv) {
case 1:
r0 := rv[0]
o0 := off[0]
d0 := delta[0*nSum:]
a0 := d0[nSum-1]
for range runs {
for range inner {
acc := T(1)
acc *= r0[o0]
out[slot] += acc
o0 += a0
}
o0 -= a0 * inner
for t := nSum - 2; t >= 0; t-- {
coord[t]++
o0 += d0[t]
if coord[t] < sizes[t] {
break
}
coord[t] = 0
o0 -= d0[t] * sizes[t]
}
}
off[0] = o0
case 2:
r0, r1 := rv[0], rv[1]
o0, o1 := off[0], off[1]
d0 := delta[0*nSum : 1*nSum]
d1 := delta[1*nSum : 2*nSum]
a0, a1 := d0[nSum-1], d1[nSum-1]
for range runs {
for range inner {
acc := T(1)
acc *= r0[o0]
acc *= r1[o1]
out[slot] += acc
o0 += a0
o1 += a1
}
o0 -= a0 * inner
o1 -= a1 * inner
for t := nSum - 2; t >= 0; t-- {
coord[t]++
o0 += d0[t]
o1 += d1[t]
if coord[t] < sizes[t] {
break
}
coord[t] = 0
o0 -= d0[t] * sizes[t]
o1 -= d1[t] * sizes[t]
}
}
off[0], off[1] = o0, o1
case 3:
r0, r1, r2 := rv[0], rv[1], rv[2]
o0, o1, o2 := off[0], off[1], off[2]
d0 := delta[0*nSum : 1*nSum]
d1 := delta[1*nSum : 2*nSum]
d2 := delta[2*nSum : 3*nSum]
a0, a1, a2 := d0[nSum-1], d1[nSum-1], d2[nSum-1]
for range runs {
for range inner {
acc := T(1)
acc *= r0[o0]
acc *= r1[o1]
acc *= r2[o2]
out[slot] += acc
o0 += a0
o1 += a1
o2 += a2
}
o0 -= a0 * inner
o1 -= a1 * inner
o2 -= a2 * inner
for t := nSum - 2; t >= 0; t-- {
coord[t]++
o0 += d0[t]
o1 += d1[t]
o2 += d2[t]
if coord[t] < sizes[t] {
break
}
coord[t] = 0
o0 -= d0[t] * sizes[t]
o1 -= d1[t] * sizes[t]
o2 -= d2[t] * sizes[t]
}
}
off[0], off[1], off[2] = o0, o1, o2
default:
for range runs {
for range inner {
acc := T(1)
for p := range rv {
acc *= rv[p][off[p]]
}
out[slot] += acc
for p := range rv {
off[p] += delta[p*nSum+nSum-1]
}
}
for p := range rv {
off[p] -= delta[p*nSum+nSum-1] * inner
}
for t := nSum - 2; t >= 0; t-- {
coord[t]++
for p := range rv {
off[p] += delta[p*nSum+t]
}
if coord[t] < sizes[t] {
break
}
coord[t] = 0
for p := range rv {
off[p] -= delta[p*nSum+t] * sizes[t]
}
}
}
}
}
// einsumSlotSumF32 is einsumSlotSum for a Float32 result: the products
// accumulate in float64 and each addend narrows into out[slot] on
// arrival, float32(float64(out[slot]) + acc), one narrowing per visit
// like the scalar walk. The innermost axis is peeled the same way, and
// the run loop is written out per operand count up to three, mirroring
// einsumSlotSum.
func einsumSlotSumF32(rv [][]float64, out []float32, slot int, base, off, coord, delta, sizes []int, total int) {
nSum := len(sizes)
copy(off, base)
clear(coord[:nSum])
if nSum == 0 {
for range total {
acc := 1.0
for p := range rv {
acc *= rv[p][off[p]]
}
out[slot] = float32(float64(out[slot]) + acc)
}
return
}
inner := sizes[nSum-1]
runs := 1
for t := range nSum - 1 {
runs *= sizes[t]
}
switch len(rv) {
case 1:
r0 := rv[0]
o0 := off[0]
d0 := delta[0*nSum:]
a0 := d0[nSum-1]
for range runs {
for range inner {
acc := 1.0
acc *= r0[o0]
out[slot] = float32(float64(out[slot]) + acc)
o0 += a0
}
o0 -= a0 * inner
for t := nSum - 2; t >= 0; t-- {
coord[t]++
o0 += d0[t]
if coord[t] < sizes[t] {
break
}
coord[t] = 0
o0 -= d0[t] * sizes[t]
}
}
off[0] = o0
case 2:
r0, r1 := rv[0], rv[1]
o0, o1 := off[0], off[1]
d0 := delta[0*nSum : 1*nSum]
d1 := delta[1*nSum : 2*nSum]
a0, a1 := d0[nSum-1], d1[nSum-1]
for range runs {
for range inner {
acc := 1.0
acc *= r0[o0]
acc *= r1[o1]
out[slot] = float32(float64(out[slot]) + acc)
o0 += a0
o1 += a1
}
o0 -= a0 * inner
o1 -= a1 * inner
for t := nSum - 2; t >= 0; t-- {
coord[t]++
o0 += d0[t]
o1 += d1[t]
if coord[t] < sizes[t] {
break
}
coord[t] = 0
o0 -= d0[t] * sizes[t]
o1 -= d1[t] * sizes[t]
}
}
off[0], off[1] = o0, o1
case 3:
r0, r1, r2 := rv[0], rv[1], rv[2]
o0, o1, o2 := off[0], off[1], off[2]
d0 := delta[0*nSum : 1*nSum]
d1 := delta[1*nSum : 2*nSum]
d2 := delta[2*nSum : 3*nSum]
a0, a1, a2 := d0[nSum-1], d1[nSum-1], d2[nSum-1]
for range runs {
for range inner {
acc := 1.0
acc *= r0[o0]
acc *= r1[o1]
acc *= r2[o2]
out[slot] = float32(float64(out[slot]) + acc)
o0 += a0
o1 += a1
o2 += a2
}
o0 -= a0 * inner
o1 -= a1 * inner
o2 -= a2 * inner
for t := nSum - 2; t >= 0; t-- {
coord[t]++
o0 += d0[t]
o1 += d1[t]
o2 += d2[t]
if coord[t] < sizes[t] {
break
}
coord[t] = 0
o0 -= d0[t] * sizes[t]
o1 -= d1[t] * sizes[t]
o2 -= d2[t] * sizes[t]
}
}
off[0], off[1], off[2] = o0, o1, o2
default:
for range runs {
for range inner {
acc := 1.0
for p := range rv {
acc *= rv[p][off[p]]
}
out[slot] = float32(float64(out[slot]) + acc)
for p := range rv {
off[p] += delta[p*nSum+nSum-1]
}
}
for p := range rv {
off[p] -= delta[p*nSum+nSum-1] * inner
}
for t := nSum - 2; t >= 0; t-- {
coord[t]++
for p := range rv {
off[p] += delta[p*nSum+t]
}
if coord[t] < sizes[t] {
break
}
coord[t] = 0
for p := range rv {
off[p] -= delta[p*nSum+t] * sizes[t]
}
}
}
}
}