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

245 lines
7.3 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import "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
}
}