feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,234 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user