919 lines
24 KiB
Go
919 lines
24 KiB
Go
package grad
|
|
|
|
import (
|
|
"math"
|
|
"testing"
|
|
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|
)
|
|
|
|
// Backward-coverage pins for the graph kernels: every test here checks
|
|
// a committed leaf gradient against a closed-form identity or a central
|
|
// difference of the summed loss, on combinations the op-level tests do
|
|
// not build.
|
|
|
|
// totalLoss sums a non-scalar loss so a central difference matches the
|
|
// all-ones seed Backward uses.
|
|
func totalLoss(l *Tensor) float64 {
|
|
s := 0.0
|
|
d := l.Data()
|
|
for i := range d.Len() {
|
|
s += d.FloatAt(i)
|
|
}
|
|
return s
|
|
}
|
|
|
|
func abs2c(z complex128) float64 { return real(z)*real(z) + imag(z)*imag(z) }
|
|
|
|
func mustRecover(vals []float64, sh ...int) *core.Array {
|
|
a, _ := core.FromFloats(vals, sh...)
|
|
return a
|
|
}
|
|
|
|
func mustRecoverComplex(zs []complex128, shape ...int) *core.Array {
|
|
if len(shape) == 0 {
|
|
shape = []int{len(zs)}
|
|
}
|
|
w := append([]complex128(nil), zs...)
|
|
a, _ := core.ComplexFromArray(w, shape...)
|
|
return a
|
|
}
|
|
|
|
// probeCheckCentral builds the loss from base, backpropagates once and
|
|
// compares every element of the committed gradient against central
|
|
// differences of the summed loss.
|
|
func probeCheckCentral(t *testing.T, base *Tensor, name string, build func() (*Tensor, error), tol float64) {
|
|
t.Helper()
|
|
base.ZeroGrad()
|
|
loss, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
grads := flatFloats(base.Grad())
|
|
orig := flatFloats(base.Data())
|
|
sh := base.Data().Shape()
|
|
for i := range orig {
|
|
const h = 1e-6
|
|
plus := append([]float64(nil), orig...)
|
|
plus[i] += h
|
|
base.ReplaceWith(mustRecover(plus, sh...))
|
|
lp, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
minus := append([]float64(nil), orig...)
|
|
minus[i] -= h
|
|
base.ReplaceWith(mustRecover(minus, sh...))
|
|
lm, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
|
|
if math.Abs(grads[i]-num) > tol*math.Max(1, math.Abs(num)) {
|
|
t.Errorf("%s grad[%d] = %g, central %g", name, i, grads[i], num)
|
|
}
|
|
base.ReplaceWith(mustRecover(orig, sh...))
|
|
}
|
|
}
|
|
|
|
// TestFFTParsevalGradient pins the FFT backward through Parseval's
|
|
// theorem: L = sum |FFT(x)|^2 has gradient n·2x for real x, because the
|
|
// unnormalised forward DFT scales the energy by n.
|
|
func TestFFTParsevalGradient(t *testing.T) {
|
|
const n = 16
|
|
vals := make([]float64, n)
|
|
for i := range vals {
|
|
vals[i] = math.Sin(0.7*float64(i)) + 0.3*float64(i%4)
|
|
}
|
|
x, _ := FromFloat64s(vals, true, n)
|
|
f, _ := x.FFT()
|
|
abs2, _ := f.Abs2()
|
|
loss, _ := abs2.Sum()
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
g := x.Grad()
|
|
for i := range n {
|
|
want := 2 * float64(n) * vals[i]
|
|
if got := g.FloatAt(i); math.Abs(got-want) > 1e-8*math.Max(1, math.Abs(want)) {
|
|
t.Fatalf("Parseval: g[%d] = %g, want %g", i, got, want)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestFFT2ParsevalGradient is Parseval's identity for the rank-2
|
|
// transform, where the backward's scale is the full element count.
|
|
func TestFFT2ParsevalGradient(t *testing.T) {
|
|
vals := []float64{1, 2, -1, 0.5, 3, -2, 0.25, 1.5, -0.5, 2, 1, -3}
|
|
x, _ := FromFloat64s(vals, true, 3, 4)
|
|
f, err := x.FFT2()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
a, err := f.Abs2()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
loss, err := a.Sum()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for i := range vals {
|
|
want := 2 * float64(12) * vals[i]
|
|
if got := x.Grad().FloatAt(i); math.Abs(got-want) > 1e-8*math.Max(1, math.Abs(want)) {
|
|
t.Fatalf("FFT2 Parseval g[%d] = %g, want %g", i, got, want)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestSpectralHalfSpectrumBackward pins the RFFT and IRFFT backwards
|
|
// against central differences, the half-spectrum combinatorics
|
|
// included: the doubled real part on the mirrored bins and the halved
|
|
// self-mirrored bins of the inverse.
|
|
func TestSpectralHalfSpectrumBackward(t *testing.T) {
|
|
x, _ := FromFloat64s([]float64{1, -2, 0.5, 3, -1.5, 0.25, 2, -0.75}, true, 8)
|
|
buildF := func() (*Tensor, error) {
|
|
h, err := x.RFFT()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
a, err := h.Abs2()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return a.Sum()
|
|
}
|
|
probeCheckCentral(t, x, "rfft", buildF, 1e-4)
|
|
|
|
// IRFFT: the leaf is the half spectrum; the loss is the summed
|
|
// square of the real signal. The complex leaf's gradient is checked
|
|
// against a Wirtinger central difference of the real loss:
|
|
// dL/dzbar_j = (dL/dRe_j + i dL/dIm_j)/2, each part probed by
|
|
// perturbing that component.
|
|
zr := []float64{3, -1, 0.5, 2, 1.5}
|
|
zi := []float64{0, 1, -0.5, 0.25, -1}
|
|
zs := make([]complex128, 5)
|
|
for i := range zs {
|
|
zs[i] = complex(zr[i], zi[i])
|
|
}
|
|
za, _ := core.ComplexFromArray(zs, 5)
|
|
z := FromArray(za, true)
|
|
buildI := func() (*Tensor, error) {
|
|
sig, err := z.IRFFT(8)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sq, err := sig.Pow(2)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return sq.Sum()
|
|
}
|
|
z.ZeroGrad()
|
|
loss, err := buildI()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
const h = 1e-6
|
|
for j := range 5 {
|
|
perturb := func(dr, di float64) float64 {
|
|
ws := append([]complex128(nil), zs...)
|
|
ws[j] = complex(zr[j]+dr, zi[j]+di)
|
|
z.ReplaceWith(mustRecoverComplex(ws, 5))
|
|
l, err := buildI()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return totalLoss(l)
|
|
}
|
|
dr := (perturb(h, 0) - perturb(-h, 0)) / (2 * h)
|
|
di := (perturb(0, h) - perturb(0, -h)) / (2 * h)
|
|
want := complex(0.5*dr, 0.5*di)
|
|
if got := z.Grad().ComplexAt(j); abs2c(got-want) > 1e-4*abs2c(want) {
|
|
t.Errorf("irfft grad[%d] = %g, want %g", j, got, want)
|
|
}
|
|
z.ReplaceWith(mustRecoverComplex(zs, 5))
|
|
}
|
|
}
|
|
|
|
// TestCompositeGraphBackward walks a graph spanning slice, matmul,
|
|
// tanh, broadcast, abs2 and an axis reduction, and checks both leaves'
|
|
// gradients against central differences of the summed loss.
|
|
func TestCompositeGraphBackward(t *testing.T) {
|
|
a, _ := FromFloat64s([]float64{1, -2, 3, -4, 5, -6}, true, 2, 3)
|
|
w, _ := FromFloat64s([]float64{0.5, -0.25, 0.125, 1, -2, 0.75}, true, 2, 3)
|
|
|
|
build := func() (*Tensor, error) {
|
|
s, err := a.Slice(1, 1, 3) // (2,2)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
mm, err := s.MatMul(w)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
th, err := mm.Tanh()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
bc, err := th.BroadcastTo(2, 2, 3)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sq, err := bc.Abs2()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return sq.SumAxis(2)
|
|
}
|
|
numCheck := func(base *Tensor, name string, grads []float64) {
|
|
t.Helper()
|
|
orig := flatFloats(base.Data())
|
|
sh := base.Data().Shape()
|
|
for i := range orig {
|
|
const h = 1e-6
|
|
plus := append([]float64(nil), orig...)
|
|
plus[i] += h
|
|
base.ReplaceWith(mustRecover(plus, sh...))
|
|
lp, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
minus := append([]float64(nil), orig...)
|
|
minus[i] -= h
|
|
base.ReplaceWith(mustRecover(minus, sh...))
|
|
lm, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
|
|
if math.Abs(grads[i]-num) > 1e-5*math.Max(1, math.Abs(num)) {
|
|
t.Fatalf("%s[%d]: backward %g, central difference %g", name, i, grads[i], num)
|
|
}
|
|
base.ReplaceWith(mustRecover(orig, sh...))
|
|
}
|
|
}
|
|
// Backward accumulates, so every read starts from a clean slate.
|
|
a.ZeroGrad()
|
|
w.ZeroGrad()
|
|
loss, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
numCheck(a, "a", flatFloats(a.Grad()))
|
|
a.ZeroGrad()
|
|
w.ZeroGrad()
|
|
loss, err = build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
numCheck(w, "w", flatFloats(w.Grad()))
|
|
}
|
|
|
|
// TestMatMulBatchedBackward checks the stacked product rule of
|
|
// MatMulBatched against central differences on both operands.
|
|
func TestMatMulBatchedBackward(t *testing.T) {
|
|
a, _ := FromFloat64s([]float64{1, 2, 3, 4, 5, 6, 7, 8}, true, 2, 2, 2)
|
|
b, _ := FromFloat64s([]float64{0.5, -1, 2, 0.25, 1.5, -0.5, -2, 1}, true, 2, 2, 2)
|
|
|
|
build := func() (*Tensor, error) {
|
|
c, err := a.MatMulBatched(b)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sq, err := c.Abs2()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return sq.Sum()
|
|
}
|
|
a.ZeroGrad()
|
|
b.ZeroGrad()
|
|
loss, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for _, tc := range []struct {
|
|
tensor *Tensor
|
|
name string
|
|
}{{a, "a"}, {b, "b"}} {
|
|
grads := flatFloats(tc.tensor.Grad())
|
|
orig := flatFloats(tc.tensor.Data())
|
|
sh := tc.tensor.Data().Shape()
|
|
for i := range orig {
|
|
const h = 1e-6
|
|
plus := append([]float64(nil), orig...)
|
|
plus[i] += h
|
|
tc.tensor.ReplaceWith(mustRecover(plus, sh...))
|
|
lp, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
minus := append([]float64(nil), orig...)
|
|
minus[i] -= h
|
|
tc.tensor.ReplaceWith(mustRecover(minus, sh...))
|
|
lm, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
|
|
if math.Abs(grads[i]-num) > 1e-4*math.Max(1, math.Abs(num)) {
|
|
t.Fatalf("%s[%d]: backward %g, central difference %g", tc.name, i, grads[i], num)
|
|
}
|
|
tc.tensor.ReplaceWith(mustRecover(orig, sh...))
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestStridedOperandGradientMatchesDenseTwin pins the graph's stride
|
|
// discipline: a MatMul over a sliced view and the same product over a
|
|
// dense twin must commit identical gradients. The view's leaf receives
|
|
// the dense twin's gradient scattered into the sliced span, zeros
|
|
// elsewhere.
|
|
func TestStridedOperandGradientMatchesDenseTwin(t *testing.T) {
|
|
wVals := []float64{0.5, -0.25, 0.125, 1, -2, 0.75}
|
|
|
|
// The strided run: leaf (2,4) -> slice -> matmul.
|
|
full, _ := FromFloat64s([]float64{9, 1, -2, 3, 9, 5, -6, 9}, true, 2, 4)
|
|
s, err := full.Slice(1, 1, 4) // (2,3) strided view holding 1,-2,3 / 5,-6,9
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
w, _ := FromFloat64s(wVals, true, 3, 2)
|
|
mm, err := s.MatMul(w)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
sq, err := mm.Abs2()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
loss, err := sq.Sum()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if full.Grad() == nil || w.Grad() == nil {
|
|
t.Fatal("both leaves must receive a gradient")
|
|
}
|
|
gs := flatFloats(full.Grad())
|
|
wStrided := flatFloats(w.Grad())
|
|
full.ZeroGrad()
|
|
w.ZeroGrad()
|
|
|
|
// The dense twin.
|
|
a2, _ := FromFloat64s(flatFloats(s.Data()), true, 2, 3)
|
|
w2, _ := FromFloat64s(wVals, true, 3, 2)
|
|
mm2, err := a2.MatMul(w2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
sq2, err := mm2.Abs2()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
loss2, err := sq2.Sum()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss2.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
aDense := flatFloats(a2.Grad())
|
|
wDense := flatFloats(w2.Grad())
|
|
|
|
for i := range wStrided {
|
|
if wStrided[i] != wDense[i] {
|
|
t.Fatalf("w gradient %v over the strided operand, want the dense twin's %v", wStrided, wDense)
|
|
}
|
|
}
|
|
if gs[0] != 0 || gs[4] != 0 {
|
|
t.Fatalf("columns outside the slice carry %g and %g, want 0", gs[0], gs[4])
|
|
}
|
|
if gs[1] != aDense[0] || gs[2] != aDense[1] || gs[3] != aDense[2] ||
|
|
gs[5] != aDense[3] || gs[6] != aDense[4] || gs[7] != aDense[5] {
|
|
t.Fatalf("leaf gradient %v does not scatter the dense gradient %v into the slice", gs, aDense)
|
|
}
|
|
}
|
|
|
|
// TestHessianMatchesAnalyticForm differentiates
|
|
// f(x, y) = x^3 y + exp(x) log(y+2) twice and compares the answer with
|
|
// the closed form:
|
|
//
|
|
// dxx = 6xy + exp(x)log(y+2), dxy = 3x^2 + exp(x)/(y+2),
|
|
// dyy = -exp(x)/(y+2)^2.
|
|
func TestHessianMatchesAnalyticForm(t *testing.T) {
|
|
f := func(v *Tensor) (*Tensor, error) {
|
|
x, err := v.Slice(0, 0, 1)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
y, err := v.Slice(0, 1, 2)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
x3, err := x.Pow(3)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
t1, err := x3.Mul(y)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
e, err := x.Exp()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
y2, err := y.Add(FromArray(mustRecover([]float64{2}, 1), false))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
lg, err := y2.Log()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
t2, err := e.Mul(lg)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return t1.Add(t2)
|
|
}
|
|
x, _ := FromFloat64s([]float64{0.4, 1.7}, true, 2)
|
|
h, err := Hessian(f, x, HessianOptions{Step: 1e-5})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
xx, yy := 0.4, 1.7
|
|
ex, ly := math.Exp(xx), math.Log(yy+2)
|
|
want := [4]float64{
|
|
6*xx*yy + ex*ly, 3*xx*xx + ex/(yy+2),
|
|
3*xx*xx + ex/(yy+2), -ex / ((yy + 2) * (yy + 2)),
|
|
}
|
|
for i := range 4 {
|
|
if math.Abs(h.FloatAt(i)-want[i]) > 1e-4*math.Max(1, math.Abs(want[i])) {
|
|
t.Fatalf("Hessian[%d] = %g, want %g", i, h.FloatAt(i), want[i])
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestTripleUseLeafGradient folds a leaf through three nodes and pins
|
|
// the accumulated gradient: dL/dz = 2(v^2+v)(2v+1) for
|
|
// L = (v^2+v)^2.
|
|
func TestTripleUseLeafGradient(t *testing.T) {
|
|
for _, v := range []float64{0.5, -1.25, 2} {
|
|
z, _ := FromFloat64s([]float64{v}, true, 1)
|
|
z2, _ := z.Mul(z)
|
|
s, err := z2.Add(z)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
sq, err := s.Pow(2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
loss, err := sq.Sum()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
want := 2 * (v*v + v) * (2*v + 1)
|
|
if got := z.Grad().FloatAt(0); math.Abs(got-want) > 1e-12*math.Max(1, math.Abs(want)) {
|
|
t.Fatalf("triple use at z=%g: g = %g, want %g", v, got, want)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestComplexAbs2LeafGradient pins the Wirtinger gradient of a summed
|
|
// |z|^2 loss: the leaf receives z itself.
|
|
func TestComplexAbs2LeafGradient(t *testing.T) {
|
|
zr := []float64{0.3, -1.2, 0.7}
|
|
zi := []float64{-0.4, 0.8, 1.1}
|
|
zs := make([]complex128, 3)
|
|
for i := range zs {
|
|
zs[i] = complex(zr[i], zi[i])
|
|
}
|
|
za, _ := core.ComplexFromArray(zs, 3)
|
|
z := FromArray(za, true)
|
|
abs2, err := z.Abs2()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
loss, err := abs2.Sum()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
g := z.Grad()
|
|
for i := range 3 {
|
|
want := complex(zr[i], zi[i])
|
|
if got := g.ComplexAt(i); got != want {
|
|
t.Fatalf("|z|^2 leaf grad[%d] = %g, want %g", i, got, want)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestComplexMatMulBackward checks the Wirtinger adjoint of MatMul
|
|
// against component-wise central differences of a real loss. The leaf
|
|
// gradient is dL/dzbar = (dL/dRe + i dL/dIm)/2 under the package
|
|
// convention.
|
|
func TestComplexMatMulBackward(t *testing.T) {
|
|
az := []complex128{complex(1, 0.5), complex(-0.5, 2), complex(0.25, -1), complex(2, 0.25)}
|
|
bz := []complex128{complex(0.5, 1), complex(-1, 0.25), complex(1.5, -0.5), complex(0.75, 1.25)}
|
|
aa, _ := core.ComplexFromArray(az, 2, 2)
|
|
ba, _ := core.ComplexFromArray(bz, 2, 2)
|
|
a := FromArray(aa, true)
|
|
b := FromArray(ba, true)
|
|
build := func() (*Tensor, error) {
|
|
m, err := a.MatMul(b)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
r, err := m.Real()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sq, err := r.Abs2()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return sq.Sum()
|
|
}
|
|
a.ZeroGrad()
|
|
b.ZeroGrad()
|
|
loss, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
const h = 1e-6
|
|
check := func(base *Tensor, vals []complex128, name string) {
|
|
t.Helper()
|
|
for j := range len(vals) {
|
|
perturb := func(dr, di float64) float64 {
|
|
ws := append([]complex128(nil), vals...)
|
|
ws[j] = complex(real(vals[j])+dr, imag(vals[j])+di)
|
|
base.ReplaceWith(mustRecoverComplex(ws, 2, 2))
|
|
l, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return totalLoss(l)
|
|
}
|
|
dr := (perturb(h, 0) - perturb(-h, 0)) / (2 * h)
|
|
di := (perturb(0, h) - perturb(0, -h)) / (2 * h)
|
|
want := complex(0.5*dr, 0.5*di)
|
|
if got := base.Grad().ComplexAt(j); abs2c(got-want) > 1e-4*abs2c(want) {
|
|
t.Errorf("complex matmul %s grad[%d] = %g, want %g", name, j, got, want)
|
|
}
|
|
base.ReplaceWith(mustRecoverComplex(vals, 2, 2))
|
|
}
|
|
}
|
|
check(a, az, "a")
|
|
check(b, bz, "b")
|
|
}
|
|
|
|
// TestConcatMixedDtypeBackward checks the join's backward on a real and
|
|
// a complex side at once: the real side receives 2 Re g through the
|
|
// narrowing, the complex side its Wirtinger gradient.
|
|
func TestConcatMixedDtypeBackward(t *testing.T) {
|
|
r, _ := FromFloat64s([]float64{1, -2, 3}, true, 3)
|
|
zr := []complex128{complex(0.5, 1), complex(-1.5, 0.25)}
|
|
za, _ := core.ComplexFromArray(zr, 2)
|
|
z := FromArray(za, true)
|
|
build := func() (*Tensor, error) {
|
|
c, err := r.Concat(z, 0)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
re, err := c.Real()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sq, err := re.Abs2()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return sq.Sum()
|
|
}
|
|
r.ZeroGrad()
|
|
z.ZeroGrad()
|
|
loss, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
grads := flatFloats(r.Grad())
|
|
orig := flatFloats(r.Data())
|
|
for i := range orig {
|
|
const h = 1e-6
|
|
plus := append([]float64(nil), orig...)
|
|
plus[i] += h
|
|
r.ReplaceWith(mustRecover(plus, 3))
|
|
lp, _ := build()
|
|
minus := append([]float64(nil), orig...)
|
|
minus[i] -= h
|
|
r.ReplaceWith(mustRecover(minus, 3))
|
|
lm, _ := build()
|
|
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
|
|
if math.Abs(grads[i]-num) > 1e-5*math.Max(1, math.Abs(num)) {
|
|
t.Errorf("concat real side grad[%d] = %g, central %g", i, grads[i], num)
|
|
}
|
|
r.ReplaceWith(mustRecover(orig, 3))
|
|
}
|
|
for j := range 2 {
|
|
perturb := func(dr, di float64) float64 {
|
|
ws := append([]complex128(nil), zr...)
|
|
ws[j] = complex(real(zr[j])+dr, imag(zr[j])+di)
|
|
z.ReplaceWith(mustRecoverComplex(ws, 2))
|
|
l, _ := build()
|
|
return totalLoss(l)
|
|
}
|
|
dr := (perturb(1e-6, 0) - perturb(-1e-6, 0)) / (2e-6)
|
|
di := (perturb(0, 1e-6) - perturb(0, -1e-6)) / (2e-6)
|
|
want := complex(0.5*dr, 0.5*di)
|
|
if got := z.Grad().ComplexAt(j); abs2c(got-want) > 1e-4*abs2c(want) {
|
|
t.Errorf("concat complex side grad[%d] = %g, want %g", j, got, want)
|
|
}
|
|
z.ReplaceWith(mustRecoverComplex(zr, 2))
|
|
}
|
|
}
|
|
|
|
// TestTransposeAxesBackward reverses an axis permutation under a
|
|
// non-linear loss and checks the gradient against central differences.
|
|
func TestTransposeAxesBackward(t *testing.T) {
|
|
x, _ := FromFloat64s([]float64{1, 2, 3, 4, 5, 6, 7, 8}, true, 2, 2, 2)
|
|
build := func() (*Tensor, error) {
|
|
p, err := x.TransposeAxes(2, 0, 1)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sq, err := p.Abs2()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
w, err := sq.Mul(p)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return w.Sum()
|
|
}
|
|
probeCheckCentral(t, x, "transposeaxes", build, 1e-5)
|
|
}
|
|
|
|
// TestUnaryKernelChainBackward pushes Sqrt, Log, Sigmoid, Scale, Neg,
|
|
// Mean and Pow through one graph and checks the committed gradient.
|
|
func TestUnaryKernelChainBackward(t *testing.T) {
|
|
x, _ := FromFloat64s([]float64{0.5, 1.5, 2.5, 3.5}, true, 4)
|
|
build := func() (*Tensor, error) {
|
|
s, err := x.Sqrt()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
l, err := s.Log()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sg, err := x.Sigmoid()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
m, err := l.Mul(sg)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sc, err := m.Scale(2.5)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
ng, err := sc.Neg()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
mn, err := ng.Mean()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return mn.Pow(2)
|
|
}
|
|
probeCheckCentral(t, x, "unary-chain", build, 1e-4)
|
|
}
|
|
|
|
// TestBackwardAccumulatesUntilZeroGrad pins the accumulation contract
|
|
// end to end: two Backward calls double the gradient, ZeroGrad resets.
|
|
func TestBackwardAccumulatesUntilZeroGrad(t *testing.T) {
|
|
x, _ := FromFloat64s([]float64{2}, true, 1)
|
|
build := func() (*Tensor, error) {
|
|
s, err := x.Pow(2)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return s.Sum()
|
|
}
|
|
l1, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := l1.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
l2, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := l2.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if got := x.Grad().FloatAt(0); got != 8 {
|
|
t.Fatalf("accumulated g = %g, want 8", got)
|
|
}
|
|
x.ZeroGrad()
|
|
l3, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := l3.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if got := x.Grad().FloatAt(0); got != 4 {
|
|
t.Fatalf("post-zero g = %g, want 4", got)
|
|
}
|
|
}
|
|
|
|
// TestFloat32LeafKeepsGradientDtype pins the dtype contract on a
|
|
// float32 leaf: the gradient narrows to the leaf's width.
|
|
func TestFloat32LeafKeepsGradientDtype(t *testing.T) {
|
|
v := []float32{1.5, -2.5}
|
|
a, _ := core.FromFloat32Slice(v, 2)
|
|
x := FromArray(a, true)
|
|
y, err := x.Pow(2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
l, err := y.Sum()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := l.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
g := x.Grad()
|
|
if g.Dtype() != core.Float32 {
|
|
t.Fatalf("grad dtype = %s, want float32", g.Dtype())
|
|
}
|
|
for i := range 2 {
|
|
want := 2 * float64(v[i])
|
|
if got := g.FloatAt(i); math.Abs(got-want) > 1e-6 {
|
|
t.Fatalf("g[%d] = %g, want %g", i, got, want)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestBroadcastToBackwardMatchesCentralDifferences and its SumAxis
|
|
// neighbour cover the reduction/expansion pair in isolation.
|
|
func TestBroadcastToBackwardMatchesCentralDifferences(t *testing.T) {
|
|
x, _ := FromFloat64s([]float64{1, -2, 3, -4, 5, -6}, true, 2, 3)
|
|
build := func() (*Tensor, error) {
|
|
bc, err := x.BroadcastTo(2, 2, 3)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sq, err := bc.Abs2()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return sq.Sum()
|
|
}
|
|
probeCheckCentral(t, x, "broadcast", build, 1e-5)
|
|
}
|
|
|
|
// TestSumAxisBackwardMatchesCentralDifferences reduces the middle axis
|
|
// with a non-scalar loss, so the backward must scatter into both rows.
|
|
func TestSumAxisBackwardMatchesCentralDifferences(t *testing.T) {
|
|
x, _ := FromFloat64s([]float64{1, -2, 3, -4, 5, -6}, true, 2, 3)
|
|
build := func() (*Tensor, error) {
|
|
sq, err := x.Abs2()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return sq.SumAxis(1)
|
|
}
|
|
probeCheckCentral(t, x, "sumaxis", build, 1e-5)
|
|
}
|
|
|
|
// TestMatMulTanhBackwardMatchesCentralDifferences checks the matmul
|
|
// product rule under a tanh on top, non-scalar loss included.
|
|
func TestMatMulTanhBackwardMatchesCentralDifferences(t *testing.T) {
|
|
a, _ := FromFloat64s([]float64{1, -2, 3, -4}, true, 2, 2)
|
|
w, _ := FromFloat64s([]float64{0.5, -0.25, 1, -2}, true, 2, 2)
|
|
build := func() (*Tensor, error) {
|
|
mm, err := a.MatMul(w)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return mm.Tanh()
|
|
}
|
|
a.ZeroGrad()
|
|
w.ZeroGrad()
|
|
loss, err := build()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := loss.Backward(); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
check := func(base *Tensor, name string) {
|
|
t.Helper()
|
|
grads := flatFloats(base.Grad())
|
|
orig := flatFloats(base.Data())
|
|
sh := base.Data().Shape()
|
|
for i := range orig {
|
|
const h = 1e-6
|
|
plus := append([]float64(nil), orig...)
|
|
plus[i] += h
|
|
base.ReplaceWith(mustRecover(plus, sh...))
|
|
lp, _ := build()
|
|
minus := append([]float64(nil), orig...)
|
|
minus[i] -= h
|
|
base.ReplaceWith(mustRecover(minus, sh...))
|
|
lm, _ := build()
|
|
var num float64
|
|
for j := range lp.Data().Len() {
|
|
num += lp.Data().FloatAt(j) - lm.Data().FloatAt(j)
|
|
}
|
|
num /= 2 * h
|
|
if math.Abs(grads[i]-num) > 1e-5*math.Max(1, math.Abs(num)) {
|
|
t.Errorf("%s[%d] backward %g, central %g", name, i, grads[i], num)
|
|
}
|
|
base.ReplaceWith(mustRecover(orig, sh...))
|
|
}
|
|
}
|
|
check(a, "a")
|
|
check(w, "w")
|
|
}
|
|
|
|
// TestNewtonCGSolvesLeastSquares pins the truncated Newton method on a
|
|
// small full-rank least-squares problem whose solution is A^-1 b.
|
|
func TestNewtonCGSolvesLeastSquares(t *testing.T) {
|
|
ar := []float64{2, 0.5, 1, 3}
|
|
br := []float64{1, -1}
|
|
A, _ := FromFloat64s(ar, false, 2, 2)
|
|
B, _ := FromFloat64s(br, false, 2)
|
|
f := func(x *Tensor) (*Tensor, error) {
|
|
ax, err := A.MatMul(x)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
d, err := ax.Sub(B)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
s, err := d.Abs2()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return s.Sum()
|
|
}
|
|
x0 := mustRecover([]float64{0, 0}, 2)
|
|
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
det := 2*3 - 0.5*1
|
|
wantX := (3*1 - 0.5*(-1)) / det
|
|
wantY := (2*(-1) - 1*1) / det
|
|
if math.Abs(x.FloatAt(0)-wantX) > 1e-5 || math.Abs(x.FloatAt(1)-wantY) > 1e-5 {
|
|
t.Fatalf("NewtonCG point (%g, %g), want (%g, %g)", x.FloatAt(0), x.FloatAt(1), wantX, wantY)
|
|
}
|
|
if fv > 1e-10 {
|
|
t.Fatalf("NewtonCG f = %g, want ~0", fv)
|
|
}
|
|
}
|