146 lines
3.9 KiB
Go
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")
|
|
}
|
|
}
|