feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -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)
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user