Files
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

235 lines
7.0 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"strings"
"testing"
)
func TestBroadcastTo(t *testing.T) {
a := mustFromInts(t, []int64{1, 2, 3}, 3, 1)
b, err := BroadcastTo(a, 3, 2)
if err != nil {
t.Fatalf("BroadcastTo: %v", err)
}
want := mustFromInts(t, []int64{1, 1, 2, 2, 3, 3}, 3, 2)
if !Equal(want, b) {
t.Fatalf("BroadcastTo: %s", b)
}
// Prepending a leading dimension.
c, err := BroadcastTo(a, 2, 3, 1)
if err != nil {
t.Fatalf("BroadcastTo prepend: %v", err)
}
if c.NDim() != 3 || c.Shape()[0] != 2 {
t.Fatalf("BroadcastTo prepend shape: %v", c.Shape())
}
if v, _ := IntAt(c, 1, 2, 0); v != 3 {
t.Fatalf("BroadcastTo prepend value: %d", v)
}
// An identical shape is a copy.
same, err := BroadcastTo(a, 3, 1)
if err != nil || !Equal(a, same) {
t.Fatalf("BroadcastTo same: %s %v", same, err)
}
if _, err := BroadcastTo(a, 4, 1); err == nil || !strings.Contains(err.Error(), "cannot broadcast") {
t.Fatalf("BroadcastTo incompatible: %v", err)
}
tall := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2)
if _, err := BroadcastTo(tall, 2, 3); err == nil || !strings.Contains(err.Error(), "cannot broadcast") {
t.Fatalf("BroadcastTo rank-2 mismatch: %v", err)
}
if _, err := BroadcastTo(a, -1); err == nil {
t.Fatalf("BroadcastTo negative must error")
}
}
func TestBroadcastWith(t *testing.T) {
a := mustFromInts(t, []int64{10, 20, 30}, 3, 1)
b := mustFromInts(t, []int64{1, 2}, 1, 2)
aw, bw, err := BroadcastWith(a, b)
if err != nil {
t.Fatalf("BroadcastWith: %v", err)
}
if shape := aw.Shape(); shape[0] != 3 || shape[1] != 2 {
t.Fatalf("BroadcastWith shape: %v", shape)
}
// (3,2): a column repeats across, b row repeats down.
wantA := mustFromInts(t, []int64{10, 10, 20, 20, 30, 30}, 3, 2)
wantB := mustFromInts(t, []int64{1, 2, 1, 2, 1, 2}, 3, 2)
if !Equal(wantA, aw) || !Equal(wantB, bw) {
t.Fatalf("BroadcastWith: %s / %s", aw, bw)
}
// The broadcast operands now satisfy the strict element-wise ops.
sum, err := Add(aw, bw)
if err != nil {
t.Fatalf("Add after broadcast: %v", err)
}
wantSum := mustFromInts(t, []int64{11, 12, 21, 22, 31, 32}, 3, 2)
if !Equal(wantSum, sum) {
t.Fatalf("Add after broadcast: %s", sum)
}
x := mustFromInts(t, []int64{1, 2}, 2)
y := mustFromInts(t, []int64{1, 2, 3}, 3)
if _, _, err := BroadcastWith(x, y); err == nil || !strings.Contains(err.Error(), "do not meet") {
t.Fatalf("BroadcastWith incompatible: %v", err)
}
}
// broadcastReference is a deliberately naive row-major walk: for every
// target position it derives the source coordinate, so it shares no code
// with the run-fill implementation. Any disagreement is a bug in one of
// them.
func broadcastReference(a *Array, shape []int) *Array {
total := 1
for _, d := range shape {
total *= d
}
out := &Array{shape: append([]int(nil), shape...), dt: a.Dtype()}
out.alloc(total)
coord := make([]int, len(shape))
for i := range total {
src := 0
for d := range a.Shape() {
td := len(shape) - len(a.Shape()) + d
c := coord[td]
if a.Shape()[d] == 1 {
c = 0
}
src = src*a.Shape()[d] + c
}
out.setFrom(i, a, src)
advanceOdometer(coord, shape)
// advanceOdometer is the helper under test elsewhere; the
// independent rebuild above is what makes this a reference.
}
return out
}
// TestBroadcastToMatchesReference drives the run-fill rewrite against
// the naive walk over the shape families it specialises: a constant
// trailing run (bias layouts), a prefixed rank, a size-1 dim mid-shape,
// and a full no-op broadcast.
func TestBroadcastToMatchesReference(t *testing.T) {
for _, tc := range []struct {
src []int
dst []int
name string
}{
{[]int{3}, []int{3}, "identity"},
{[]int{1, 3, 1, 1}, []int{2, 3, 4, 5}, "bias NCHW"},
{[]int{1, 4}, []int{6, 4}, "bias NC"},
{[]int{2, 3}, []int{5, 2, 3}, "prefix prepend"},
{[]int{2, 1, 3}, []int{2, 4, 3}, "middle size-1"},
{[]int{1, 1, 1}, []int{2, 3, 4}, "scalar-ish"},
{[]int{5, 1}, []int{5, 7}, "trailing replicate"},
{[]int{1}, []int{4, 1}, "single prepend"},
{[]int{2, 3, 1}, []int{1, 2, 3, 6}, "mixed"},
} {
n := 1
for _, d := range tc.src {
n *= d
}
vals := make([]float64, n)
for i := range vals {
vals[i] = float64(i)*1.5 - 3
}
src, err := FromFloats(vals, tc.src...)
if err != nil {
t.Fatalf("%s: %v", tc.name, err)
}
got, err := BroadcastTo(src, tc.dst...)
if err != nil {
t.Fatalf("%s: %v", tc.name, err)
}
want := broadcastReference(src, tc.dst)
if got.Len() != want.Len() {
t.Fatalf("%s: length %d, want %d", tc.name, got.Len(), want.Len())
}
for i := range want.Len() {
if got.FloatAt(i) != want.FloatAt(i) {
t.Fatalf("%s: element %d = %v, want %v", tc.name, i, got.FloatAt(i), want.FloatAt(i))
}
}
}
}
// TestBroadcastToSplitMatchesSerial pins the parallel outer walk's chunk
// cursor. The fill splits the outer positions across workers and each
// worker rebuilds the source coordinate of its own first position, so
// the result must not depend on where the chunk boundaries fall. Every
// outer position below reads a different source element, and the split
// is asserted before it is compared so a shrunken workload cannot
// quietly fall back to the single-chunk fill.
func TestBroadcastToSplitMatchesSerial(t *testing.T) {
const workers = 4
cases := []struct {
src, target []int
run int // trailing target elements one outer position fills
}{
{[]int{4096, 1}, []int{4096, 8}, 8},
{[]int{512, 1, 1}, []int{512, 4, 4}, 16},
}
prev := NumWorkers()
defer SetNumCPU(prev)
for _, tc := range cases {
n := 1
for _, d := range tc.src {
n *= d
}
vals := make([]float64, n)
for i := range vals {
vals[i] = float64(i) + 0.5
}
src := mustFromFloats(t, vals, tc.src...)
outer := 1
for _, d := range tc.target {
outer *= d
}
outer /= tc.run
parMin := max(1, broadcastMinPerWorker/tc.run)
if chunk := (outer + workers - 1) / workers; chunk < parMin {
t.Fatalf("src %v to %v: %d outer positions no longer split at %d workers: chunk %d below the %d-element floor",
tc.src, tc.target, outer, workers, chunk, parMin)
}
SetNumCPU(1)
want, err := BroadcastTo(src, tc.target...)
if err != nil {
t.Fatalf("BroadcastTo %v: %v", tc.target, err)
}
SetNumCPU(workers)
got, err := BroadcastTo(src, tc.target...)
if err != nil {
t.Fatalf("BroadcastTo %v split: %v", tc.target, err)
}
wf, gf := want.RawFloats(), got.RawFloats()
if len(gf) != len(wf) {
t.Fatalf("BroadcastTo %v: split result has %d elements, the serial fill %d",
tc.target, len(gf), len(wf))
}
for i := range wf {
if gf[i] != wf[i] {
t.Fatalf("BroadcastTo %v: split element %d = %v, the serial fill %v",
tc.target, i, gf[i], wf[i])
}
}
// The last outer position is the one a moved chunk boundary
// reads from the wrong source slot: it must be its own row.
if last := float64(outer-1) + 0.5; gf[len(gf)-1] != last {
t.Fatalf("BroadcastTo %v: last element = %v, want %v",
tc.target, gf[len(gf)-1], last)
}
}
}