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