1654 lines
49 KiB
Go
1654 lines
49 KiB
Go
// 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]
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|