Files
tensor/grad/grad_guard_pins_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

876 lines
28 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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)
})
}