876 lines
28 KiB
Go
876 lines
28 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
|
|
// SPDX-License-Identifier: MIT
|
|||
|
|
|
|||
|
|
package grad
|
|||
|
|
|
|||
|
|
import (
|
|||
|
|
"math"
|
|||
|
|
"strings"
|
|||
|
|
"testing"
|
|||
|
|
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// Regression pins for the grad package guards: refusals and gradients
|
|||
|
|
// that used to panic or silently pass through. Each test names
|
|||
|
|
// the defect it pins and fails without its fix.
|
|||
|
|
|
|||
|
|
// TestConcatMixedDtypeBackwardNarrows pins the mixed real/complex
|
|||
|
|
// Concat backward. The join promotes along the dtype ladder, so a
|
|||
|
|
// complex gradient reaches a real operand; that operand must receive
|
|||
|
|
// 2·Re of its span, the package's real-operand rule, instead of dying
|
|||
|
|
// in copyElem, which used to read the complex source through FloatAt (a
|
|||
|
|
// nil int payload) and panic. Both operand orders and the constant
|
|||
|
|
// complex operand are covered, and the real side is checked against
|
|||
|
|
// central differences as well as its closed form.
|
|||
|
|
func TestConcatMixedDtypeBackwardNarrows(t *testing.T) {
|
|||
|
|
xv := []float64{1.5, -0.75, 2.25, 0.5, -1, 3}
|
|||
|
|
zv := []complex128{1 + 1i, 2 - 0.5i, -0.25 + 0.75i, 3, 0.5 - 2i, -1.5i}
|
|||
|
|
|
|||
|
|
// L = Σ|Concat(a, b)|² = Σx² + Σ|z|², so the real operand's
|
|||
|
|
// gradient is 2x and the complex one's is z, whatever the axis.
|
|||
|
|
cases := []struct {
|
|||
|
|
name string
|
|||
|
|
dim int
|
|||
|
|
xSh []int
|
|||
|
|
zSh []int
|
|||
|
|
}{
|
|||
|
|
{name: "dim0", dim: 0, xSh: []int{2, 3}, zSh: []int{2, 3}},
|
|||
|
|
{name: "dim1", dim: 1, xSh: []int{3, 2}, zSh: []int{3, 2}},
|
|||
|
|
}
|
|||
|
|
for _, tc := range cases {
|
|||
|
|
t.Run(tc.name, func(t *testing.T) {
|
|||
|
|
x, err := core.FromFloats(xv, tc.xSh...)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
z, err := core.FromComplexes(zv, tc.zSh...)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromComplexes: %v", err)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// The A side of a real x complex join is the panic the
|
|||
|
|
// report reproduced; the B side of a complex x real join
|
|||
|
|
// fails identically.
|
|||
|
|
concatChecked(t, "real||cplx", FromArray(x, true), FromArray(z, true), tc.dim, xv, zv)
|
|||
|
|
concatChecked(t, "cplx||real", FromArray(z, true), FromArray(x, true), tc.dim, xv, zv)
|
|||
|
|
|
|||
|
|
// A constant complex operand takes the same span loop, so
|
|||
|
|
// the real side must still narrow.
|
|||
|
|
concatChecked(t, "real||cplxconst", FromArray(x, true), FromArray(z, false), tc.dim, xv, nil)
|
|||
|
|
|
|||
|
|
// Finite differences confirm the 2·Re rule itself, not just
|
|||
|
|
// its agreement with the closed form.
|
|||
|
|
ref := numericGrad(func(v *core.Array) float64 {
|
|||
|
|
realSide := FromArray(v, false)
|
|||
|
|
cat, cerr := realSide.Concat(FromArray(z, false), tc.dim)
|
|||
|
|
if cerr != nil {
|
|||
|
|
return math.NaN()
|
|||
|
|
}
|
|||
|
|
sq, cerr := cat.Abs2()
|
|||
|
|
if cerr != nil {
|
|||
|
|
return math.NaN()
|
|||
|
|
}
|
|||
|
|
s, cerr := sq.Sum()
|
|||
|
|
if cerr != nil {
|
|||
|
|
return math.NaN()
|
|||
|
|
}
|
|||
|
|
return s.Data().FloatAt(0)
|
|||
|
|
}, x)
|
|||
|
|
xt := FromArray(x, true)
|
|||
|
|
concatChecked(t, "fd", xt, FromArray(z, false), tc.dim, xv, nil)
|
|||
|
|
if d := maxAbsDiff(flatFloats(xt.Grad()), ref); d > 1e-8 {
|
|||
|
|
t.Errorf("real operand gradient differs from central differences by %g", d)
|
|||
|
|
}
|
|||
|
|
})
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// A float32 real operand beside a complex one: the narrowed side
|
|||
|
|
// keeps the operand's width, and the value is still 2x.
|
|||
|
|
t.Run("float32RealSide", func(t *testing.T) {
|
|||
|
|
x32 := []float32{1.5, -0.75, 2.25, 0.5}
|
|||
|
|
z32 := []complex128{1 + 1i, 2 - 0.5i, -0.25 + 0.75i, 3}
|
|||
|
|
a, err := core.FromFloat32s(x32, 2, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat32s: %v", err)
|
|||
|
|
}
|
|||
|
|
z, err := core.FromComplexes(z32, 2, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromComplexes: %v", err)
|
|||
|
|
}
|
|||
|
|
xt := FromArray(a, true)
|
|||
|
|
cat, err := xt.Concat(FromArray(z, true), 0)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Concat: %v", err)
|
|||
|
|
}
|
|||
|
|
sq, err := cat.Abs2()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Abs2: %v", err)
|
|||
|
|
}
|
|||
|
|
loss, err := sq.Sum()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Sum: %v", err)
|
|||
|
|
}
|
|||
|
|
if err := loss.Backward(); err != nil {
|
|||
|
|
t.Fatalf("Backward: %v", err)
|
|||
|
|
}
|
|||
|
|
g := xt.Grad()
|
|||
|
|
if got := g.Dtype(); got != core.Float32 {
|
|||
|
|
t.Errorf("float32 real operand gradient dtype = %s, want float32", got)
|
|||
|
|
}
|
|||
|
|
for i, v := range x32 {
|
|||
|
|
if got, want := g.FloatAt(i), 2*float64(v); math.Abs(got-want) > 1e-6 {
|
|||
|
|
t.Errorf("float32 gradient[%d] = %v, want %v (2·Re rule)", i, got, want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
// A rebased-view complex operand: the view aliases a longer payload,
|
|||
|
|
// and the gradient must scatter back through the view's own slots.
|
|||
|
|
t.Run("rebasedViewOperand", func(t *testing.T) {
|
|||
|
|
big, err := core.FromComplexes(zv[:4], 4)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromComplexes: %v", err)
|
|||
|
|
}
|
|||
|
|
zb := FromArray(big, true)
|
|||
|
|
view, err := zb.Slice(0, 1, 3)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Slice: %v", err)
|
|||
|
|
}
|
|||
|
|
x, err := core.FromFloats(xv[:2], 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
xt := FromArray(x, true)
|
|||
|
|
cat, err := xt.Concat(view, 0)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Concat: %v", err)
|
|||
|
|
}
|
|||
|
|
sq, err := cat.Abs2()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Abs2: %v", err)
|
|||
|
|
}
|
|||
|
|
loss, err := sq.Sum()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Sum: %v", err)
|
|||
|
|
}
|
|||
|
|
if err := loss.Backward(); err != nil {
|
|||
|
|
t.Fatalf("Backward: %v", err)
|
|||
|
|
}
|
|||
|
|
for i := range xv[:2] {
|
|||
|
|
if got, want := xt.Grad().FloatAt(i), 2*xv[i]; got != want {
|
|||
|
|
t.Errorf("view case: real gradient[%d] = %v, want %v", i, got, want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
g := zb.Grad()
|
|||
|
|
if g == nil {
|
|||
|
|
t.Fatal("the view's parent received no gradient")
|
|||
|
|
}
|
|||
|
|
// Only the aliased slots carry the view's own values; the rest
|
|||
|
|
// stays zero through the Slice backward.
|
|||
|
|
want := []complex128{0, zv[1], zv[2], 0}
|
|||
|
|
for i := range want {
|
|||
|
|
if got := g.ComplexAt(i); got != want[i] {
|
|||
|
|
t.Errorf("view case: parent gradient[%d] = %v, want %v", i, got, want[i])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
// The mirror image: a rebased-view REAL operand beside a complex one.
|
|||
|
|
// The narrowed span (2·Re) reaches the view as float and Slice's
|
|||
|
|
// backward scatters it into the parent's own slots, leaving the rest
|
|||
|
|
// zero. This is the one shape that exercises both fixes together.
|
|||
|
|
t.Run("rebasedViewRealOperand", func(t *testing.T) {
|
|||
|
|
full := []float64{9, 1.5, -0.75, 7}
|
|||
|
|
fa, err := core.FromFloats(full, 4)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
fb := FromArray(fa, true)
|
|||
|
|
view, err := fb.Slice(0, 1, 3)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Slice: %v", err)
|
|||
|
|
}
|
|||
|
|
z, err := core.FromComplexes(zv[:2], 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromComplexes: %v", err)
|
|||
|
|
}
|
|||
|
|
cat, err := view.Concat(FromArray(z, true), 0)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Concat: %v", err)
|
|||
|
|
}
|
|||
|
|
sq, err := cat.Abs2()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Abs2: %v", err)
|
|||
|
|
}
|
|||
|
|
loss, err := sq.Sum()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Sum: %v", err)
|
|||
|
|
}
|
|||
|
|
if err := loss.Backward(); err != nil {
|
|||
|
|
t.Fatalf("Backward: %v", err)
|
|||
|
|
}
|
|||
|
|
g := fb.Grad()
|
|||
|
|
if g == nil {
|
|||
|
|
t.Fatal("the view's parent received no gradient")
|
|||
|
|
}
|
|||
|
|
want := []float64{0, 2 * full[1], 2 * full[2], 0}
|
|||
|
|
for i := range want {
|
|||
|
|
if got := g.FloatAt(i); got != want[i] {
|
|||
|
|
t.Errorf("view case: parent gradient[%d] = %v, want %v", i, got, want[i])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
})
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// concatChecked builds Σ|first.Concat(second, dim)|², runs the
|
|||
|
|
// backward and checks every gradient-carrying operand against its
|
|||
|
|
// closed form by dtype: a real side against 2x (the 2·Re rule), a
|
|||
|
|
// complex side against z. A nil wantZ marks a constant complex operand,
|
|||
|
|
// and a non-grad operand has nothing to check.
|
|||
|
|
func concatChecked(t *testing.T, label string, first, second *Tensor, dim int, xv []float64, wantZ []complex128) {
|
|||
|
|
t.Helper()
|
|||
|
|
cat, err := first.Concat(second, dim)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("%s: Concat: %v", label, err)
|
|||
|
|
}
|
|||
|
|
sq, err := cat.Abs2()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("%s: Abs2: %v", label, err)
|
|||
|
|
}
|
|||
|
|
loss, err := sq.Sum()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("%s: Sum: %v", label, err)
|
|||
|
|
}
|
|||
|
|
if err := loss.Backward(); err != nil {
|
|||
|
|
t.Fatalf("%s: Backward: %v", label, err)
|
|||
|
|
}
|
|||
|
|
for _, side := range []*Tensor{first, second} {
|
|||
|
|
if !side.RequiresGrad() {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
g := side.Grad()
|
|||
|
|
if g == nil {
|
|||
|
|
t.Fatalf("%s: an operand received no gradient", label)
|
|||
|
|
}
|
|||
|
|
if g.Dtype() == core.Complex {
|
|||
|
|
if wantZ == nil {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
for i := range wantZ {
|
|||
|
|
if got := g.ComplexAt(i); got != wantZ[i] {
|
|||
|
|
t.Errorf("%s: complex operand gradient[%d] = %v, want %v", label, i, got, wantZ[i])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
for i := range xv {
|
|||
|
|
if got, want := g.FloatAt(i), 2*xv[i]; got != want {
|
|||
|
|
t.Errorf("%s: real operand gradient[%d] = %v, want %v (the 2·Re rule)", label, i, got, want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestHessianAndHVPRejectComplexOperands pins the dtype guard on the
|
|||
|
|
// two second-order helpers: a complex point or direction is refused
|
|||
|
|
// with an error naming the dtype, as MinimiseNewtonCG, SampleHMC and
|
|||
|
|
// AdjointODE already do, never a panic out of flatFloats.
|
|||
|
|
func TestHessianAndHVPRejectComplexOperands(t *testing.T) {
|
|||
|
|
z, err := core.FromComplexes([]complex128{1 + 1i, 2 - 1i}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromComplexes: %v", err)
|
|||
|
|
}
|
|||
|
|
zt := FromArray(z, false)
|
|||
|
|
xf, err := core.FromFloats([]float64{1, 2}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
xt := FromArray(xf, false)
|
|||
|
|
objective := func(q *Tensor) (*Tensor, error) {
|
|||
|
|
sq, err := q.Abs2()
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return sq.Sum()
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
refusesComplex(t, "Hessian on a complex point", func() error {
|
|||
|
|
_, err := Hessian(objective, zt, HessianOptions{})
|
|||
|
|
return err
|
|||
|
|
})
|
|||
|
|
refusesComplex(t, "HessianVectorProduct on a complex point", func() error {
|
|||
|
|
_, err := HessianVectorProduct(objective, zt, xt, HessianOptions{})
|
|||
|
|
return err
|
|||
|
|
})
|
|||
|
|
refusesComplex(t, "HessianVectorProduct on a complex direction", func() error {
|
|||
|
|
_, err := HessianVectorProduct(objective, xt, zt, HessianOptions{})
|
|||
|
|
return err
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
// The guards must not narrow the accepted surface: a real point
|
|||
|
|
// still differentiates, and the real callers keep working.
|
|||
|
|
h, err := Hessian(objective, xt, HessianOptions{})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Hessian of a real point: %v", err)
|
|||
|
|
}
|
|||
|
|
// Σ|q|² over a real q is Σq², whose Hessian is 2·I.
|
|||
|
|
for i := range 2 {
|
|||
|
|
if got := h.FloatAt(i*2 + i); math.Abs(got-2) > 1e-6 {
|
|||
|
|
t.Errorf("real Hessian diagonal[%d] = %v, want 2", i, got)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// refusesComplex runs fn and requires an error that names the
|
|||
|
|
// complex dtype, treating a panic as the failure it reports.
|
|||
|
|
func refusesComplex(t *testing.T, label string, fn func() error) {
|
|||
|
|
t.Helper()
|
|||
|
|
defer func() {
|
|||
|
|
if r := recover(); r != nil {
|
|||
|
|
t.Errorf("%s panicked: %v", label, r)
|
|||
|
|
}
|
|||
|
|
}()
|
|||
|
|
err := fn()
|
|||
|
|
if err == nil {
|
|||
|
|
t.Errorf("%s was accepted", label)
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
if !strings.Contains(err.Error(), "complex") {
|
|||
|
|
t.Errorf("%s error %q does not name the dtype", label, err)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestSecondOrderHelpersLeaveCallerGradients pins the internal reverse
|
|||
|
|
// passes of Hessian, HessianVectorProduct and MinimiseNewtonCG: the
|
|||
|
|
// objectives close over the caller's trainable tensors, and no path
|
|||
|
|
// (success or error) may mutate their accumulated gradients. AdjointODE
|
|||
|
|
// documents and implements the same guarantee.
|
|||
|
|
func TestSecondOrderHelpersLeaveCallerGradients(t *testing.T) {
|
|||
|
|
thetaVals := []float64{2, 3}
|
|||
|
|
presetVals := []float64{0.5, 0.25}
|
|||
|
|
thetaArr, err := core.FromFloats(thetaVals, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
xArr, err := core.FromFloats([]float64{1, -1}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
// Σ z_i²·θ_i: the Hessian is diag(2θ) = diag(4, 6).
|
|||
|
|
weighted := func(theta *Tensor) func(*Tensor) (*Tensor, error) {
|
|||
|
|
return func(z *Tensor) (*Tensor, error) {
|
|||
|
|
sq, err := z.Mul(z)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
w, err := sq.Mul(theta)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return w.Sum()
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
// Σ θ_i²: independent of the probes, the disconnected-objective case.
|
|||
|
|
constant := func(theta *Tensor) func(*Tensor) (*Tensor, error) {
|
|||
|
|
return func(*Tensor) (*Tensor, error) {
|
|||
|
|
sq, err := theta.Mul(theta)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return sq.Sum()
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
t.Run("hessian success", func(t *testing.T) {
|
|||
|
|
theta := FromArray(thetaArr, true)
|
|||
|
|
preset, presetBits := presetGrad(t, theta, presetVals)
|
|||
|
|
h, err := Hessian(weighted(theta), FromArray(xArr, false), HessianOptions{})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Hessian: %v", err)
|
|||
|
|
}
|
|||
|
|
for i := range 2 {
|
|||
|
|
if got, want := h.FloatAt(i*2+i), 2*thetaVals[i]; math.Abs(got-want) > 1e-8 {
|
|||
|
|
t.Errorf("Hessian diagonal[%d] = %v, want %v", i, got, want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
requireGradUntouched(t, "Hessian", theta, preset, presetBits)
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
t.Run("hessian error", func(t *testing.T) {
|
|||
|
|
theta := FromArray(thetaArr, true)
|
|||
|
|
preset, presetBits := presetGrad(t, theta, presetVals)
|
|||
|
|
if _, err := Hessian(constant(theta), FromArray(xArr, false), HessianOptions{}); err == nil {
|
|||
|
|
t.Fatal("expected the disconnected-objective error")
|
|||
|
|
}
|
|||
|
|
requireGradUntouched(t, "Hessian error path", theta, preset, presetBits)
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
t.Run("hvp success", func(t *testing.T) {
|
|||
|
|
theta := FromArray(thetaArr, true)
|
|||
|
|
preset, presetBits := presetGrad(t, theta, presetVals)
|
|||
|
|
v, err := core.FromFloats([]float64{1, 0.5}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
hv, err := HessianVectorProduct(weighted(theta), FromArray(xArr, false), FromArray(v, false), HessianOptions{})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("HessianVectorProduct: %v", err)
|
|||
|
|
}
|
|||
|
|
for i := range 2 {
|
|||
|
|
if got, want := hv.FloatAt(i), 2*thetaVals[i]*v.FloatAt(i); math.Abs(got-want) > 1e-6 {
|
|||
|
|
t.Errorf("H·v[%d] = %v, want %v", i, got, want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
requireGradUntouched(t, "HessianVectorProduct", theta, preset, presetBits)
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
t.Run("hvp error", func(t *testing.T) {
|
|||
|
|
theta := FromArray(thetaArr, true)
|
|||
|
|
preset, presetBits := presetGrad(t, theta, presetVals)
|
|||
|
|
v, err := core.FromFloats([]float64{1, 0.5}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
if _, err := HessianVectorProduct(constant(theta), FromArray(xArr, false), FromArray(v, false), HessianOptions{}); err == nil {
|
|||
|
|
t.Fatal("expected the disconnected-objective error")
|
|||
|
|
}
|
|||
|
|
requireGradUntouched(t, "HessianVectorProduct error path", theta, preset, presetBits)
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
t.Run("newtoncg success", func(t *testing.T) {
|
|||
|
|
theta := FromArray(thetaArr, true)
|
|||
|
|
preset, presetBits := presetGrad(t, theta, presetVals)
|
|||
|
|
// Σ (z − θ)² minimises at z = θ, where the gradient vanishes.
|
|||
|
|
objective := func(z *Tensor) (*Tensor, error) {
|
|||
|
|
d, err := z.Sub(theta)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
sq, err := d.Mul(d)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return sq.Sum()
|
|||
|
|
}
|
|||
|
|
x0, err := core.FromFloats([]float64{0, 0}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
got, fv, err := MinimiseNewtonCG(objective, x0, NewtonCGOptions{})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("MinimiseNewtonCG: %v", err)
|
|||
|
|
}
|
|||
|
|
for i := range 2 {
|
|||
|
|
if math.Abs(got.FloatAt(i)-thetaVals[i]) > 1e-8 {
|
|||
|
|
t.Errorf("minimiser[%d] = %v, want %v", i, got.FloatAt(i), thetaVals[i])
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
if math.Abs(fv) > 1e-12 {
|
|||
|
|
t.Errorf("value = %v, want 0", fv)
|
|||
|
|
}
|
|||
|
|
requireGradUntouched(t, "MinimiseNewtonCG", theta, preset, presetBits)
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
t.Run("newtoncg error", func(t *testing.T) {
|
|||
|
|
theta := FromArray(thetaArr, true)
|
|||
|
|
preset, presetBits := presetGrad(t, theta, presetVals)
|
|||
|
|
x0, err := core.FromFloats([]float64{1}, 1)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
if _, _, err := MinimiseNewtonCG(constant(theta), x0, NewtonCGOptions{}); err == nil {
|
|||
|
|
t.Fatal("expected the disconnected-objective error")
|
|||
|
|
}
|
|||
|
|
requireGradUntouched(t, "MinimiseNewtonCG error path", theta, preset, presetBits)
|
|||
|
|
})
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// presetGrad installs vals as theta's accumulated gradient and
|
|||
|
|
// returns the array the caller set together with a byte-exact snapshot
|
|||
|
|
// of its payload, so the check below can prove both that the very array
|
|||
|
|
// survived and that nothing wrote through it.
|
|||
|
|
func presetGrad(t *testing.T, theta *Tensor, vals []float64) (*core.Array, []uint64) {
|
|||
|
|
t.Helper()
|
|||
|
|
g, err := core.FromFloats(vals, theta.Data().Len())
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
theta.SetGrad(g)
|
|||
|
|
return g, gradBits(g)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// gradBits copies the raw float bits of a gradient array.
|
|||
|
|
func gradBits(a *core.Array) []uint64 {
|
|||
|
|
fs := a.RawFloats()
|
|||
|
|
out := make([]uint64, len(fs))
|
|||
|
|
for i, v := range fs {
|
|||
|
|
out[i] = math.Float64bits(v)
|
|||
|
|
}
|
|||
|
|
return out
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// requireGradUntouched requires the preset gradient to be the very
|
|||
|
|
// array that was set and to hold the exact bits the snapshot captured:
|
|||
|
|
// the helper under test must not write into it, replace it or clear it.
|
|||
|
|
func requireGradUntouched(t *testing.T, label string, theta *Tensor, want *core.Array, before []uint64) {
|
|||
|
|
t.Helper()
|
|||
|
|
got := theta.Grad()
|
|||
|
|
if got == nil {
|
|||
|
|
t.Errorf("%s: the caller's gradient was cleared", label)
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
if got != want {
|
|||
|
|
t.Errorf("%s: the caller's gradient array was replaced", label)
|
|||
|
|
}
|
|||
|
|
after := gradBits(got)
|
|||
|
|
if len(after) != len(before) {
|
|||
|
|
t.Errorf("%s: the caller's gradient changed length, %d to %d", label, len(before), len(after))
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
for i := range before {
|
|||
|
|
if after[i] != before[i] {
|
|||
|
|
t.Errorf("%s: gradient[%d] = %v, want the preset bits of %v", label, i,
|
|||
|
|
got.FloatAt(i), math.Float64frombits(before[i]))
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestBroadcastToRefusesIntOperands pins the dtype gate on BroadcastTo,
|
|||
|
|
// the one shape op that used to accept an int tensor and record a graph
|
|||
|
|
// node for it. The shape is validated first, so an impossible target
|
|||
|
|
// keeps reporting the mismatch (the probe below asserts the same
|
|||
|
|
// order).
|
|||
|
|
func TestBroadcastToRefusesIntOperands(t *testing.T) {
|
|||
|
|
src, err := core.FromInts([]int64{1, 2, 3}, 3)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromInts: %v", err)
|
|||
|
|
}
|
|||
|
|
if out, err := FromArray(src, true).BroadcastTo(2, 3); err == nil {
|
|||
|
|
t.Errorf("BroadcastTo accepted an int tensor: shape %v dtype %s",
|
|||
|
|
out.Data().Shape(), out.Data().Dtype())
|
|||
|
|
} else if !strings.Contains(err.Error(), "needs a float") {
|
|||
|
|
t.Errorf("int refusal = %v, want the dtype error", err)
|
|||
|
|
}
|
|||
|
|
one, err := core.FromInts([]int64{7}, 1)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromInts: %v", err)
|
|||
|
|
}
|
|||
|
|
if _, err := FromArray(one, true).BroadcastTo(2, 3); err == nil {
|
|||
|
|
t.Error("BroadcastTo accepted an int (1,) tensor")
|
|||
|
|
}
|
|||
|
|
square, err := core.FromInts([]int64{1, 2, 3, 4}, 2, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromInts: %v", err)
|
|||
|
|
}
|
|||
|
|
if _, err := FromArray(square, true).BroadcastTo(2, 3); err == nil {
|
|||
|
|
t.Error("an impossible broadcast target was accepted")
|
|||
|
|
} else if !strings.Contains(err.Error(), "cannot broadcast") {
|
|||
|
|
t.Errorf("impossible int broadcast = %v, want the shape refusal", err)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// A float tensor still broadcasts and differentiates; a complex one
|
|||
|
|
// remains inside the accepted dtypes.
|
|||
|
|
xf, err := core.FromFloats([]float64{1, 2, 3}, 3)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
xt := FromArray(xf, true)
|
|||
|
|
y, err := xt.BroadcastTo(2, 3)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("float BroadcastTo: %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 xt.Grad().Len() {
|
|||
|
|
if got := xt.Grad().FloatAt(i); got != 2 {
|
|||
|
|
t.Errorf("float broadcast gradient[%d] = %v, want 2", i, got)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
zc, err := core.FromComplexes([]complex128{1 + 1i, 2 - 0.5i}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromComplexes: %v", err)
|
|||
|
|
}
|
|||
|
|
zt := FromArray(zc, true)
|
|||
|
|
cz, err := zt.BroadcastTo(3, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("complex BroadcastTo: %v", err)
|
|||
|
|
}
|
|||
|
|
csq, err := cz.Abs2()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Abs2: %v", err)
|
|||
|
|
}
|
|||
|
|
closs, err := csq.Sum()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Sum: %v", err)
|
|||
|
|
}
|
|||
|
|
if err := closs.Backward(); err != nil {
|
|||
|
|
t.Fatalf("Backward: %v", err)
|
|||
|
|
}
|
|||
|
|
// L = Σ|broadcast(z)|² = 3·Σ|z|², so dL/dz̄ = 3z.
|
|||
|
|
for i := range 2 {
|
|||
|
|
if got, want := zt.Grad().ComplexAt(i), 3*zc.ComplexAt(i); got != want {
|
|||
|
|
t.Errorf("complex broadcast gradient[%d] = %v, want %v", i, got, want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestComplexMeanOfEmptyIsAnError pins the degenerate reduction: the
|
|||
|
|
// complex branch of Mean divides by the element count, so an empty
|
|||
|
|
// tensor used to answer 0/0 = NaN while the real branch errors loudly.
|
|||
|
|
func TestComplexMeanOfEmptyIsAnError(t *testing.T) {
|
|||
|
|
ec, err := core.FromComplexes([]complex128{}, 0)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromComplexes: %v", err)
|
|||
|
|
}
|
|||
|
|
m, cerr := FromArray(ec, true).Mean()
|
|||
|
|
if cerr == nil {
|
|||
|
|
t.Fatalf("complex Mean of an empty tensor returned %v instead of an error", m.Data().ComplexAt(0))
|
|||
|
|
}
|
|||
|
|
if !strings.Contains(cerr.Error(), "empty array has no mean") {
|
|||
|
|
t.Errorf("complex Mean error = %v, want the empty-reduction refusal", cerr)
|
|||
|
|
}
|
|||
|
|
er, err := core.FromFloats([]float64{}, 0)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
_, rerr := FromArray(er, true).Mean()
|
|||
|
|
if rerr == nil {
|
|||
|
|
t.Fatal("real Mean of an empty tensor was accepted")
|
|||
|
|
}
|
|||
|
|
// The two dtypes answer with one message.
|
|||
|
|
if rerr.Error() != cerr.Error() {
|
|||
|
|
t.Errorf("messages disagree: real %q, complex %q", rerr, cerr)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// A non-empty complex mean still reduces and differentiates.
|
|||
|
|
z, err := core.FromComplexes([]complex128{1 + 1i, 3 - 1i}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromComplexes: %v", err)
|
|||
|
|
}
|
|||
|
|
zt := FromArray(z, true)
|
|||
|
|
mean, err := zt.Mean()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Mean: %v", err)
|
|||
|
|
}
|
|||
|
|
if got, want := mean.Data().ComplexAt(0), complex(2, 0); got != want {
|
|||
|
|
t.Errorf("complex mean = %v, want %v", got, want)
|
|||
|
|
}
|
|||
|
|
loss, err := mean.Abs2()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Abs2: %v", err)
|
|||
|
|
}
|
|||
|
|
if err := loss.Backward(); err != nil {
|
|||
|
|
t.Fatalf("Backward: %v", err)
|
|||
|
|
}
|
|||
|
|
// L = |mean|², so dL/dz̄ = mean/2 per element.
|
|||
|
|
for i := range 2 {
|
|||
|
|
if got, want := zt.Grad().ComplexAt(i), complex(1, 0); got != want {
|
|||
|
|
t.Errorf("gradient[%d] = %v, want %v", i, got, want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestAbs2KeepsFloat32Width pins Abs2's real branch: a float32 operand
|
|||
|
|
// squares in float64 and stays float32, exactly as Pow(2) does,
|
|||
|
|
// instead of promoting the forward result to float64.
|
|||
|
|
func TestAbs2KeepsFloat32Width(t *testing.T) {
|
|||
|
|
vals := []float32{1.3, -2.7, 0.5}
|
|||
|
|
a, err := core.FromFloat32s(vals, 3)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat32s: %v", err)
|
|||
|
|
}
|
|||
|
|
xt := FromArray(a, true)
|
|||
|
|
sq, err := xt.Abs2()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Abs2: %v", err)
|
|||
|
|
}
|
|||
|
|
if got := sq.Data().Dtype(); got != core.Float32 {
|
|||
|
|
t.Errorf("float32 Abs2 output dtype = %s, want float32", got)
|
|||
|
|
}
|
|||
|
|
pw, err := xt.Pow(2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Pow: %v", err)
|
|||
|
|
}
|
|||
|
|
if got := pw.Data().Dtype(); got != core.Float32 {
|
|||
|
|
t.Errorf("float32 Pow(2) output dtype = %s, want float32", got)
|
|||
|
|
}
|
|||
|
|
for i, v := range vals {
|
|||
|
|
want := float32(float64(v) * float64(v))
|
|||
|
|
if got := sq.Data().FloatAt(i); got != float64(want) {
|
|||
|
|
t.Errorf("Abs2[%d] = %v, want the once-rounded %v", i, got, want)
|
|||
|
|
}
|
|||
|
|
if got := pw.Data().FloatAt(i); got != float64(want) {
|
|||
|
|
t.Errorf("Pow(2)[%d] = %v, want %v", i, got, want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// L = Σx² has dL/dx = 2x, and the float32 leaf keeps its width.
|
|||
|
|
loss, err := sq.Sum()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Sum: %v", err)
|
|||
|
|
}
|
|||
|
|
if err := loss.Backward(); err != nil {
|
|||
|
|
t.Fatalf("Backward: %v", err)
|
|||
|
|
}
|
|||
|
|
g := xt.Grad()
|
|||
|
|
if got := g.Dtype(); got != core.Float32 {
|
|||
|
|
t.Errorf("float32 leaf gradient dtype = %s, want float32", got)
|
|||
|
|
}
|
|||
|
|
for i, v := range vals {
|
|||
|
|
if got, want := g.FloatAt(i), 2*float64(v); math.Abs(got-want) > 1e-6 {
|
|||
|
|
t.Errorf("gradient[%d] = %v, want %v", i, got, want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// A float64 operand is untouched: float64 out, exact squares.
|
|||
|
|
af, err := core.FromFloats([]float64{1.3, -2.7}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
sqf, err := FromArray(af, true).Abs2()
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("Abs2: %v", err)
|
|||
|
|
}
|
|||
|
|
if got := sqf.Data().Dtype(); got != core.Float {
|
|||
|
|
t.Errorf("float64 Abs2 output dtype = %s, want float", got)
|
|||
|
|
}
|
|||
|
|
for i := range 2 {
|
|||
|
|
if got, want := sqf.Data().FloatAt(i), af.FloatAt(i)*af.FloatAt(i); got != want {
|
|||
|
|
t.Errorf("float64 Abs2[%d] = %v, want %v", i, got, want)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestNewtonCGConstantObjectiveHitsDisconnectedGuard covers
|
|||
|
|
// MinimiseNewtonCG's g == nil guard, which the suite missed: its
|
|||
|
|
// "constant objective" case slices the point itself, so a gradient
|
|||
|
|
// exists and the guard never fires. A constant built from an
|
|||
|
|
// independent graduated tensor (the disconnected-objective pattern) has no path to
|
|||
|
|
// the starting point at all.
|
|||
|
|
func TestNewtonCGConstantObjectiveHitsDisconnectedGuard(t *testing.T) {
|
|||
|
|
c, err := FromFloat64s([]float64{3}, true, 1)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloat64s: %v", err)
|
|||
|
|
}
|
|||
|
|
constant := func(*Tensor) (*Tensor, error) { return c.Mul(c) }
|
|||
|
|
x0, err := core.FromFloats([]float64{1}, 1)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
if _, _, err := MinimiseNewtonCG(constant, x0, NewtonCGOptions{MaxIterations: 2}); err == nil {
|
|||
|
|
t.Fatal("a genuinely constant objective minimised without error")
|
|||
|
|
} else if !strings.Contains(err.Error(), "does not depend on the starting point") {
|
|||
|
|
t.Fatalf("error = %v, want the disconnected-graph refusal", err)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// TestSampleHMCLeavesCallerGradients pins SampleHMC's internal gradient
|
|||
|
|
// evaluations: one leapfrog step runs one reverse pass, so a density
|
|||
|
|
// closing over a graduated tensor used to add a contribution per step
|
|||
|
|
// and one per proposal, on the success and the error path alike. The
|
|||
|
|
// pass commits nothing, so the closed-over tensor keeps its accumulated
|
|||
|
|
// gradient bit for bit, the guarantee Hessian, HessianVectorProduct,
|
|||
|
|
// MinimiseNewtonCG and AdjointODE document.
|
|||
|
|
func TestSampleHMCLeavesCallerGradients(t *testing.T) {
|
|||
|
|
thetaVals := []float64{2, 1.5}
|
|||
|
|
presetVals := []float64{0.5, -0.25}
|
|||
|
|
thetaArr, err := core.FromFloats(thetaVals, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
q0, err := core.FromFloats([]float64{0.5, -0.5}, 2)
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("FromFloats: %v", err)
|
|||
|
|
}
|
|||
|
|
// log π(q) = −½·Σ θ_i q_i²: a density that closes over a graduated
|
|||
|
|
// tensor the caller owns, so any committed pass shows up in θ.
|
|||
|
|
density := func(theta *Tensor) func(*Tensor) (*Tensor, error) {
|
|||
|
|
return func(q *Tensor) (*Tensor, error) {
|
|||
|
|
sq, err := q.Mul(q)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
w, err := sq.Mul(theta)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
s, err := w.Sum()
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return s.Scale(-0.5)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
t.Run("success", func(t *testing.T) {
|
|||
|
|
theta := FromArray(thetaArr, true)
|
|||
|
|
preset, bits := presetGrad(t, theta, presetVals)
|
|||
|
|
calls := 0
|
|||
|
|
counted := func(q *Tensor) (*Tensor, error) {
|
|||
|
|
calls++
|
|||
|
|
return density(theta)(q)
|
|||
|
|
}
|
|||
|
|
samples, err := SampleHMC(counted, q0,
|
|||
|
|
HMCOptions{Step: 0.1, Steps: 2, Samples: 2, Thin: 1, Seed: 11})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("SampleHMC: %v", err)
|
|||
|
|
}
|
|||
|
|
if got := samples.Shape(); got[0] != 2 || got[1] != 2 {
|
|||
|
|
t.Fatalf("samples shape %v, want [2 2]", got)
|
|||
|
|
}
|
|||
|
|
// One evaluation at q0 plus one per leapfrog step: the run really
|
|||
|
|
// did differentiate the density several times.
|
|||
|
|
if calls < 3 {
|
|||
|
|
t.Fatalf("SampleHMC ran %d gradient evaluations, want at least 3", calls)
|
|||
|
|
}
|
|||
|
|
requireGradUntouched(t, "SampleHMC", theta, preset, bits)
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
t.Run("midChainError", func(t *testing.T) {
|
|||
|
|
theta := FromArray(thetaArr, true)
|
|||
|
|
preset, bits := presetGrad(t, theta, presetVals)
|
|||
|
|
// The first evaluation at q0 succeeds, so a reverse pass has run;
|
|||
|
|
// every proposal then reports the state as outside the support,
|
|||
|
|
// which rejects the trajectory instead of aborting the run.
|
|||
|
|
calls, refused := 0, 0
|
|||
|
|
failing := func(q *Tensor) (*Tensor, error) {
|
|||
|
|
calls++
|
|||
|
|
if calls > 1 {
|
|||
|
|
refused++
|
|||
|
|
return nil, errf("outside the support")
|
|||
|
|
}
|
|||
|
|
return density(theta)(q)
|
|||
|
|
}
|
|||
|
|
samples, err := SampleHMC(failing, q0,
|
|||
|
|
HMCOptions{Step: 0.1, Steps: 3, Samples: 1, Thin: 1, Seed: 12})
|
|||
|
|
if err != nil {
|
|||
|
|
t.Fatalf("SampleHMC: %v", err)
|
|||
|
|
}
|
|||
|
|
if got := samples.Shape(); got[0] != 1 || got[1] != 2 {
|
|||
|
|
t.Fatalf("samples shape %v, want [1 2]", got)
|
|||
|
|
}
|
|||
|
|
if refused == 0 {
|
|||
|
|
t.Fatal("the density was never forced to fail mid-chain")
|
|||
|
|
}
|
|||
|
|
requireGradUntouched(t, "SampleHMC mid-chain error", theta, preset, bits)
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
t.Run("startError", func(t *testing.T) {
|
|||
|
|
theta := FromArray(thetaArr, true)
|
|||
|
|
preset, bits := presetGrad(t, theta, presetVals)
|
|||
|
|
fails := func(*Tensor) (*Tensor, error) { return nil, errf("no density at the start") }
|
|||
|
|
if _, err := SampleHMC(fails, q0,
|
|||
|
|
HMCOptions{Step: 0.1, Steps: 2, Samples: 1, Thin: 1, Seed: 13}); err == nil {
|
|||
|
|
t.Fatal("expected the start-time density error to be fatal")
|
|||
|
|
}
|
|||
|
|
requireGradUntouched(t, "SampleHMC start error", theta, preset, bits)
|
|||
|
|
})
|
|||
|
|
}
|