Files
tensor/grad/tensor_test.go
T

344 lines
8.4 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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)
}
}