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