Files
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

1654 lines
49 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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]
}
}
}
}
}