344 lines
8.4 KiB
Go
344 lines
8.4 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"
|
||
"strings"
|
||
"testing"
|
||
)
|
||
|
||
func mustTensor(t *testing.T, vals []float64, shape ...int) *Tensor {
|
||
t.Helper()
|
||
tt, err := FromFloat64s(vals, true, shape...)
|
||
if err != nil {
|
||
t.Fatalf("FromFloat64s(%v, %v): %v", vals, shape, err)
|
||
}
|
||
return tt
|
||
}
|
||
|
||
func TestAutogradBasicChain(t *testing.T) {
|
||
x := mustTensor(t, []float64{2, 3}, 2)
|
||
y := mustTensor(t, []float64{4, 5}, 2)
|
||
z, err := x.Add(y)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
sq, err := z.Mul(z)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
s, err := sq.Sum()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := s.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// d/dx (x+y)^2 summed = 2(x+y); at x=2: 12, at x=3: 16.
|
||
if gx := x.Grad().FloatAt(0); math.Abs(gx-12) > 1e-9 {
|
||
t.Errorf("grad x[0]: got %v, want 12", gx)
|
||
}
|
||
if gx := x.Grad().FloatAt(1); math.Abs(gx-16) > 1e-9 {
|
||
t.Errorf("grad x[1]: got %v, want 16", gx)
|
||
}
|
||
// y's gradient matches x's, symmetric in the sum.
|
||
if gy := y.Grad().FloatAt(0); math.Abs(gy-12) > 1e-9 {
|
||
t.Errorf("grad y[0]: got %v, want 12", gy)
|
||
}
|
||
}
|
||
|
||
func TestAutogradMatMul(t *testing.T) {
|
||
a := mustTensor(t, []float64{1, 2, 3, 4}, 2, 2)
|
||
b := mustTensor(t, []float64{5, 6, 7, 8}, 2, 2)
|
||
p, err := a.MatMul(b)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
s, err := p.Sum()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := s.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// d/dA sum(A·B) = J·Bᵀ; Bᵀ = [[5,7],[6,8]], so J·Bᵀ =
|
||
// [[11,15],[11,15]] (each row is the column sums of Bᵀ).
|
||
wantA := []float64{11, 15, 11, 15}
|
||
for i := range 4 {
|
||
if g := a.Grad().FloatAt(i); math.Abs(g-wantA[i]) > 1e-9 {
|
||
t.Errorf("grad A[%d]: got %v, want %v", i, g, wantA[i])
|
||
}
|
||
}
|
||
// d/dB sum(A·B) = Aᵀ·J; Aᵀ = [[1,3],[2,4]], row sums: 4, 6, so
|
||
// Aᵀ·J = [[4,4],[6,6]].
|
||
wantB := []float64{4, 4, 6, 6}
|
||
for i := range 4 {
|
||
if g := b.Grad().FloatAt(i); math.Abs(g-wantB[i]) > 1e-9 {
|
||
t.Errorf("grad B[%d]: got %v, want %v", i, g, wantB[i])
|
||
}
|
||
}
|
||
}
|
||
|
||
func TestAutogradMatMulVector(t *testing.T) {
|
||
// 2-D × 1-D: y = A·x.
|
||
a := mustTensor(t, []float64{1, 2, 3, 4}, 2, 2)
|
||
x := mustTensor(t, []float64{2, 3}, 2)
|
||
y, err := a.MatMul(x)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
s, err := y.Sum()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := s.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// grad x = Aᵀ·1 = column sums of A: 4, 6.
|
||
if g := x.Grad().FloatAt(0); math.Abs(g-4) > 1e-9 {
|
||
t.Errorf("grad x[0]: got %v, want 4", g)
|
||
}
|
||
if g := x.Grad().FloatAt(1); math.Abs(g-6) > 1e-9 {
|
||
t.Errorf("grad x[1]: got %v, want 6", g)
|
||
}
|
||
// grad A = outer(1, x): [[2,3],[2,3]].
|
||
if g := a.Grad().FloatAt(2); math.Abs(g-2) > 1e-9 {
|
||
t.Errorf("grad A[2]: got %v, want 2", g)
|
||
}
|
||
}
|
||
|
||
func TestAutogradActivations(t *testing.T) {
|
||
// Sigmoid at 0: σ(0)=0.5, σ' = 0.25.
|
||
x3 := mustTensor(t, []float64{0}, 1)
|
||
sg, _ := x3.Sigmoid()
|
||
if err := sg.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if g := x3.Grad().FloatAt(0); math.Abs(g-0.25) > 1e-9 {
|
||
t.Errorf("Sigmoid grad at 0: got %v, want 0.25", g)
|
||
}
|
||
|
||
// Exp and Log compose to identity: grad log(exp(x)) = 1.
|
||
x4 := mustTensor(t, []float64{2}, 1)
|
||
e, _ := x4.Exp()
|
||
l, _ := e.Log()
|
||
if err := l.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if g := x4.Grad().FloatAt(0); math.Abs(g-1) > 1e-9 {
|
||
t.Errorf("log(exp) grad: got %v, want 1", g)
|
||
}
|
||
}
|
||
|
||
func TestAutogradGradientAccumulation(t *testing.T) {
|
||
x := mustTensor(t, []float64{1}, 1)
|
||
a, _ := x.Mul(x)
|
||
b, _ := x.Mul(x)
|
||
s, err := a.Add(b)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := s.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// d/dx (x² + x²) at x=1 = 4.
|
||
if g := x.Grad().FloatAt(0); math.Abs(g-4) > 1e-9 {
|
||
t.Errorf("accumulated grad: got %v, want 4", g)
|
||
}
|
||
// A second Backward without ZeroGrad accumulates into the leaf.
|
||
if err := s.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if g := x.Grad().FloatAt(0); math.Abs(g-8) > 1e-9 {
|
||
t.Errorf("accumulated grad after 2nd pass: got %v, want 8", g)
|
||
}
|
||
x.ZeroGrad()
|
||
if x.Grad() != nil {
|
||
t.Error("ZeroGrad did not clear the gradient")
|
||
}
|
||
}
|
||
|
||
func TestAutogradRejectsNonFloat(t *testing.T) {
|
||
i, err := core.FromInts([]int64{1, 2}, 2)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
it := FromArray(i, true)
|
||
if _, err := it.Sum(); err == nil || !strings.Contains(err.Error(), "float") {
|
||
t.Errorf("int Sum: %v", err)
|
||
}
|
||
c, _ := core.FromComplexes([]complex128{1 + 2i}, 1)
|
||
ct := FromArray(c, true)
|
||
// Complex Exp is differentiable (the Wirtinger graph); the
|
||
// real-only kernels are the ones that must still refuse it.
|
||
if _, err := ct.Exp(); err != nil {
|
||
t.Errorf("complex Exp must differentiate: %v", err)
|
||
}
|
||
if _, err := ct.Log(); err == nil {
|
||
t.Error("complex Log must error")
|
||
}
|
||
if _, err := ct.Tanh(); err == nil {
|
||
t.Error("complex Tanh must error")
|
||
}
|
||
}
|
||
|
||
func TestAutogradDivTanhNeg(t *testing.T) {
|
||
// d/dx (x/y) at x=4, y=2 = 1/2.
|
||
x := mustTensor(t, []float64{4}, 1)
|
||
y := mustTensor(t, []float64{2}, 1)
|
||
q, err := x.Div(y)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := q.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if g := x.Grad().FloatAt(0); math.Abs(g-0.5) > 1e-9 {
|
||
t.Errorf("Div grad x: got %v, want 0.5", g)
|
||
}
|
||
// d/dy (x/y) at y=2 = -x/y² = -1.
|
||
if g := y.Grad().FloatAt(0); math.Abs(g+1) > 1e-9 {
|
||
t.Errorf("Div grad y: got %v, want -1", g)
|
||
}
|
||
|
||
// tanh'(0) = 1.
|
||
t0 := mustTensor(t, []float64{0}, 1)
|
||
th, _ := t0.Tanh()
|
||
if err := th.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if g := t0.Grad().FloatAt(0); math.Abs(g-1) > 1e-9 {
|
||
t.Errorf("Tanh grad at 0: got %v, want 1", g)
|
||
}
|
||
|
||
// d/dx (-x) = -1.
|
||
n := mustTensor(t, []float64{3}, 1)
|
||
neg, _ := n.Neg()
|
||
if err := neg.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if g := n.Grad().FloatAt(0); g != -1 {
|
||
t.Errorf("Neg grad: got %v, want -1", g)
|
||
}
|
||
|
||
// Accessors and Mean grad.
|
||
m := mustTensor(t, []float64{1, 2, 3, 4}, 2, 2)
|
||
if m.Data() != m.Data() || m.RequiresGrad() != true {
|
||
t.Error("accessors wrong")
|
||
}
|
||
mean, err := m.Mean()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := mean.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// d/dx mean(x) = 1/n = 1/4.
|
||
if g := m.Grad().FloatAt(0); math.Abs(g-0.25) > 1e-9 {
|
||
t.Errorf("Mean grad: got %v, want 0.25", g)
|
||
}
|
||
if m.Grad() == nil {
|
||
t.Error("Grad() must be non-nil after Backward")
|
||
}
|
||
}
|
||
|
||
func TestAutogradReshape(t *testing.T) {
|
||
x, err := FromFloat64s([]float64{1, 2, 3, 4}, true, 2, 2)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
r, err := x.Reshape(4)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
sumT, err := r.Sum()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := sumT.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
g, err := x.Grad().Elements[float64]()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
for i := range g {
|
||
if g[i] != 1 {
|
||
t.Fatalf("Reshape grad[%d]: %v, want 1", i, g[i])
|
||
}
|
||
}
|
||
}
|
||
|
||
func TestAutogradTransposeBackward(t *testing.T) {
|
||
x, err := FromFloat64s([]float64{1, 2, 3, 4}, true, 2, 2)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
tr, err := x.Transpose()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
s, err := tr.Sum()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := s.Backward(); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
g, err := x.Grad().Elements[float64]()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
for i := range g {
|
||
if g[i] != 1 {
|
||
t.Errorf("Transpose grad[%d]: %v, want 1", i, g[i])
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestAutogradPowZeroGradient pins the exponent-0 backward: d/dx x⁰
|
||
// is the zero gradient everywhere, including at x = 0 where the
|
||
// chain rule would evaluate 0·∞ and produce NaN.
|
||
func TestAutogradPowZeroGradient(t *testing.T) {
|
||
x := mustTensor(t, []float64{0, 2}, 2)
|
||
y, err := x.Pow(0)
|
||
if err != nil {
|
||
t.Fatalf("Pow(0): %v", err)
|
||
}
|
||
loss, err := y.Sum()
|
||
if err != nil {
|
||
t.Fatalf("Sum: %v", err)
|
||
}
|
||
if err := loss.Backward(); err != nil {
|
||
t.Fatalf("Backward: %v", err)
|
||
}
|
||
for i := range 2 {
|
||
g := x.Grad().FloatAt(i)
|
||
if math.IsNaN(g) || g != 0 {
|
||
t.Errorf("d/dx x⁰ at %g = %v, want exactly 0", x.Data().FloatAt(i), g)
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestAutogradLeafBackwardAccumulates pins that Backward on a leaf
|
||
// accumulates into the existing gradient like any other backward
|
||
// pass, instead of overwriting it.
|
||
func TestAutogradLeafBackwardAccumulates(t *testing.T) {
|
||
x := mustTensor(t, []float64{3}, 1)
|
||
if err := x.Backward(); err != nil {
|
||
t.Fatalf("Backward: %v", err)
|
||
}
|
||
if got := x.Grad().FloatAt(0); got != 1 {
|
||
t.Fatalf("first leaf Backward: grad %v, want 1", got)
|
||
}
|
||
if err := x.Backward(); err != nil {
|
||
t.Fatalf("second Backward: %v", err)
|
||
}
|
||
if got := x.Grad().FloatAt(0); got != 2 {
|
||
t.Fatalf("second leaf Backward: grad %v, want 2 (accumulated)", got)
|
||
}
|
||
}
|