// Copyright (c) 2026 Petr Balvín (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) }) }