Files
tensor/grad/broadcast_backward_pins_test.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

146 lines
3.9 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
// Regression pins for the broadcast and shape backwards: a broadcast
// gradient must collapse to its source's shape, and every widened op
// must hand each operand its own correctly shaped gradient buffer.
// TestBroadcastRank1Gradient pins the rank-1 broadcast backward: the
// gradient of a size-1 source broadcast to length m must collapse back
// to a single sum, not arrive with the broadcast shape.
func TestBroadcastRank1Gradient(t *testing.T) {
x, err := core.FromFloats([]float64{0}, 1)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(x, true)
y, err := xt.BroadcastTo(3)
if err != nil {
t.Fatalf("BroadcastTo: %v", err)
}
w, err := core.FromFloats([]float64{1, 2, 3}, 3)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
wt := FromArray(w, false)
prod, err := y.Mul(wt)
if err != nil {
t.Fatalf("Mul: %v", err)
}
loss, err := prod.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
g := xt.Grad()
if g.NDim() != 1 || g.Shape()[0] != 1 {
t.Fatalf("gradient shape %v, want [1]", g.Shape())
}
if math.Abs(g.FloatAt(0)-6) > 1e-12 {
t.Fatalf("gradient = %g, want 6", g.FloatAt(0))
}
}
// TestBroadcastRank1ToMatrix pins the (1,) to (m, n) broadcast backward
// against central differences.
func TestBroadcastRank1ToMatrix(t *testing.T) {
x, err := core.FromFloats([]float64{2}, 1)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(x, true)
y, err := xt.BroadcastTo(2, 3)
if err != nil {
t.Fatalf("BroadcastTo: %v", err)
}
w, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
wt := FromArray(w, false)
prod, err := y.Mul(wt)
if err != nil {
t.Fatalf("Mul: %v", err)
}
loss, err := prod.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
g := xt.Grad()
if g.NDim() != 1 || g.Shape()[0] != 1 {
t.Fatalf("gradient shape %v, want [1]", g.Shape())
}
if math.Abs(g.FloatAt(0)-21) > 1e-12 {
t.Fatalf("gradient = %g, want 21", g.FloatAt(0))
}
}
// TestL2NormAxisEmptyDim pins the backward against the integer division
// by zero an empty reduced dimension used to hit.
func TestL2NormAxisEmptyDim(t *testing.T) {
x, err := core.FromFloats([]float64{}, 3, 0)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
xt := FromArray(x, true)
n, err := xt.L2NormAxis(1)
if err != nil {
t.Fatalf("L2NormAxis: %v", err)
}
loss, err := n.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
g := xt.Grad()
if g.Len() != 0 {
t.Fatalf("gradient length %d, want 0", g.Len())
}
}
// TestShapeOpsRejectNonFloat pins the dtype contract on the shape ops
// that used to record graph nodes without validating the dtype.
func TestShapeOpsRejectNonFloat(t *testing.T) {
i, err := core.FromInts([]int64{1, 2, 3, 4}, 2, 2)
if err != nil {
t.Fatalf("FromInts: %v", err)
}
it := FromArray(i, true)
if _, err := it.Transpose(); err == nil {
t.Error("Transpose accepted an int tensor")
}
if _, err := it.Squeeze(0); err == nil {
t.Error("Squeeze accepted an int tensor")
}
if _, err := it.Unsqueeze(0); err == nil {
t.Error("Unsqueeze accepted an int tensor")
}
if _, err := it.Reshape(4); err == nil {
t.Error("Reshape accepted an int tensor")
}
if _, err := it.TransposeAxes(1, 0); err == nil {
t.Error("TransposeAxes accepted an int tensor")
}
if _, err := it.Floor(); err == nil {
t.Error("Floor accepted an int tensor")
}
if _, err := it.BroadcastTo(2, 3); err == nil {
t.Error("BroadcastTo accepted an int tensor")
}
}