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