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) } }