245 lines
7.3 KiB
Go
245 lines
7.3 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
|
|
// SPDX-License-Identifier: MIT
|
|||
|
|
|
|||
|
|
package core
|
|||
|
|
|
|||
|
|
import "sourcedock.dev/petrbalvin/tensor/internal/engine"
|
|||
|
|
|
|||
|
|
// Explicit broadcasting. Silent broadcasting is banned; the
|
|||
|
|
// escape hatch is a named operation: BroadcastTo expands one array to a
|
|||
|
|
// target shape, BroadcastWith expands two operands to their common shape.
|
|||
|
|
// Both copy: results never alias their receiver.
|
|||
|
|
|
|||
|
|
// broadcastMinPerWorker is the element count every spawned broadcast
|
|||
|
|
// worker must carry. Its inner loop is one run fill or one element write
|
|||
|
|
// per step, cheaper per element than an arithmetic map, but a broadcast
|
|||
|
|
// below a few thousand elements still measures several times faster on
|
|||
|
|
// the calling goroutine than fanned out: a 256-element expansion paid
|
|||
|
|
// more in thirty-two spawns than the whole fill cost.
|
|||
|
|
const broadcastMinPerWorker = 1024
|
|||
|
|
|
|||
|
|
// BroadcastTo returns a new array expanded to the target shape: size-1
|
|||
|
|
// dimensions replicate, missing leading dimensions prepend, and
|
|||
|
|
// anything else is a loud error naming both shapes.
|
|||
|
|
func BroadcastTo(a *Array, shape ...int) (*Array, error) {
|
|||
|
|
target, _, err := checkedDims(shape)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
if err := broadcastable(a.shape, shape); err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
sh := make([]int, len(shape))
|
|||
|
|
copy(sh, shape)
|
|||
|
|
out := &Array{shape: sh, dt: a.dt}
|
|||
|
|
out.alloc(target)
|
|||
|
|
if target == 0 {
|
|||
|
|
return out, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// BroadcastTo used to walk the target one element at a time: every
|
|||
|
|
// element recomputed its source index over the whole rank and took
|
|||
|
|
// an odometer step. Profiling the MLP training step put that walk at
|
|||
|
|
// the top of the profile, so the output is filled in runs instead.
|
|||
|
|
//
|
|||
|
|
// For each target dimension, srcStride is the stride the source
|
|||
|
|
// advances by (0 where the source dimension is 1 or prepended). A
|
|||
|
|
// trailing block of zero strides means the source element is
|
|||
|
|
// constant across the whole block, so each run is a plain fill: the
|
|||
|
|
// bias-shaped broadcast (1, C, 1, 1) to (N, C, H, W) collapses to one
|
|||
|
|
// fill per channel rather than an odometer step per cell.
|
|||
|
|
off := len(shape) - len(a.shape)
|
|||
|
|
srcStride := make([]int, len(shape))
|
|||
|
|
srcDense := denseStrides(a.shape)
|
|||
|
|
for d := range shape {
|
|||
|
|
if d >= off && a.shape[d-off] != 1 {
|
|||
|
|
srcStride[d] = srcDense[d-off]
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
tail := 0
|
|||
|
|
for d := len(shape) - 1; d >= 0 && srcStride[d] == 0; d-- {
|
|||
|
|
tail++
|
|||
|
|
}
|
|||
|
|
run := 1
|
|||
|
|
for d := len(shape) - tail; d < len(shape); d++ {
|
|||
|
|
run *= shape[d]
|
|||
|
|
}
|
|||
|
|
outerShape := shape[:len(shape)-tail]
|
|||
|
|
outer := target / run
|
|||
|
|
|
|||
|
|
// The target offsets are disjoint, so the walk parallelises.
|
|||
|
|
// A worker needs broadcastMinPerWorker output elements before its
|
|||
|
|
// spawn pays for itself, and one outer step writes a whole run, so
|
|||
|
|
// the item floor is the element floor over the run length.
|
|||
|
|
parMin := max(1, broadcastMinPerWorker/run)
|
|||
|
|
engine.ParallelMin(outer, parMin, func(s, e int) {
|
|||
|
|
coord := make([]int, len(outerShape))
|
|||
|
|
src := 0
|
|||
|
|
rem := s
|
|||
|
|
for d := len(outerShape) - 1; d >= 0; d-- {
|
|||
|
|
if outerShape[d] == 0 {
|
|||
|
|
coord[d] = 0
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
c := rem % outerShape[d]
|
|||
|
|
rem /= outerShape[d]
|
|||
|
|
coord[d] = c
|
|||
|
|
src += c * srcStride[d]
|
|||
|
|
}
|
|||
|
|
if run == 1 {
|
|||
|
|
// No constant run to fill; keep the plain element write so a
|
|||
|
|
// pure prefix broadcast does not pay for the run machinery.
|
|||
|
|
for o := s; o < e; o++ {
|
|||
|
|
out.setFrom(o, a, src)
|
|||
|
|
src += advanceStride(coord, outerShape, srcStride)
|
|||
|
|
}
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
for o := s; o < e; o++ {
|
|||
|
|
out.fillRun(o*run, run, a, src)
|
|||
|
|
src += advanceStride(coord, outerShape, srcStride)
|
|||
|
|
}
|
|||
|
|
})
|
|||
|
|
return out, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// advanceStride steps the odometer one position and returns the change
|
|||
|
|
// in Σ coord[d]·stride[d], so a caller can carry the source index
|
|||
|
|
// forward instead of rebuilding it from the coordinates each element.
|
|||
|
|
func advanceStride(coord, shape, stride []int) int {
|
|||
|
|
delta := 0
|
|||
|
|
for d := len(coord) - 1; d >= 0; d-- {
|
|||
|
|
coord[d]++
|
|||
|
|
if coord[d] < shape[d] {
|
|||
|
|
return delta + stride[d]
|
|||
|
|
}
|
|||
|
|
// Rolled from shape[d]−1 back to 0: that dim's whole
|
|||
|
|
// contribution leaves the sum.
|
|||
|
|
delta -= (shape[d] - 1) * stride[d]
|
|||
|
|
coord[d] = 0
|
|||
|
|
}
|
|||
|
|
return delta
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// fillRun writes n copies of element srcFlat of src starting at dst of a.
|
|||
|
|
// The destination is freshly allocated and contiguous; the source may be
|
|||
|
|
// a view.
|
|||
|
|
func (a *Array) fillRun(dst, n int, src *Array, srcFlat int) {
|
|||
|
|
p := src.physIndex(srcFlat)
|
|||
|
|
switch a.dt {
|
|||
|
|
case Int:
|
|||
|
|
fillRepeat(a.ints[dst:dst+n], src.ints[p])
|
|||
|
|
case Bool:
|
|||
|
|
fillRepeat(a.bools[dst:dst+n], src.bools[p])
|
|||
|
|
case Int8:
|
|||
|
|
fillRepeat(a.i8s[dst:dst+n], src.i8s[p])
|
|||
|
|
case Uint8:
|
|||
|
|
fillRepeat(a.u8s[dst:dst+n], src.u8s[p])
|
|||
|
|
case Int16:
|
|||
|
|
fillRepeat(a.i16s[dst:dst+n], src.i16s[p])
|
|||
|
|
case Uint16:
|
|||
|
|
fillRepeat(a.u16s[dst:dst+n], src.u16s[p])
|
|||
|
|
case Int32:
|
|||
|
|
fillRepeat(a.i32s[dst:dst+n], src.i32s[p])
|
|||
|
|
case Uint32:
|
|||
|
|
fillRepeat(a.u32s[dst:dst+n], src.u32s[p])
|
|||
|
|
case Float16:
|
|||
|
|
fillRepeat(a.halves[dst:dst+n], src.halves[p])
|
|||
|
|
case Float32:
|
|||
|
|
fillRepeat(a.floats32[dst:dst+n], src.floats32[p])
|
|||
|
|
case Float:
|
|||
|
|
fillRepeat(a.floats[dst:dst+n], src.floats[p])
|
|||
|
|
default:
|
|||
|
|
fillRepeat(a.complexes[dst:dst+n], src.complexes[p])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// fillRepeat writes v into every slot of dst. The first element seeds the
|
|||
|
|
// run and each further round doubles the stretched prefix with one copy,
|
|||
|
|
// so a long constant run is a logarithmic number of memmoves rather than
|
|||
|
|
// one store per element. Source and destination never overlap in the
|
|||
|
|
// direction of travel: the copied prefix ends exactly where the
|
|||
|
|
// destination starts.
|
|||
|
|
func fillRepeat[T any](dst []T, v T) {
|
|||
|
|
if len(dst) == 0 {
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
dst[0] = v
|
|||
|
|
for i := 1; i < len(dst); i *= 2 {
|
|||
|
|
copy(dst[i:], dst[:i])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// BroadcastWith broadcasts both operands to their common shape, so
|
|||
|
|
// `aw, bw, err := tensor.BroadcastWith(a, b)` can precede an ordinary
|
|||
|
|
// element-wise operation.
|
|||
|
|
func BroadcastWith(a, b *Array) (*Array, *Array, error) {
|
|||
|
|
common, err := commonShape(a.shape, b.shape)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, nil, err
|
|||
|
|
}
|
|||
|
|
aw, err := BroadcastTo(a, common...)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, nil, err
|
|||
|
|
}
|
|||
|
|
bw, err := BroadcastTo(b, common...)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, nil, err
|
|||
|
|
}
|
|||
|
|
return aw, bw, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// broadcastable reports whether from can broadcast to target, erroring
|
|||
|
|
// with both shapes otherwise.
|
|||
|
|
func broadcastable(from, target []int) error {
|
|||
|
|
if len(from) > len(target) {
|
|||
|
|
return errf("BroadcastTo: cannot broadcast %s to %s", shapeText(from), shapeText(target))
|
|||
|
|
}
|
|||
|
|
off := len(target) - len(from)
|
|||
|
|
for d := range from {
|
|||
|
|
if from[d] != 1 && from[d] != target[off+d] {
|
|||
|
|
return errf("BroadcastTo: cannot broadcast %s to %s", shapeText(from), shapeText(target))
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// commonShape returns the shape two arrays broadcast to together.
|
|||
|
|
func commonShape(a, b []int) ([]int, error) {
|
|||
|
|
n := max(len(b), len(a))
|
|||
|
|
out := make([]int, n)
|
|||
|
|
for d := range n {
|
|||
|
|
av, bv := 1, 1
|
|||
|
|
if d >= n-len(a) {
|
|||
|
|
av = a[d-(n-len(a))]
|
|||
|
|
}
|
|||
|
|
if d >= n-len(b) {
|
|||
|
|
bv = b[d-(n-len(b))]
|
|||
|
|
}
|
|||
|
|
switch {
|
|||
|
|
case av == bv:
|
|||
|
|
out[d] = av
|
|||
|
|
case av == 1:
|
|||
|
|
out[d] = bv
|
|||
|
|
case bv == 1:
|
|||
|
|
out[d] = av
|
|||
|
|
default:
|
|||
|
|
return nil, errf("BroadcastWith: shapes %s and %s do not meet",
|
|||
|
|
shapeText(a), shapeText(b))
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return out, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// advanceOdometer increments a row-major coordinate over shape.
|
|||
|
|
func advanceOdometer(coord, shape []int) {
|
|||
|
|
for d := len(coord) - 1; d >= 0; d-- {
|
|||
|
|
coord[d]++
|
|||
|
|
if coord[d] < shape[d] {
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
coord[d] = 0
|
|||
|
|
}
|
|||
|
|
}
|