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