feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+875
View File
@@ -0,0 +1,875 @@
// 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)
})
}