// Copyright (c) 2026 Petr Balvín (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] } } } } }