feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
+313
@@ -0,0 +1,313 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"slices"
|
||||
|
||||
ode "sourcedock.dev/petrbalvin/tensor/integrate"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Adjoint sensitivities of an initial value problem. Fitting the
|
||||
// parameters of a differential equation to data asks for dL/dθ when
|
||||
// the trajectory y(t; θ) is produced by an ODE solve, and
|
||||
// differentiating through every solver step is neither necessary nor
|
||||
// cheap. The adjoint method runs the dynamics once forward, then
|
||||
// integrates the adjoint state λ(t) = ∂L/∂y(t) backward along
|
||||
// λ' = −(∂f/∂y)ᵀλ from the loss gradient at the endpoint, carrying a
|
||||
// per-parameter accumulator β' = −(∂f/∂θ)ᵀλ alongside; both
|
||||
// Jacobian-vector products come from one automatic-differentiation
|
||||
// backward pass per evaluation. The cost is one more ODE solve
|
||||
// whatever the parameter count, which is what makes whole-trajectory
|
||||
// fitting tractable.
|
||||
|
||||
// AdjointODE differentiates the solution of y' = f(t, y) at t1 with
|
||||
// respect to the initial state and to the parameters the given
|
||||
// function closes over. lossGrad is ∂L/∂y(t1), the seed the loss
|
||||
// itself contributes; the return values are ∂L/∂y0 and, parallel to
|
||||
// params, ∂L/∂θk, each shaped like its parameter. Backward-in-time
|
||||
// problems (t1 < t0) work.
|
||||
//
|
||||
// The forward trajectory is recorded at the adaptive solver's accepted
|
||||
// steps and handed to the backward pass through cubic Hermite
|
||||
// interpolation, fourth-order accurate like the Dormand-Prince pair
|
||||
// that produced it; the augmented adjoint system is integrated by the
|
||||
// same adaptive solver in reverse. The parameters' own accumulated
|
||||
// gradients are left untouched: every evaluation runs a reverse pass
|
||||
// that computes into a local map and commits nothing, and the answers
|
||||
// travel out as return values instead.
|
||||
//
|
||||
// A nil function, a non-vector start or loss seed, a parameter that
|
||||
// does not require grad, and an f that ignores its state or returns a
|
||||
// wrong shape are errors, never silent zeros.
|
||||
func AdjointODE(f func(t float64, y *Tensor) (*Tensor, error),
|
||||
params []*Tensor, t0, t1 float64, y0 *core.Array,
|
||||
lossGrad *core.Array, opts ode.ODEOptions) (*core.Array, []*core.Array, error) {
|
||||
const name = "AdjointODE"
|
||||
if f == nil {
|
||||
return nil, nil, errf("%s: f must not be nil", name)
|
||||
}
|
||||
if y0 == nil || y0.NDim() != 1 || y0.Len() == 0 {
|
||||
return nil, nil, errf("%s: the state must be a non-empty vector", name)
|
||||
}
|
||||
if y0.Dtype() == core.Complex {
|
||||
return nil, nil, errf("%s: complex states are not supported", name)
|
||||
}
|
||||
dim := y0.Len()
|
||||
if lossGrad == nil || lossGrad.NDim() != 1 || lossGrad.Len() != dim {
|
||||
return nil, nil, errf("%s: lossGrad must be a vector of length %d", name, dim)
|
||||
}
|
||||
sizes := make([]int, len(params))
|
||||
total := 0
|
||||
for k, p := range params {
|
||||
if p == nil || !p.RequiresGrad() {
|
||||
return nil, nil, errf("%s: parameter %d does not require grad", name, k)
|
||||
}
|
||||
if p.Data().NDim() == 0 || p.Data().Len() == 0 {
|
||||
return nil, nil, errf("%s: parameter %d must be a non-empty tensor of rank at least 1", name, k)
|
||||
}
|
||||
sizes[k] = p.Data().Len()
|
||||
total += sizes[k]
|
||||
}
|
||||
if t0 == t1 {
|
||||
// No dynamics: the endpoint is the initial state.
|
||||
seed := flatFloats(lossGrad)
|
||||
gradY0, err := core.FromFloats(seed, dim)
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
return gradY0, zeroBlocks(params, sizes), nil
|
||||
}
|
||||
|
||||
// Forward pass with the dynamics evaluated for value only.
|
||||
forward := func(t float64, ya *core.Array) (*core.Array, error) {
|
||||
out, err := f(t, FromArray(ya, false))
|
||||
if err != nil {
|
||||
return nil, errf("%s: %w", name, err)
|
||||
}
|
||||
return out.Data(), nil
|
||||
}
|
||||
times, states, err := ode.IntegrateODESteps(forward, t0, t1, y0, opts)
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
// Cubic Hermite interpolation wants the derivative of the dynamics
|
||||
// at every recorded node; one detached pass supplies them all.
|
||||
slopes := make([][]float64, len(times))
|
||||
for i, node := range states {
|
||||
out, err := f(times[i], FromArray(node, false))
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
if out.Data().NDim() != 1 || out.Data().Len() != dim {
|
||||
return nil, nil, errf("%s: f returned shape %s, want a vector of length %d",
|
||||
name, prettyShape(out.Data().Shape()), dim)
|
||||
}
|
||||
slopes[i] = flatFloats(out.Data())
|
||||
}
|
||||
trace := newODETrace(times, states, slopes, dim)
|
||||
|
||||
vjp := func(t float64, y []float64, lam []float64) ([]float64, []float64, error) {
|
||||
data, err := core.FromFloats(y, dim)
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
yLeaf := FromArray(data, true)
|
||||
out, err := f(t, yLeaf)
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
if out.Data().NDim() != 1 || out.Data().Len() != dim {
|
||||
return nil, nil, errf("%s: f returned shape %s, want a vector of length %d",
|
||||
name, prettyShape(out.Data().Shape()), dim)
|
||||
}
|
||||
seedData, err := core.FromFloats(lam, dim)
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
weighted, err := out.Mul(FromArray(seedData, false))
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
summed, err := weighted.Sum()
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
// The Jacobian-vector products come from the reverse pass
|
||||
// alone: reverseGrads computes every reached tensor's gradient
|
||||
// into a local map and commits nothing, so no leaf the dynamics
|
||||
// close over, in params or out, is written at all, and the
|
||||
// callers' own accumulated gradients survive untouched.
|
||||
grads, err := summed.reverseGrads()
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
g := grads[yLeaf]
|
||||
if g == nil || g.Len() != dim {
|
||||
return nil, nil, errf("%s: f did not yield a state gradient of length %d", name, dim)
|
||||
}
|
||||
gY := flatFloats(g)
|
||||
gTheta := make([]float64, 0, total)
|
||||
for k, p := range params {
|
||||
now := grads[p]
|
||||
if now == nil {
|
||||
return nil, nil, errf("%s: parameter %d is disconnected from the dynamics, no gradient flowed", name, k)
|
||||
}
|
||||
nowFloats := flatFloats(now)
|
||||
gTheta = append(gTheta, nowFloats...)
|
||||
}
|
||||
return gY, gTheta, nil
|
||||
}
|
||||
|
||||
rhs := func(t float64, z *core.Array) (*core.Array, error) {
|
||||
lam := make([]float64, dim)
|
||||
if !z.Strided() && z.Dtype() == core.Float {
|
||||
copy(lam, z.RawFloats()[:dim])
|
||||
} else {
|
||||
for i := range lam {
|
||||
lam[i] = z.FloatAt(i)
|
||||
}
|
||||
}
|
||||
yArr, err := trace.at(t)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
gY, gTheta, err := vjp(t, yArr, lam)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]float64, dim+total)
|
||||
for i := range gY {
|
||||
out[i] = -gY[i]
|
||||
}
|
||||
for i := range gTheta {
|
||||
out[dim+i] = -gTheta[i]
|
||||
}
|
||||
return core.FromFloats(out, dim+total)
|
||||
}
|
||||
zStart := make([]float64, dim+total)
|
||||
if !lossGrad.Strided() && lossGrad.Dtype() == core.Float {
|
||||
copy(zStart, lossGrad.RawFloats()[:dim])
|
||||
} else {
|
||||
for i := range dim {
|
||||
zStart[i] = lossGrad.FloatAt(i)
|
||||
}
|
||||
}
|
||||
zEnd, err := ode.IntegrateODE(rhs, t1, t0, mustFromFloats(zStart, dim+total), opts)
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
seed := make([]float64, dim)
|
||||
for i := range seed {
|
||||
seed[i] = zEnd.FloatAt(i)
|
||||
}
|
||||
gradY0, err := core.FromFloats(seed, dim)
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
blocks := make([]*core.Array, len(params))
|
||||
offset := dim
|
||||
for k, p := range params {
|
||||
block := make([]float64, sizes[k])
|
||||
if !zEnd.Strided() && zEnd.Dtype() == core.Float {
|
||||
copy(block, zEnd.RawFloats()[offset:offset+sizes[k]])
|
||||
} else {
|
||||
for j := range block {
|
||||
block[j] = zEnd.FloatAt(offset + j)
|
||||
}
|
||||
}
|
||||
shaped, err := core.FromFloats(block, p.Data().Shape()...)
|
||||
if err != nil {
|
||||
return nil, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
blocks[k] = shaped
|
||||
offset += sizes[k]
|
||||
}
|
||||
return gradY0, blocks, nil
|
||||
}
|
||||
|
||||
// zeroBlocks builds zero arrays shaped like each parameter, the
|
||||
// degenerate answer of a span with no dynamics.
|
||||
func zeroBlocks(params []*Tensor, sizes []int) []*core.Array {
|
||||
blocks := make([]*core.Array, len(params))
|
||||
for k, p := range params {
|
||||
zeros := make([]float64, sizes[k])
|
||||
blocks[k], _ = core.FromFloats(zeros, p.Data().Shape()...)
|
||||
}
|
||||
return blocks
|
||||
}
|
||||
|
||||
// mustFromFloats wraps a plain construction that cannot fail: the
|
||||
// length always matches the single-dimension shape.
|
||||
func mustFromFloats(vals []float64, n int) *core.Array {
|
||||
a, _ := core.FromFloats(vals, n)
|
||||
return a
|
||||
}
|
||||
|
||||
// odeTrace is a recorded forward trajectory, kept ascending in time,
|
||||
// with the dynamics' derivative at every node so the interpolation is
|
||||
// cubic Hermite.
|
||||
type odeTrace struct {
|
||||
times []float64
|
||||
states [][]float64
|
||||
slopes [][]float64
|
||||
dim int
|
||||
// buf is the interpolation buffer at reuses: the solver calls at
|
||||
// once per right-hand-side evaluation, so the array is borrowed
|
||||
// rather than allocated each time. Callers must consume the answer
|
||||
// before the next call.
|
||||
buf []float64
|
||||
}
|
||||
|
||||
// newODETrace flattens and time-orders the recorded nodes.
|
||||
func newODETrace(times []float64, states []*core.Array, slopes [][]float64, dim int) *odeTrace {
|
||||
tr := &odeTrace{times: append([]float64{}, times...), slopes: slopes, dim: dim}
|
||||
tr.states = make([][]float64, len(states))
|
||||
for i, s := range states {
|
||||
tr.states[i] = flatFloats(s)
|
||||
}
|
||||
if len(times) > 1 && times[1] < times[0] {
|
||||
// A backward-in-time pass records descending nodes; flip to
|
||||
// ascending so the search below has one convention.
|
||||
for i, j := 0, len(tr.times)-1; i < j; i, j = i+1, j-1 {
|
||||
tr.times[i], tr.times[j] = tr.times[j], tr.times[i]
|
||||
tr.states[i], tr.states[j] = tr.states[j], tr.states[i]
|
||||
tr.slopes[i], tr.slopes[j] = tr.slopes[j], tr.slopes[i]
|
||||
}
|
||||
}
|
||||
return tr
|
||||
}
|
||||
|
||||
// at interpolates the state at t by cubic Hermite over the enclosing
|
||||
// recorded interval, fourth-order accurate in the step size and never
|
||||
// leaving the recorded span. A trace with fewer than two nodes offers
|
||||
// no interval to interpolate over, so the accessor refuses instead of
|
||||
// indexing out of range. The returned slice belongs to the trace and
|
||||
// stays valid only until the next call.
|
||||
func (tr *odeTrace) at(t float64) ([]float64, error) {
|
||||
n := len(tr.times)
|
||||
if n < 2 {
|
||||
return nil, errf("AdjointODE: the solver recorded %d trajectory nodes, the adjoint interpolation needs at least two", n)
|
||||
}
|
||||
idx, _ := slices.BinarySearch(tr.times, t)
|
||||
i := min(max(idx-1, 0), n-2)
|
||||
h := tr.times[i+1] - tr.times[i]
|
||||
s := (t - tr.times[i]) / h
|
||||
h00 := s*s*(2*s-3) + 1
|
||||
h10 := s * (s - 1) * (s - 1)
|
||||
h01 := s * s * (3 - 2*s)
|
||||
h11 := s * s * (s - 1)
|
||||
if len(tr.buf) != tr.dim {
|
||||
tr.buf = make([]float64, tr.dim)
|
||||
}
|
||||
out := tr.buf
|
||||
si, si1 := tr.states[i], tr.states[i+1]
|
||||
li, li1 := tr.slopes[i], tr.slopes[i+1]
|
||||
for j := range out {
|
||||
out[j] = h00*si[j] + h01*si1[j] +
|
||||
h*(h10*li[j]+h11*li1[j])
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
@@ -0,0 +1,349 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
ode "sourcedock.dev/petrbalvin/tensor/integrate"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// decayWith builds f for y' = −θy with θ a scalar parameter leaf.
|
||||
func decayWith(theta *Tensor) func(t float64, y *Tensor) (*Tensor, error) {
|
||||
return func(t float64, y *Tensor) (*Tensor, error) {
|
||||
rate, err := theta.Neg()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return y.Mul(rate)
|
||||
}
|
||||
}
|
||||
|
||||
// scalarVector returns a length-1 real array.
|
||||
func scalarVector(t *testing.T, v float64) *core.Array {
|
||||
t.Helper()
|
||||
a, err := core.FromFloats([]float64{v}, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// TestAdjointODEDecay differentiates y' = −θy with L = y(1): the exact
|
||||
// sensitivities are dL/dy0 = e^{−θ} and dL/dθ = −e^{−θ}, and a run
|
||||
// with no parameters at all must still answer the initial-state one.
|
||||
func TestAdjointODEDecay(t *testing.T) {
|
||||
const want = 0.4965853037914095 // e^{−0.7}
|
||||
theta, err := FromFloat64s([]float64{0.7}, true, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0, 1,
|
||||
scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||||
if err != nil {
|
||||
t.Fatalf("AdjointODE: %v", err)
|
||||
}
|
||||
if math.Abs(gradY0.FloatAt(0)-want) > 1e-6 {
|
||||
t.Fatalf("dL/dy0 = %.14g, want %.14g", gradY0.FloatAt(0), want)
|
||||
}
|
||||
if len(blocks) != 1 || math.Abs(blocks[0].FloatAt(0)+want) > 1e-6 {
|
||||
t.Fatalf("dL/dθ = %v, want −%.14g", blocks[0].FloatAt(0), want)
|
||||
}
|
||||
// The parameter's own accumulated gradient must be untouched.
|
||||
if theta.Grad() != nil {
|
||||
t.Fatal("AdjointODE must leave the parameters' gradients untouched")
|
||||
}
|
||||
solo, soloBlocks, err := AdjointODE(decayWith(theta), nil, 0, 1,
|
||||
scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||||
if err != nil {
|
||||
t.Fatalf("AdjointODE without parameters: %v", err)
|
||||
}
|
||||
if len(soloBlocks) != 0 || math.Abs(solo.FloatAt(0)-want) > 1e-6 {
|
||||
t.Fatalf("parameter-free run: dL/dy0 = %.14g with %d blocks", solo.FloatAt(0), len(soloBlocks))
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjointODEOscillator differentiates the harmonic oscillator
|
||||
// y” = −ω²y with L = y(1) and y0 = (1, 0): dL/dω = −sin(ω),
|
||||
// dL/dy0 = (cos ω, sin ω/ω).
|
||||
func TestAdjointODEOscillator(t *testing.T) {
|
||||
omega, err := FromFloat64s([]float64{1.3}, true, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
f := func(t float64, y *Tensor) (*Tensor, error) {
|
||||
u, err := y.Slice(0, 0, 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
v, err := y.Slice(0, 1, 2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := omega.Mul(omega)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
acc, err := u.Mul(sq)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
drag, err := acc.Scale(-1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return v.Concat(drag, 0)
|
||||
}
|
||||
y0, err := core.FromFloats([]float64{1, 0}, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
seed, err := core.FromFloats([]float64{1, 0}, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
gradY0, blocks, err := AdjointODE(f, []*Tensor{omega}, 0, 1, y0, seed,
|
||||
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||||
if err != nil {
|
||||
t.Fatalf("AdjointODE: %v", err)
|
||||
}
|
||||
if math.Abs(gradY0.FloatAt(0)-math.Cos(1.3)) > 1e-6 {
|
||||
t.Fatalf("dL/du0 = %.14g, want cos(1.3) = %.14g", gradY0.FloatAt(0), math.Cos(1.3))
|
||||
}
|
||||
if math.Abs(gradY0.FloatAt(1)-math.Sin(1.3)/1.3) > 1e-6 {
|
||||
t.Fatalf("dL/dv0 = %.14g, want sin(1.3)/1.3 = %.14g",
|
||||
gradY0.FloatAt(1), math.Sin(1.3)/1.3)
|
||||
}
|
||||
if math.Abs(blocks[0].FloatAt(0)+math.Sin(1.3)) > 1e-6 {
|
||||
t.Fatalf("dL/dω = %.14g, want −sin(1.3) = %.14g", blocks[0].FloatAt(0), -math.Sin(1.3))
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjointODEBackwardTime runs the forward pass itself backwards
|
||||
// (t1 < t0): y(t) = e^{−θ(t−1)} from y(1) = 1 has y(0) = e^θ and both
|
||||
// sensitivities equal e^θ.
|
||||
func TestAdjointODEBackwardTime(t *testing.T) {
|
||||
const want = 1.6487212707001282 // e^{0.5}
|
||||
theta, _ := FromFloat64s([]float64{0.5}, true, 1)
|
||||
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 1, 0,
|
||||
scalarVector(t, 1), scalarVector(t, 1), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||||
if err != nil {
|
||||
t.Fatalf("AdjointODE: %v", err)
|
||||
}
|
||||
if math.Abs(gradY0.FloatAt(0)-want) > 1e-6 {
|
||||
t.Fatalf("dL/dy0 = %.14g, want %.14g", gradY0.FloatAt(0), want)
|
||||
}
|
||||
if math.Abs(blocks[0].FloatAt(0)-want) > 1e-6 {
|
||||
t.Fatalf("dL/dθ = %.14g, want %.14g", blocks[0].FloatAt(0), want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjointODEFiniteDifference checks a nonlinear two-parameter
|
||||
// system against central differences of the forward solve itself,
|
||||
// perturbing the parameter leaves around the adjoint run.
|
||||
func TestAdjointODEFiniteDifference(t *testing.T) {
|
||||
th1, _ := FromFloat64s([]float64{0.8}, true, 1)
|
||||
th2, _ := FromFloat64s([]float64{1.1}, true, 1)
|
||||
f := func(t float64, y *Tensor) (*Tensor, error) {
|
||||
y1, err := y.Slice(0, 0, 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
y2, err := y.Slice(0, 1, 2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
drag, err := y1.Mul(th1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
drive, err := y2.Mul(th2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r1, err := drive.Sub(drag)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pool, err := y1.Mul(y2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r2, err := pool.Scale(-1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return r1.Concat(r2, 0)
|
||||
}
|
||||
y0, _ := core.FromFloats([]float64{1, 0.5}, 2)
|
||||
seed, _ := core.FromFloats([]float64{1, 2}, 2)
|
||||
gradY0, blocks, err := AdjointODE(f, []*Tensor{th1, th2}, 0, 1, y0, seed,
|
||||
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-13})
|
||||
if err != nil {
|
||||
t.Fatalf("AdjointODE: %v", err)
|
||||
}
|
||||
// Central differences on the loss y1(1) + 2·y2(1), with the leaf
|
||||
// data swapped out for the perturbed values.
|
||||
forwardLoss := func() float64 {
|
||||
end, err := ode.IntegrateODE(func(t float64, ya *core.Array) (*core.Array, error) {
|
||||
out, err := f(t, FromArray(ya, false))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out.Data(), nil
|
||||
}, 0, 1, y0, ode.ODEOptions{RelTol: 1e-11, AbsTol: 1e-15})
|
||||
if err != nil {
|
||||
t.Fatalf("forward solve: %v", err)
|
||||
}
|
||||
return end.FloatAt(0) + 2*end.FloatAt(1)
|
||||
}
|
||||
perturb := func(p *Tensor, eps float64) {
|
||||
swapped, _ := core.FromFloats([]float64{p.Data().FloatAt(0) + eps}, 1)
|
||||
p.ReplaceWith(swapped)
|
||||
}
|
||||
restore := func(p *Tensor, v float64) {
|
||||
orig, _ := core.FromFloats([]float64{v}, 1)
|
||||
p.ReplaceWith(orig)
|
||||
}
|
||||
const eps = 1e-5
|
||||
for k, p := range []*Tensor{th1, th2} {
|
||||
orig := p.Data().FloatAt(0)
|
||||
perturb(p, eps)
|
||||
up := forwardLoss()
|
||||
perturb(p, -2*eps)
|
||||
down := forwardLoss()
|
||||
restore(p, orig)
|
||||
fd := (up - down) / (2 * eps)
|
||||
got := blocks[k].FloatAt(0)
|
||||
if math.Abs(got-fd) > 1e-3*math.Max(1, math.Abs(fd)) {
|
||||
t.Fatalf("dL/dθ%d: adjoint %.8g, finite difference %.8g", k+1, got, fd)
|
||||
}
|
||||
}
|
||||
if gradY0.Len() != 2 {
|
||||
t.Fatalf("dL/dy0 has length %d, want 2", gradY0.Len())
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjointODEDegenerateSpan returns the loss seed unchanged and
|
||||
// zero parameter gradients when the span carries no dynamics.
|
||||
func TestAdjointODEDegenerateSpan(t *testing.T) {
|
||||
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
|
||||
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0.5, 0.5,
|
||||
scalarVector(t, 1), scalarVector(t, 2), ode.ODEOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("AdjointODE: %v", err)
|
||||
}
|
||||
if gradY0.FloatAt(0) != 2 {
|
||||
t.Fatalf("dL/dy0 = %.14g, want the seed 2", gradY0.FloatAt(0))
|
||||
}
|
||||
if blocks[0].FloatAt(0) != 0 {
|
||||
t.Fatalf("dL/dθ = %.14g, want 0", blocks[0].FloatAt(0))
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjointODEErrors pins the validation contract.
|
||||
func TestAdjointODEErrors(t *testing.T) {
|
||||
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
|
||||
y0 := scalarVector(t, 1)
|
||||
seed := scalarVector(t, 1)
|
||||
rank2, _ := core.FromFloats([]float64{1, 1}, 1, 2)
|
||||
if _, _, err := AdjointODE(nil, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
|
||||
t.Fatal("expected an error for a nil function")
|
||||
}
|
||||
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, nil, seed, ode.ODEOptions{}); err == nil {
|
||||
t.Fatal("expected an error for a nil state")
|
||||
}
|
||||
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, rank2, seed, ode.ODEOptions{}); err == nil {
|
||||
t.Fatal("expected an error for a rank-2 state")
|
||||
}
|
||||
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, y0, nil, ode.ODEOptions{}); err == nil {
|
||||
t.Fatal("expected an error for a nil loss seed")
|
||||
}
|
||||
badSeed, _ := core.FromFloats([]float64{1, 1}, 2)
|
||||
if _, _, err := AdjointODE(decayWith(theta), nil, 0, 1, y0, badSeed, ode.ODEOptions{}); err == nil {
|
||||
t.Fatal("expected an error for a loss seed of the wrong length")
|
||||
}
|
||||
frozen := FromArray(scalarVector(t, 0.7), false)
|
||||
if _, _, err := AdjointODE(decayWith(frozen), []*Tensor{frozen}, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
|
||||
t.Fatal("expected an error for a parameter that does not require grad")
|
||||
}
|
||||
wrongShape := func(t float64, y *Tensor) (*Tensor, error) {
|
||||
three, _ := core.FromFloats([]float64{1, 2, 3}, 3)
|
||||
return FromArray(three, false), nil
|
||||
}
|
||||
if _, _, err := AdjointODE(wrongShape, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
|
||||
t.Fatal("expected an error for a wrong-shaped derivative")
|
||||
}
|
||||
ignoresState := func(t float64, y *Tensor) (*Tensor, error) {
|
||||
one, _ := core.FromFloats([]float64{1}, 1)
|
||||
return FromArray(one, false), nil
|
||||
}
|
||||
if _, _, err := AdjointODE(ignoresState, nil, 0, 1, y0, seed, ode.ODEOptions{}); err == nil {
|
||||
t.Fatal("expected an error when f ignores its state")
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjointODERestoresGradientsOnError pins the error-path contract:
|
||||
// when f consumes the parameter but ignores the state, no state
|
||||
// gradient can flow and the run errors, and the caller's own
|
||||
// accumulated parameter gradient must come back untouched instead of
|
||||
// being polluted by the aborted pass: the reverse sweep is not run at
|
||||
// all, so nothing writes the leaf gradients on the way out.
|
||||
func TestAdjointODERestoresGradientsOnError(t *testing.T) {
|
||||
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
|
||||
preset, _ := core.FromFloats([]float64{3}, 1)
|
||||
theta.SetGrad(preset)
|
||||
f := func(t float64, y *Tensor) (*Tensor, error) {
|
||||
return theta.Scale(2)
|
||||
}
|
||||
if _, _, err := AdjointODE(f, []*Tensor{theta}, 0, 1, scalarVector(t, 1),
|
||||
scalarVector(t, 1), ode.ODEOptions{}); err == nil {
|
||||
t.Fatal("expected an error when f ignores its state")
|
||||
}
|
||||
if theta.Grad() == nil || theta.Grad().FloatAt(0) != 3 {
|
||||
t.Fatalf("the parameter's gradient was not restored: %v, want 3", theta.Grad())
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjointODEDisconnectedParameter pins that a parameter the
|
||||
// dynamics never touch is an error, not a silent zero gradient block.
|
||||
func TestAdjointODEDisconnectedParameter(t *testing.T) {
|
||||
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
|
||||
f := func(t float64, y *Tensor) (*Tensor, error) {
|
||||
return y.Scale(2)
|
||||
}
|
||||
_, blocks, err := AdjointODE(f, []*Tensor{theta}, 0, 1, scalarVector(t, 1),
|
||||
scalarVector(t, 1), ode.ODEOptions{})
|
||||
if err == nil {
|
||||
t.Fatal("expected an error for a parameter disconnected from the dynamics")
|
||||
}
|
||||
if blocks != nil {
|
||||
t.Fatalf("an errored run returned blocks: %v", blocks)
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjointODESuccessKeepsPresetGradient pins the same restore on
|
||||
// the success path: a preset accumulated gradient survives the run.
|
||||
func TestAdjointODESuccessKeepsPresetGradient(t *testing.T) {
|
||||
theta, _ := FromFloat64s([]float64{0.5}, true, 1)
|
||||
preset, _ := core.FromFloats([]float64{7}, 1)
|
||||
theta.SetGrad(preset)
|
||||
gradY0, blocks, err := AdjointODE(decayWith(theta), []*Tensor{theta}, 0, 1,
|
||||
scalarVector(t, 1), scalarVector(t, 2), ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||||
if err != nil {
|
||||
t.Fatalf("AdjointODE: %v", err)
|
||||
}
|
||||
const want = 0.6065306597126334 // e^{−0.5}
|
||||
if math.Abs(gradY0.FloatAt(0)-2*want) > 1e-6 {
|
||||
t.Fatalf("dL/dy0 = %.14g, want 2e^{−0.5}", gradY0.FloatAt(0))
|
||||
}
|
||||
if theta.Grad() == nil || theta.Grad().FloatAt(0) != 7 {
|
||||
t.Fatalf("the preset gradient did not survive a successful run: %v", theta.Grad())
|
||||
}
|
||||
if blocks[0].FloatAt(0) >= 0 {
|
||||
t.Fatalf("dL/dθ = %v, want negative", blocks[0].FloatAt(0))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,918 @@
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Backward-coverage pins for the graph kernels: every test here checks
|
||||
// a committed leaf gradient against a closed-form identity or a central
|
||||
// difference of the summed loss, on combinations the op-level tests do
|
||||
// not build.
|
||||
|
||||
// totalLoss sums a non-scalar loss so a central difference matches the
|
||||
// all-ones seed Backward uses.
|
||||
func totalLoss(l *Tensor) float64 {
|
||||
s := 0.0
|
||||
d := l.Data()
|
||||
for i := range d.Len() {
|
||||
s += d.FloatAt(i)
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func abs2c(z complex128) float64 { return real(z)*real(z) + imag(z)*imag(z) }
|
||||
|
||||
func mustRecover(vals []float64, sh ...int) *core.Array {
|
||||
a, _ := core.FromFloats(vals, sh...)
|
||||
return a
|
||||
}
|
||||
|
||||
func mustRecoverComplex(zs []complex128, shape ...int) *core.Array {
|
||||
if len(shape) == 0 {
|
||||
shape = []int{len(zs)}
|
||||
}
|
||||
w := append([]complex128(nil), zs...)
|
||||
a, _ := core.ComplexFromArray(w, shape...)
|
||||
return a
|
||||
}
|
||||
|
||||
// probeCheckCentral builds the loss from base, backpropagates once and
|
||||
// compares every element of the committed gradient against central
|
||||
// differences of the summed loss.
|
||||
func probeCheckCentral(t *testing.T, base *Tensor, name string, build func() (*Tensor, error), tol float64) {
|
||||
t.Helper()
|
||||
base.ZeroGrad()
|
||||
loss, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
grads := flatFloats(base.Grad())
|
||||
orig := flatFloats(base.Data())
|
||||
sh := base.Data().Shape()
|
||||
for i := range orig {
|
||||
const h = 1e-6
|
||||
plus := append([]float64(nil), orig...)
|
||||
plus[i] += h
|
||||
base.ReplaceWith(mustRecover(plus, sh...))
|
||||
lp, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
minus := append([]float64(nil), orig...)
|
||||
minus[i] -= h
|
||||
base.ReplaceWith(mustRecover(minus, sh...))
|
||||
lm, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
|
||||
if math.Abs(grads[i]-num) > tol*math.Max(1, math.Abs(num)) {
|
||||
t.Errorf("%s grad[%d] = %g, central %g", name, i, grads[i], num)
|
||||
}
|
||||
base.ReplaceWith(mustRecover(orig, sh...))
|
||||
}
|
||||
}
|
||||
|
||||
// TestFFTParsevalGradient pins the FFT backward through Parseval's
|
||||
// theorem: L = sum |FFT(x)|^2 has gradient n·2x for real x, because the
|
||||
// unnormalised forward DFT scales the energy by n.
|
||||
func TestFFTParsevalGradient(t *testing.T) {
|
||||
const n = 16
|
||||
vals := make([]float64, n)
|
||||
for i := range vals {
|
||||
vals[i] = math.Sin(0.7*float64(i)) + 0.3*float64(i%4)
|
||||
}
|
||||
x, _ := FromFloat64s(vals, true, n)
|
||||
f, _ := x.FFT()
|
||||
abs2, _ := f.Abs2()
|
||||
loss, _ := abs2.Sum()
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
g := x.Grad()
|
||||
for i := range n {
|
||||
want := 2 * float64(n) * vals[i]
|
||||
if got := g.FloatAt(i); math.Abs(got-want) > 1e-8*math.Max(1, math.Abs(want)) {
|
||||
t.Fatalf("Parseval: g[%d] = %g, want %g", i, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestFFT2ParsevalGradient is Parseval's identity for the rank-2
|
||||
// transform, where the backward's scale is the full element count.
|
||||
func TestFFT2ParsevalGradient(t *testing.T) {
|
||||
vals := []float64{1, 2, -1, 0.5, 3, -2, 0.25, 1.5, -0.5, 2, 1, -3}
|
||||
x, _ := FromFloat64s(vals, true, 3, 4)
|
||||
f, err := x.FFT2()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
a, err := f.Abs2()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loss, err := a.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range vals {
|
||||
want := 2 * float64(12) * vals[i]
|
||||
if got := x.Grad().FloatAt(i); math.Abs(got-want) > 1e-8*math.Max(1, math.Abs(want)) {
|
||||
t.Fatalf("FFT2 Parseval g[%d] = %g, want %g", i, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestSpectralHalfSpectrumBackward pins the RFFT and IRFFT backwards
|
||||
// against central differences, the half-spectrum combinatorics
|
||||
// included: the doubled real part on the mirrored bins and the halved
|
||||
// self-mirrored bins of the inverse.
|
||||
func TestSpectralHalfSpectrumBackward(t *testing.T) {
|
||||
x, _ := FromFloat64s([]float64{1, -2, 0.5, 3, -1.5, 0.25, 2, -0.75}, true, 8)
|
||||
buildF := func() (*Tensor, error) {
|
||||
h, err := x.RFFT()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
a, err := h.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return a.Sum()
|
||||
}
|
||||
probeCheckCentral(t, x, "rfft", buildF, 1e-4)
|
||||
|
||||
// IRFFT: the leaf is the half spectrum; the loss is the summed
|
||||
// square of the real signal. The complex leaf's gradient is checked
|
||||
// against a Wirtinger central difference of the real loss:
|
||||
// dL/dzbar_j = (dL/dRe_j + i dL/dIm_j)/2, each part probed by
|
||||
// perturbing that component.
|
||||
zr := []float64{3, -1, 0.5, 2, 1.5}
|
||||
zi := []float64{0, 1, -0.5, 0.25, -1}
|
||||
zs := make([]complex128, 5)
|
||||
for i := range zs {
|
||||
zs[i] = complex(zr[i], zi[i])
|
||||
}
|
||||
za, _ := core.ComplexFromArray(zs, 5)
|
||||
z := FromArray(za, true)
|
||||
buildI := func() (*Tensor, error) {
|
||||
sig, err := z.IRFFT(8)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := sig.Pow(2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sq.Sum()
|
||||
}
|
||||
z.ZeroGrad()
|
||||
loss, err := buildI()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
const h = 1e-6
|
||||
for j := range 5 {
|
||||
perturb := func(dr, di float64) float64 {
|
||||
ws := append([]complex128(nil), zs...)
|
||||
ws[j] = complex(zr[j]+dr, zi[j]+di)
|
||||
z.ReplaceWith(mustRecoverComplex(ws, 5))
|
||||
l, err := buildI()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return totalLoss(l)
|
||||
}
|
||||
dr := (perturb(h, 0) - perturb(-h, 0)) / (2 * h)
|
||||
di := (perturb(0, h) - perturb(0, -h)) / (2 * h)
|
||||
want := complex(0.5*dr, 0.5*di)
|
||||
if got := z.Grad().ComplexAt(j); abs2c(got-want) > 1e-4*abs2c(want) {
|
||||
t.Errorf("irfft grad[%d] = %g, want %g", j, got, want)
|
||||
}
|
||||
z.ReplaceWith(mustRecoverComplex(zs, 5))
|
||||
}
|
||||
}
|
||||
|
||||
// TestCompositeGraphBackward walks a graph spanning slice, matmul,
|
||||
// tanh, broadcast, abs2 and an axis reduction, and checks both leaves'
|
||||
// gradients against central differences of the summed loss.
|
||||
func TestCompositeGraphBackward(t *testing.T) {
|
||||
a, _ := FromFloat64s([]float64{1, -2, 3, -4, 5, -6}, true, 2, 3)
|
||||
w, _ := FromFloat64s([]float64{0.5, -0.25, 0.125, 1, -2, 0.75}, true, 2, 3)
|
||||
|
||||
build := func() (*Tensor, error) {
|
||||
s, err := a.Slice(1, 1, 3) // (2,2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
mm, err := s.MatMul(w)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
th, err := mm.Tanh()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bc, err := th.BroadcastTo(2, 2, 3)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := bc.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sq.SumAxis(2)
|
||||
}
|
||||
numCheck := func(base *Tensor, name string, grads []float64) {
|
||||
t.Helper()
|
||||
orig := flatFloats(base.Data())
|
||||
sh := base.Data().Shape()
|
||||
for i := range orig {
|
||||
const h = 1e-6
|
||||
plus := append([]float64(nil), orig...)
|
||||
plus[i] += h
|
||||
base.ReplaceWith(mustRecover(plus, sh...))
|
||||
lp, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
minus := append([]float64(nil), orig...)
|
||||
minus[i] -= h
|
||||
base.ReplaceWith(mustRecover(minus, sh...))
|
||||
lm, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
|
||||
if math.Abs(grads[i]-num) > 1e-5*math.Max(1, math.Abs(num)) {
|
||||
t.Fatalf("%s[%d]: backward %g, central difference %g", name, i, grads[i], num)
|
||||
}
|
||||
base.ReplaceWith(mustRecover(orig, sh...))
|
||||
}
|
||||
}
|
||||
// Backward accumulates, so every read starts from a clean slate.
|
||||
a.ZeroGrad()
|
||||
w.ZeroGrad()
|
||||
loss, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
numCheck(a, "a", flatFloats(a.Grad()))
|
||||
a.ZeroGrad()
|
||||
w.ZeroGrad()
|
||||
loss, err = build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
numCheck(w, "w", flatFloats(w.Grad()))
|
||||
}
|
||||
|
||||
// TestMatMulBatchedBackward checks the stacked product rule of
|
||||
// MatMulBatched against central differences on both operands.
|
||||
func TestMatMulBatchedBackward(t *testing.T) {
|
||||
a, _ := FromFloat64s([]float64{1, 2, 3, 4, 5, 6, 7, 8}, true, 2, 2, 2)
|
||||
b, _ := FromFloat64s([]float64{0.5, -1, 2, 0.25, 1.5, -0.5, -2, 1}, true, 2, 2, 2)
|
||||
|
||||
build := func() (*Tensor, error) {
|
||||
c, err := a.MatMulBatched(b)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := c.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sq.Sum()
|
||||
}
|
||||
a.ZeroGrad()
|
||||
b.ZeroGrad()
|
||||
loss, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, tc := range []struct {
|
||||
tensor *Tensor
|
||||
name string
|
||||
}{{a, "a"}, {b, "b"}} {
|
||||
grads := flatFloats(tc.tensor.Grad())
|
||||
orig := flatFloats(tc.tensor.Data())
|
||||
sh := tc.tensor.Data().Shape()
|
||||
for i := range orig {
|
||||
const h = 1e-6
|
||||
plus := append([]float64(nil), orig...)
|
||||
plus[i] += h
|
||||
tc.tensor.ReplaceWith(mustRecover(plus, sh...))
|
||||
lp, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
minus := append([]float64(nil), orig...)
|
||||
minus[i] -= h
|
||||
tc.tensor.ReplaceWith(mustRecover(minus, sh...))
|
||||
lm, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
|
||||
if math.Abs(grads[i]-num) > 1e-4*math.Max(1, math.Abs(num)) {
|
||||
t.Fatalf("%s[%d]: backward %g, central difference %g", tc.name, i, grads[i], num)
|
||||
}
|
||||
tc.tensor.ReplaceWith(mustRecover(orig, sh...))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestStridedOperandGradientMatchesDenseTwin pins the graph's stride
|
||||
// discipline: a MatMul over a sliced view and the same product over a
|
||||
// dense twin must commit identical gradients. The view's leaf receives
|
||||
// the dense twin's gradient scattered into the sliced span, zeros
|
||||
// elsewhere.
|
||||
func TestStridedOperandGradientMatchesDenseTwin(t *testing.T) {
|
||||
wVals := []float64{0.5, -0.25, 0.125, 1, -2, 0.75}
|
||||
|
||||
// The strided run: leaf (2,4) -> slice -> matmul.
|
||||
full, _ := FromFloat64s([]float64{9, 1, -2, 3, 9, 5, -6, 9}, true, 2, 4)
|
||||
s, err := full.Slice(1, 1, 4) // (2,3) strided view holding 1,-2,3 / 5,-6,9
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
w, _ := FromFloat64s(wVals, true, 3, 2)
|
||||
mm, err := s.MatMul(w)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sq, err := mm.Abs2()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loss, err := sq.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if full.Grad() == nil || w.Grad() == nil {
|
||||
t.Fatal("both leaves must receive a gradient")
|
||||
}
|
||||
gs := flatFloats(full.Grad())
|
||||
wStrided := flatFloats(w.Grad())
|
||||
full.ZeroGrad()
|
||||
w.ZeroGrad()
|
||||
|
||||
// The dense twin.
|
||||
a2, _ := FromFloat64s(flatFloats(s.Data()), true, 2, 3)
|
||||
w2, _ := FromFloat64s(wVals, true, 3, 2)
|
||||
mm2, err := a2.MatMul(w2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sq2, err := mm2.Abs2()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loss2, err := sq2.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss2.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
aDense := flatFloats(a2.Grad())
|
||||
wDense := flatFloats(w2.Grad())
|
||||
|
||||
for i := range wStrided {
|
||||
if wStrided[i] != wDense[i] {
|
||||
t.Fatalf("w gradient %v over the strided operand, want the dense twin's %v", wStrided, wDense)
|
||||
}
|
||||
}
|
||||
if gs[0] != 0 || gs[4] != 0 {
|
||||
t.Fatalf("columns outside the slice carry %g and %g, want 0", gs[0], gs[4])
|
||||
}
|
||||
if gs[1] != aDense[0] || gs[2] != aDense[1] || gs[3] != aDense[2] ||
|
||||
gs[5] != aDense[3] || gs[6] != aDense[4] || gs[7] != aDense[5] {
|
||||
t.Fatalf("leaf gradient %v does not scatter the dense gradient %v into the slice", gs, aDense)
|
||||
}
|
||||
}
|
||||
|
||||
// TestHessianMatchesAnalyticForm differentiates
|
||||
// f(x, y) = x^3 y + exp(x) log(y+2) twice and compares the answer with
|
||||
// the closed form:
|
||||
//
|
||||
// dxx = 6xy + exp(x)log(y+2), dxy = 3x^2 + exp(x)/(y+2),
|
||||
// dyy = -exp(x)/(y+2)^2.
|
||||
func TestHessianMatchesAnalyticForm(t *testing.T) {
|
||||
f := func(v *Tensor) (*Tensor, error) {
|
||||
x, err := v.Slice(0, 0, 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
y, err := v.Slice(0, 1, 2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
x3, err := x.Pow(3)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
t1, err := x3.Mul(y)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
e, err := x.Exp()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
y2, err := y.Add(FromArray(mustRecover([]float64{2}, 1), false))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
lg, err := y2.Log()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
t2, err := e.Mul(lg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return t1.Add(t2)
|
||||
}
|
||||
x, _ := FromFloat64s([]float64{0.4, 1.7}, true, 2)
|
||||
h, err := Hessian(f, x, HessianOptions{Step: 1e-5})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
xx, yy := 0.4, 1.7
|
||||
ex, ly := math.Exp(xx), math.Log(yy+2)
|
||||
want := [4]float64{
|
||||
6*xx*yy + ex*ly, 3*xx*xx + ex/(yy+2),
|
||||
3*xx*xx + ex/(yy+2), -ex / ((yy + 2) * (yy + 2)),
|
||||
}
|
||||
for i := range 4 {
|
||||
if math.Abs(h.FloatAt(i)-want[i]) > 1e-4*math.Max(1, math.Abs(want[i])) {
|
||||
t.Fatalf("Hessian[%d] = %g, want %g", i, h.FloatAt(i), want[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestTripleUseLeafGradient folds a leaf through three nodes and pins
|
||||
// the accumulated gradient: dL/dz = 2(v^2+v)(2v+1) for
|
||||
// L = (v^2+v)^2.
|
||||
func TestTripleUseLeafGradient(t *testing.T) {
|
||||
for _, v := range []float64{0.5, -1.25, 2} {
|
||||
z, _ := FromFloat64s([]float64{v}, true, 1)
|
||||
z2, _ := z.Mul(z)
|
||||
s, err := z2.Add(z)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sq, err := s.Pow(2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loss, err := sq.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := 2 * (v*v + v) * (2*v + 1)
|
||||
if got := z.Grad().FloatAt(0); math.Abs(got-want) > 1e-12*math.Max(1, math.Abs(want)) {
|
||||
t.Fatalf("triple use at z=%g: g = %g, want %g", v, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestComplexAbs2LeafGradient pins the Wirtinger gradient of a summed
|
||||
// |z|^2 loss: the leaf receives z itself.
|
||||
func TestComplexAbs2LeafGradient(t *testing.T) {
|
||||
zr := []float64{0.3, -1.2, 0.7}
|
||||
zi := []float64{-0.4, 0.8, 1.1}
|
||||
zs := make([]complex128, 3)
|
||||
for i := range zs {
|
||||
zs[i] = complex(zr[i], zi[i])
|
||||
}
|
||||
za, _ := core.ComplexFromArray(zs, 3)
|
||||
z := FromArray(za, true)
|
||||
abs2, err := z.Abs2()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loss, err := abs2.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
g := z.Grad()
|
||||
for i := range 3 {
|
||||
want := complex(zr[i], zi[i])
|
||||
if got := g.ComplexAt(i); got != want {
|
||||
t.Fatalf("|z|^2 leaf grad[%d] = %g, want %g", i, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestComplexMatMulBackward checks the Wirtinger adjoint of MatMul
|
||||
// against component-wise central differences of a real loss. The leaf
|
||||
// gradient is dL/dzbar = (dL/dRe + i dL/dIm)/2 under the package
|
||||
// convention.
|
||||
func TestComplexMatMulBackward(t *testing.T) {
|
||||
az := []complex128{complex(1, 0.5), complex(-0.5, 2), complex(0.25, -1), complex(2, 0.25)}
|
||||
bz := []complex128{complex(0.5, 1), complex(-1, 0.25), complex(1.5, -0.5), complex(0.75, 1.25)}
|
||||
aa, _ := core.ComplexFromArray(az, 2, 2)
|
||||
ba, _ := core.ComplexFromArray(bz, 2, 2)
|
||||
a := FromArray(aa, true)
|
||||
b := FromArray(ba, true)
|
||||
build := func() (*Tensor, error) {
|
||||
m, err := a.MatMul(b)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r, err := m.Real()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := r.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sq.Sum()
|
||||
}
|
||||
a.ZeroGrad()
|
||||
b.ZeroGrad()
|
||||
loss, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
const h = 1e-6
|
||||
check := func(base *Tensor, vals []complex128, name string) {
|
||||
t.Helper()
|
||||
for j := range len(vals) {
|
||||
perturb := func(dr, di float64) float64 {
|
||||
ws := append([]complex128(nil), vals...)
|
||||
ws[j] = complex(real(vals[j])+dr, imag(vals[j])+di)
|
||||
base.ReplaceWith(mustRecoverComplex(ws, 2, 2))
|
||||
l, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return totalLoss(l)
|
||||
}
|
||||
dr := (perturb(h, 0) - perturb(-h, 0)) / (2 * h)
|
||||
di := (perturb(0, h) - perturb(0, -h)) / (2 * h)
|
||||
want := complex(0.5*dr, 0.5*di)
|
||||
if got := base.Grad().ComplexAt(j); abs2c(got-want) > 1e-4*abs2c(want) {
|
||||
t.Errorf("complex matmul %s grad[%d] = %g, want %g", name, j, got, want)
|
||||
}
|
||||
base.ReplaceWith(mustRecoverComplex(vals, 2, 2))
|
||||
}
|
||||
}
|
||||
check(a, az, "a")
|
||||
check(b, bz, "b")
|
||||
}
|
||||
|
||||
// TestConcatMixedDtypeBackward checks the join's backward on a real and
|
||||
// a complex side at once: the real side receives 2 Re g through the
|
||||
// narrowing, the complex side its Wirtinger gradient.
|
||||
func TestConcatMixedDtypeBackward(t *testing.T) {
|
||||
r, _ := FromFloat64s([]float64{1, -2, 3}, true, 3)
|
||||
zr := []complex128{complex(0.5, 1), complex(-1.5, 0.25)}
|
||||
za, _ := core.ComplexFromArray(zr, 2)
|
||||
z := FromArray(za, true)
|
||||
build := func() (*Tensor, error) {
|
||||
c, err := r.Concat(z, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
re, err := c.Real()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := re.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sq.Sum()
|
||||
}
|
||||
r.ZeroGrad()
|
||||
z.ZeroGrad()
|
||||
loss, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
grads := flatFloats(r.Grad())
|
||||
orig := flatFloats(r.Data())
|
||||
for i := range orig {
|
||||
const h = 1e-6
|
||||
plus := append([]float64(nil), orig...)
|
||||
plus[i] += h
|
||||
r.ReplaceWith(mustRecover(plus, 3))
|
||||
lp, _ := build()
|
||||
minus := append([]float64(nil), orig...)
|
||||
minus[i] -= h
|
||||
r.ReplaceWith(mustRecover(minus, 3))
|
||||
lm, _ := build()
|
||||
num := (totalLoss(lp) - totalLoss(lm)) / (2 * h)
|
||||
if math.Abs(grads[i]-num) > 1e-5*math.Max(1, math.Abs(num)) {
|
||||
t.Errorf("concat real side grad[%d] = %g, central %g", i, grads[i], num)
|
||||
}
|
||||
r.ReplaceWith(mustRecover(orig, 3))
|
||||
}
|
||||
for j := range 2 {
|
||||
perturb := func(dr, di float64) float64 {
|
||||
ws := append([]complex128(nil), zr...)
|
||||
ws[j] = complex(real(zr[j])+dr, imag(zr[j])+di)
|
||||
z.ReplaceWith(mustRecoverComplex(ws, 2))
|
||||
l, _ := build()
|
||||
return totalLoss(l)
|
||||
}
|
||||
dr := (perturb(1e-6, 0) - perturb(-1e-6, 0)) / (2e-6)
|
||||
di := (perturb(0, 1e-6) - perturb(0, -1e-6)) / (2e-6)
|
||||
want := complex(0.5*dr, 0.5*di)
|
||||
if got := z.Grad().ComplexAt(j); abs2c(got-want) > 1e-4*abs2c(want) {
|
||||
t.Errorf("concat complex side grad[%d] = %g, want %g", j, got, want)
|
||||
}
|
||||
z.ReplaceWith(mustRecoverComplex(zr, 2))
|
||||
}
|
||||
}
|
||||
|
||||
// TestTransposeAxesBackward reverses an axis permutation under a
|
||||
// non-linear loss and checks the gradient against central differences.
|
||||
func TestTransposeAxesBackward(t *testing.T) {
|
||||
x, _ := FromFloat64s([]float64{1, 2, 3, 4, 5, 6, 7, 8}, true, 2, 2, 2)
|
||||
build := func() (*Tensor, error) {
|
||||
p, err := x.TransposeAxes(2, 0, 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := p.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
w, err := sq.Mul(p)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return w.Sum()
|
||||
}
|
||||
probeCheckCentral(t, x, "transposeaxes", build, 1e-5)
|
||||
}
|
||||
|
||||
// TestUnaryKernelChainBackward pushes Sqrt, Log, Sigmoid, Scale, Neg,
|
||||
// Mean and Pow through one graph and checks the committed gradient.
|
||||
func TestUnaryKernelChainBackward(t *testing.T) {
|
||||
x, _ := FromFloat64s([]float64{0.5, 1.5, 2.5, 3.5}, true, 4)
|
||||
build := func() (*Tensor, error) {
|
||||
s, err := x.Sqrt()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
l, err := s.Log()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sg, err := x.Sigmoid()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
m, err := l.Mul(sg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sc, err := m.Scale(2.5)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ng, err := sc.Neg()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
mn, err := ng.Mean()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return mn.Pow(2)
|
||||
}
|
||||
probeCheckCentral(t, x, "unary-chain", build, 1e-4)
|
||||
}
|
||||
|
||||
// TestBackwardAccumulatesUntilZeroGrad pins the accumulation contract
|
||||
// end to end: two Backward calls double the gradient, ZeroGrad resets.
|
||||
func TestBackwardAccumulatesUntilZeroGrad(t *testing.T) {
|
||||
x, _ := FromFloat64s([]float64{2}, true, 1)
|
||||
build := func() (*Tensor, error) {
|
||||
s, err := x.Pow(2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.Sum()
|
||||
}
|
||||
l1, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := l1.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
l2, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := l2.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := x.Grad().FloatAt(0); got != 8 {
|
||||
t.Fatalf("accumulated g = %g, want 8", got)
|
||||
}
|
||||
x.ZeroGrad()
|
||||
l3, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := l3.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := x.Grad().FloatAt(0); got != 4 {
|
||||
t.Fatalf("post-zero g = %g, want 4", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFloat32LeafKeepsGradientDtype pins the dtype contract on a
|
||||
// float32 leaf: the gradient narrows to the leaf's width.
|
||||
func TestFloat32LeafKeepsGradientDtype(t *testing.T) {
|
||||
v := []float32{1.5, -2.5}
|
||||
a, _ := core.FromFloat32Slice(v, 2)
|
||||
x := FromArray(a, true)
|
||||
y, err := x.Pow(2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
l, err := y.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := l.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
g := x.Grad()
|
||||
if g.Dtype() != core.Float32 {
|
||||
t.Fatalf("grad dtype = %s, want float32", g.Dtype())
|
||||
}
|
||||
for i := range 2 {
|
||||
want := 2 * float64(v[i])
|
||||
if got := g.FloatAt(i); math.Abs(got-want) > 1e-6 {
|
||||
t.Fatalf("g[%d] = %g, want %g", i, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestBroadcastToBackwardMatchesCentralDifferences and its SumAxis
|
||||
// neighbour cover the reduction/expansion pair in isolation.
|
||||
func TestBroadcastToBackwardMatchesCentralDifferences(t *testing.T) {
|
||||
x, _ := FromFloat64s([]float64{1, -2, 3, -4, 5, -6}, true, 2, 3)
|
||||
build := func() (*Tensor, error) {
|
||||
bc, err := x.BroadcastTo(2, 2, 3)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := bc.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sq.Sum()
|
||||
}
|
||||
probeCheckCentral(t, x, "broadcast", build, 1e-5)
|
||||
}
|
||||
|
||||
// TestSumAxisBackwardMatchesCentralDifferences reduces the middle axis
|
||||
// with a non-scalar loss, so the backward must scatter into both rows.
|
||||
func TestSumAxisBackwardMatchesCentralDifferences(t *testing.T) {
|
||||
x, _ := FromFloat64s([]float64{1, -2, 3, -4, 5, -6}, true, 2, 3)
|
||||
build := func() (*Tensor, error) {
|
||||
sq, err := x.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sq.SumAxis(1)
|
||||
}
|
||||
probeCheckCentral(t, x, "sumaxis", build, 1e-5)
|
||||
}
|
||||
|
||||
// TestMatMulTanhBackwardMatchesCentralDifferences checks the matmul
|
||||
// product rule under a tanh on top, non-scalar loss included.
|
||||
func TestMatMulTanhBackwardMatchesCentralDifferences(t *testing.T) {
|
||||
a, _ := FromFloat64s([]float64{1, -2, 3, -4}, true, 2, 2)
|
||||
w, _ := FromFloat64s([]float64{0.5, -0.25, 1, -2}, true, 2, 2)
|
||||
build := func() (*Tensor, error) {
|
||||
mm, err := a.MatMul(w)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return mm.Tanh()
|
||||
}
|
||||
a.ZeroGrad()
|
||||
w.ZeroGrad()
|
||||
loss, err := build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
check := func(base *Tensor, name string) {
|
||||
t.Helper()
|
||||
grads := flatFloats(base.Grad())
|
||||
orig := flatFloats(base.Data())
|
||||
sh := base.Data().Shape()
|
||||
for i := range orig {
|
||||
const h = 1e-6
|
||||
plus := append([]float64(nil), orig...)
|
||||
plus[i] += h
|
||||
base.ReplaceWith(mustRecover(plus, sh...))
|
||||
lp, _ := build()
|
||||
minus := append([]float64(nil), orig...)
|
||||
minus[i] -= h
|
||||
base.ReplaceWith(mustRecover(minus, sh...))
|
||||
lm, _ := build()
|
||||
var num float64
|
||||
for j := range lp.Data().Len() {
|
||||
num += lp.Data().FloatAt(j) - lm.Data().FloatAt(j)
|
||||
}
|
||||
num /= 2 * h
|
||||
if math.Abs(grads[i]-num) > 1e-5*math.Max(1, math.Abs(num)) {
|
||||
t.Errorf("%s[%d] backward %g, central %g", name, i, grads[i], num)
|
||||
}
|
||||
base.ReplaceWith(mustRecover(orig, sh...))
|
||||
}
|
||||
}
|
||||
check(a, "a")
|
||||
check(w, "w")
|
||||
}
|
||||
|
||||
// TestNewtonCGSolvesLeastSquares pins the truncated Newton method on a
|
||||
// small full-rank least-squares problem whose solution is A^-1 b.
|
||||
func TestNewtonCGSolvesLeastSquares(t *testing.T) {
|
||||
ar := []float64{2, 0.5, 1, 3}
|
||||
br := []float64{1, -1}
|
||||
A, _ := FromFloat64s(ar, false, 2, 2)
|
||||
B, _ := FromFloat64s(br, false, 2)
|
||||
f := func(x *Tensor) (*Tensor, error) {
|
||||
ax, err := A.MatMul(x)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
d, err := ax.Sub(B)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s, err := d.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.Sum()
|
||||
}
|
||||
x0 := mustRecover([]float64{0, 0}, 2)
|
||||
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
det := 2*3 - 0.5*1
|
||||
wantX := (3*1 - 0.5*(-1)) / det
|
||||
wantY := (2*(-1) - 1*1) / det
|
||||
if math.Abs(x.FloatAt(0)-wantX) > 1e-5 || math.Abs(x.FloatAt(1)-wantY) > 1e-5 {
|
||||
t.Fatalf("NewtonCG point (%g, %g), want (%g, %g)", x.FloatAt(0), x.FloatAt(1), wantX, wantY)
|
||||
}
|
||||
if fv > 1e-10 {
|
||||
t.Fatalf("NewtonCG f = %g, want ~0", fv)
|
||||
}
|
||||
}
|
||||
+165
@@ -0,0 +1,165 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// MatMulBatched multiplies stacked 3-D matrices batch-wise:
|
||||
// (N, M, K) · (N, K, P) gives (N, M, P). The backward runs the classic
|
||||
// product rule inside every batch slot, dA·Bᵀ and Aᵀ·dB, so batched
|
||||
// sequence models can one day drop their per-slice fan-out without
|
||||
// leaving the graph.
|
||||
func (t *Tensor) MatMulBatched(u *Tensor) (*Tensor, error) {
|
||||
if err := t.checkFloat("MatMulBatched"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := u.checkFloat("MatMulBatched"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ta, tu := t.data.Shape(), u.data.Shape()
|
||||
if len(ta) != 3 || len(tu) != 3 {
|
||||
return nil, errf("MatMulBatched: needs rank-3 operands, got %s and %s",
|
||||
prettyShape(ta), prettyShape(tu))
|
||||
}
|
||||
if ta[0] != tu[0] || ta[2] != tu[1] {
|
||||
return nil, errf("MatMulBatched: batch or inner dimension mismatch for %s · %s",
|
||||
prettyShape(ta), prettyShape(tu))
|
||||
}
|
||||
n := ta[0]
|
||||
|
||||
slices := make([]*core.Array, n)
|
||||
// The dtype follows the promotion ladder even for an empty batch,
|
||||
// where no product runs to derive it: an Int output for float
|
||||
// inputs would leak the zero value.
|
||||
dt := t.data.Dtype()
|
||||
if u.data.Dtype() == core.Complex || dt == core.Complex {
|
||||
dt = core.Complex
|
||||
} else if u.data.Dtype() == core.Float && dt != core.Float {
|
||||
dt = core.Float
|
||||
}
|
||||
for i := range n {
|
||||
aSlot, _ := core.Slice(t.data, 0, i, i+1)
|
||||
bSlot, _ := core.Slice(u.data, 0, i, i+1)
|
||||
aMat, _ := core.Reshape(aSlot, ta[1], ta[2])
|
||||
bMat, _ := core.Reshape(bSlot, tu[1], tu[2])
|
||||
prod, err := core.MatMul2D(aMat, bMat)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
slices[i] = prod
|
||||
dt = prod.Dtype()
|
||||
}
|
||||
out := zeros(dt, []int{n, ta[1], tu[2]})
|
||||
slot := ta[1] * tu[2]
|
||||
// Each slot lands in a fresh contiguous product, so the scatter is
|
||||
// a raw slice move per batch row, widened nowhere: out carries the
|
||||
// products' own dtype.
|
||||
for i := range n {
|
||||
switch dt {
|
||||
case core.Float32:
|
||||
copy(out.RawFloat32s()[i*slot:(i+1)*slot], slices[i].RawFloat32s())
|
||||
case core.Float:
|
||||
copy(out.RawFloats()[i*slot:(i+1)*slot], slices[i].RawFloats())
|
||||
default:
|
||||
for j := range slot {
|
||||
out.SetFloatAt(i*slot+j, slices[i].FloatAt(j))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
at, au := t.data, u.data
|
||||
return binaryResult("MatMulBatched", t, u, out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
da := gradSlot{arr: ar.borrowGrad(core.Float, ta), sh: ta}
|
||||
db := gradSlot{arr: ar.borrowGrad(core.Float, tu), sh: tu}
|
||||
m, k, p := ta[1], ta[2], tu[2]
|
||||
slotLen := m * p
|
||||
for i := range n {
|
||||
gMat, err := core.Reshape(
|
||||
mustWindow(g.arr, i*slotLen, slotLen), m, p)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
bMat, err := windowMatrix(au, i, k, p)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
bT := core.Transpose(bMat)
|
||||
daPart, err := core.MatMul2D(gMat, bT)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
copyInto(da.arr, daPart, i*m*k)
|
||||
aMat, err := windowMatrix(at, i, m, k)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
aT := core.Transpose(aMat)
|
||||
dbPart, err := core.MatMul2D(aT, gMat)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
copyInto(db.arr, dbPart, i*k*p)
|
||||
}
|
||||
dst[0], dst[1] = da, db
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// mustWindow flattens row-major window [from, from+len) into an
|
||||
// (rows, cols) matrix view materialisation. The window hands its slice
|
||||
// to FloatsFromArray, which takes ownership, so contiguous gradients
|
||||
// copy nothing at all.
|
||||
func mustWindow(g *core.Array, from, length int) *core.Array {
|
||||
if !g.Strided() && g.Dtype() == core.Float {
|
||||
arr, _ := core.FloatsFromArray(g.RawFloats()[from:from+length], length)
|
||||
return arr
|
||||
}
|
||||
vals := make([]float64, length)
|
||||
for i := range vals {
|
||||
vals[i] = g.FloatAt(from + i)
|
||||
}
|
||||
arr, _ := core.FromFloats(vals, length)
|
||||
return arr
|
||||
}
|
||||
|
||||
// windowMatrix reads batch slot n as an (r, c) float64 matrix, the
|
||||
// backward's native arithmetic domain. A contiguous float64 operand is
|
||||
// aliased rather than copied; anything else is read through the
|
||||
// widening accessor.
|
||||
func windowMatrix(a *core.Array, n, r, c int) (*core.Array, error) {
|
||||
base := n * r * c
|
||||
if !a.Strided() && a.Dtype() == core.Float {
|
||||
arr, err := core.FloatsFromArray(a.RawFloats()[base:base+r*c], r, c)
|
||||
return arr, err
|
||||
}
|
||||
vals := make([]float64, r*c)
|
||||
for i := range vals {
|
||||
vals[i] = a.FloatAt(base + i)
|
||||
}
|
||||
return core.FromFloats(vals, r, c)
|
||||
}
|
||||
|
||||
// copyInto writes src's elements at dst's flat offset. The
|
||||
// destination is a fresh contiguous float64 accumulator; a contiguous
|
||||
// float64 source moves with one copy, a float32 one widens in place.
|
||||
func copyInto(dst, src *core.Array, offset int) {
|
||||
switch {
|
||||
case !src.Strided() && src.Dtype() == core.Float:
|
||||
copy(dst.RawFloats()[offset:], src.RawFloats())
|
||||
case !src.Strided() && src.Dtype() == core.Float32:
|
||||
// Bounded by the source: the destination tail runs on to the
|
||||
// end of the accumulator, which is longer for every batch but
|
||||
// the last.
|
||||
ss, ds := src.RawFloat32s(), dst.RawFloats()[offset:]
|
||||
for i := range src.Len() {
|
||||
ds[i] = float64(ss[i])
|
||||
}
|
||||
default:
|
||||
for i := range src.Len() {
|
||||
dst.SetFloatAt(offset+i, src.FloatAt(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,188 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
func TestMatMulBatchedForward(t *testing.T) {
|
||||
a, _ := core.FromFloats([]float64{
|
||||
1, 0,
|
||||
0, 1,
|
||||
3, 4,
|
||||
5, 6,
|
||||
}, 2, 2, 2)
|
||||
b, _ := core.FromFloats([]float64{
|
||||
1, 1,
|
||||
1, 0,
|
||||
2, 0,
|
||||
0, 2,
|
||||
}, 2, 2, 2)
|
||||
|
||||
out, err := FromArray(a, false).MatMulBatched(FromArray(b, false))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := []float64{1, 1, 1, 0, 6, 8, 10, 12}
|
||||
for i := range want {
|
||||
if g := out.Data().FloatAt(i); g != want[i] {
|
||||
t.Fatalf("slot %d = %v, want %v", i, g, want[i])
|
||||
}
|
||||
}
|
||||
|
||||
// Rank and batch mismatches error loudly.
|
||||
flat, _ := core.Reshape(a, 8)
|
||||
if _, err := FromArray(flat, false).MatMulBatched(FromArray(b, false)); err == nil {
|
||||
t.Fatal("rank-2 operand accepted")
|
||||
}
|
||||
c, _ := core.FromFloats(make([]float64, 4), 1, 2, 2)
|
||||
if _, err := FromArray(a, false).MatMulBatched(FromArray(c, false)); err == nil {
|
||||
t.Fatal("batch-size mismatch accepted")
|
||||
}
|
||||
d, _ := core.FromFloats(make([]float64, 12), 2, 3, 2)
|
||||
if _, err := FromArray(a, false).MatMulBatched(FromArray(d, false)); err == nil {
|
||||
t.Fatal("inner-dimension mismatch accepted")
|
||||
}
|
||||
}
|
||||
|
||||
// TestMatMulBatchedGradients finite-difference checks both operands on
|
||||
// a weighted sum objective so every batch slot earns its own weight.
|
||||
func TestMatMulBatchedGradients(t *testing.T) {
|
||||
aVal := []float64{0.5, -1, 2, 0.25, 1.5, -0.5, 0.75, 1.25, -0.25, 0.5, -1.5, 2}
|
||||
bVal := []float64{1, -0.25, 0.75, 2, -1, 0.125, 0.5, -2}
|
||||
weight := sweepPattern(12) // covers (2, 3, 2) outputs
|
||||
|
||||
aArr, _ := core.FromFloats(aVal, 2, 3, 2)
|
||||
bArr, _ := core.FromFloats(bVal, 2, 2, 2)
|
||||
mArr, _ := core.FromFloats(weight, 2, 3, 2)
|
||||
|
||||
at := FromArray(aArr, true)
|
||||
bt := FromArray(bArr, true)
|
||||
out, err := at.MatMulBatched(bt)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
scaled, err := out.Mul(FromArray(mArr, false))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loss, err := scaled.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
objective := func(av, bv []float64) float64 {
|
||||
x, _ := core.FromFloats(av, 2, 3, 2)
|
||||
y, _ := core.FromFloats(bv, 2, 2, 2)
|
||||
o, rerr := FromArray(x, false).MatMulBatched(FromArray(y, false))
|
||||
if rerr != nil {
|
||||
return math.NaN()
|
||||
}
|
||||
total := 0.0
|
||||
for i := range weight {
|
||||
total += mArr.FloatAt(i) * o.Data().FloatAt(i)
|
||||
}
|
||||
return total
|
||||
}
|
||||
checkSpan(t, at.Grad(), numericGrad(func(v *core.Array) float64 {
|
||||
return objective(flatten(v), bVal)
|
||||
}, aArr))
|
||||
checkSpan(t, bt.Grad(), numericGrad(func(v *core.Array) float64 {
|
||||
return objective(aVal, flatten(v))
|
||||
}, bArr))
|
||||
}
|
||||
|
||||
func flatten(v *core.Array) []float64 {
|
||||
out := make([]float64, v.Len())
|
||||
for i := range v.Len() {
|
||||
out[i] = v.FloatAt(i)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestMatMulBatchedGradientsFloat32 runs the weighted-sum
|
||||
// finite-difference check on float32 operands: the forward rounds to
|
||||
// float32 while the backward widens the accessors and answers float64
|
||||
// gradients, and both agree with the float64 reference.
|
||||
func TestMatMulBatchedGradientsFloat32(t *testing.T) {
|
||||
aVal := []float32{0.5, -1, 2, 0.25, 1.5, -0.5, 0.75, 1.25, -0.25, 0.5, -1.5, 2}
|
||||
bVal := []float32{1, -0.25, 0.75, 2, -1, 0.125, 0.5, -2}
|
||||
weight := sweepPattern(12)
|
||||
|
||||
aArr, err := core.FromFloat32s(aVal, 2, 3, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bArr, err := core.FromFloat32s(bVal, 2, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
mArr, err := core.FromFloats(weight, 2, 3, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
at := FromArray(aArr, true)
|
||||
bt := FromArray(bArr, true)
|
||||
out, err := at.MatMulBatched(bt)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
scaled, err := out.Mul(FromArray(mArr, false))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loss, err := scaled.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if at.Grad().Dtype() != core.Float32 || bt.Grad().Dtype() != core.Float32 {
|
||||
t.Fatalf("float32 leaves carry %s and %s gradients, want float32",
|
||||
at.Grad().Dtype(), bt.Grad().Dtype())
|
||||
}
|
||||
|
||||
// The reference differentiates the same batched product evaluated
|
||||
// in float64 over the identical operand values.
|
||||
objective := func(av, bv []float64) float64 {
|
||||
x, _ := core.FromFloats(av, 2, 3, 2)
|
||||
y, _ := core.FromFloats(bv, 2, 2, 2)
|
||||
o, rerr := FromArray(x, false).MatMulBatched(FromArray(y, false))
|
||||
if rerr != nil {
|
||||
return math.NaN()
|
||||
}
|
||||
total := 0.0
|
||||
for i := range weight {
|
||||
total += mArr.FloatAt(i) * o.Data().FloatAt(i)
|
||||
}
|
||||
return total
|
||||
}
|
||||
aRef, _ := core.FromFloats(widen32(aVal), 2, 3, 2)
|
||||
bRef, _ := core.FromFloats(widen32(bVal), 2, 2, 2)
|
||||
checkSpan(t, at.Grad(), numericGrad(func(v *core.Array) float64 {
|
||||
return objective(flatten(v), widen32(bVal))
|
||||
}, aRef))
|
||||
checkSpan(t, bt.Grad(), numericGrad(func(v *core.Array) float64 {
|
||||
return objective(widen32(aVal), flatten(v))
|
||||
}, bRef))
|
||||
}
|
||||
|
||||
// widen32 widens a float32 slice exactly, the view the backward's own
|
||||
// accessors read.
|
||||
func widen32(v []float32) []float64 {
|
||||
out := make([]float64, len(v))
|
||||
for i, x := range v {
|
||||
out[i] = float64(x)
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// The benchmarks below pin the costs the graph machinery adds around
|
||||
// the kernels: one deep chain of small tensors (per-node tape cost), one
|
||||
// wide element-wise graph (per-edge accumulation cost), the mid-size
|
||||
// sweeps whose per-element work decides their parallel policy, and the
|
||||
// L2-norm axis backward. Inputs are fixed literals, so a run is
|
||||
// deterministic.
|
||||
|
||||
// benchLit builds a leaf of n elements from fixed literals, keeping the
|
||||
// arithmetic well inside the domain of every op used here.
|
||||
func benchLit(b *testing.B, seed, n int) *Tensor {
|
||||
b.Helper()
|
||||
v := make([]float64, n)
|
||||
for i := range v {
|
||||
v[i] = 0.25 + float64((i*7+seed)%13)*0.125
|
||||
}
|
||||
a, err := core.FromFloats(v, n)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
return FromArray(a, true)
|
||||
}
|
||||
|
||||
// mustReshape reshapes an array for a benchmark fixture.
|
||||
func mustReshape(b *testing.B, a *core.Array, shape ...int) *core.Array {
|
||||
b.Helper()
|
||||
out, err := core.Reshape(a, shape...)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// deepChain builds a chain of n element-wise nodes over x and w and
|
||||
// reduces it to a scalar, the shape a training loop's tape has.
|
||||
func deepChain(x, w *Tensor, n int) (*Tensor, error) {
|
||||
h := x
|
||||
for i := range n {
|
||||
var err error
|
||||
switch i % 4 {
|
||||
case 0:
|
||||
h, err = h.Add(w)
|
||||
case 1:
|
||||
h, err = h.Mul(w)
|
||||
case 2:
|
||||
h, err = h.Tanh()
|
||||
default:
|
||||
h, err = h.Scale(0.25)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return h.Sum()
|
||||
}
|
||||
|
||||
// wideFan multiplies x by n independent leaves and sums the products,
|
||||
// so the backward folds n contributions into x's gradient.
|
||||
func wideFan(x *Tensor, leaves []*Tensor) (*Tensor, error) {
|
||||
acc, err := x.Mul(leaves[0])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, l := range leaves[1:] {
|
||||
p, err := x.Mul(l)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if acc, err = acc.Add(p); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return acc.Sum()
|
||||
}
|
||||
|
||||
// BenchmarkTapeDeepForward measures the forward pass alone: one node
|
||||
// per element-wise op over an 8-element tensor.
|
||||
func BenchmarkTapeDeepForward(b *testing.B) {
|
||||
x, w := benchLit(b, 1, 8), benchLit(b, 2, 8)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
s, err := deepChain(x, w, 128)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if s.Data().Len() != 1 {
|
||||
b.Fatal("unexpected shape")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkTapeDeepBackward measures the same chain with the reverse
|
||||
// sweep, where every node reads its gradient and folds into the two
|
||||
// shared leaves.
|
||||
func BenchmarkTapeDeepBackward(b *testing.B) {
|
||||
x, w := benchLit(b, 1, 8), benchLit(b, 2, 8)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
s, err := deepChain(x, w, 128)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
x.ZeroGrad()
|
||||
w.ZeroGrad()
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkTapeWideForward builds a 64-way fan over x, one node per
|
||||
// leaf, and reduces it.
|
||||
func BenchmarkTapeWideForward(b *testing.B) {
|
||||
x := benchLit(b, 3, 16)
|
||||
leaves := make([]*Tensor, 64)
|
||||
for i := range leaves {
|
||||
leaves[i] = benchLit(b, 10+i, 16)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
s, err := wideFan(x, leaves)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if s.Data().Len() != 1 {
|
||||
b.Fatal("unexpected shape")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkTapeWideBackward runs the same 64-way fan with the reverse
|
||||
// sweep: 64 edges fold into x's gradient through the reduction tree.
|
||||
func BenchmarkTapeWideBackward(b *testing.B) {
|
||||
x := benchLit(b, 3, 16)
|
||||
leaves := make([]*Tensor, 64)
|
||||
for i := range leaves {
|
||||
leaves[i] = benchLit(b, 10+i, 16)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
s, err := wideFan(x, leaves)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
x.ZeroGrad()
|
||||
for _, l := range leaves {
|
||||
l.ZeroGrad()
|
||||
}
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// midElems is the sweep size the transcendental benchmarks use: below
|
||||
// the element-wise floor of 1024 per worker on a 32-worker machine, so
|
||||
// a sweep of this size runs on the calling goroutine under that policy
|
||||
// and splits under a floor scaled to its per-element cost.
|
||||
const midElems = 20000
|
||||
|
||||
// BenchmarkPowForwardMid measures the integer-exponent power over a
|
||||
// mid-size sweep, one math.Pow per element.
|
||||
func BenchmarkPowForwardMid(b *testing.B) {
|
||||
x := benchLit(b, 5, midElems)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
y, err := x.Pow(3)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if y.Data().Len() != midElems {
|
||||
b.Fatal("unexpected shape")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkPowBackwardMid measures the power backward over the same
|
||||
// size: one Pow and one multiply per element, plus the reduction's
|
||||
// fill.
|
||||
func BenchmarkPowBackwardMid(b *testing.B) {
|
||||
x := benchLit(b, 5, midElems)
|
||||
y, err := x.Pow(3)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
s, err := y.Sum()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
x.ZeroGrad()
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkSqrtBackwardMid measures the square-root backward over a
|
||||
// mid-size sweep, one divide per element.
|
||||
func BenchmarkSqrtBackwardMid(b *testing.B) {
|
||||
x := benchLit(b, 6, midElems)
|
||||
y, err := x.Sqrt()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
s, err := y.Sum()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
x.ZeroGrad()
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkMeanAxisBackwardMid measures the axis-mean backward, whose
|
||||
// first stage divides every element of the incoming gradient.
|
||||
func BenchmarkMeanAxisBackwardMid(b *testing.B) {
|
||||
x := benchLit(b, 7, midElems)
|
||||
xt := FromArray(mustReshape(b, x.Data(), 200, 100), true)
|
||||
y, err := xt.MeanAxis(0)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
s, err := y.Sum()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
xt.ZeroGrad()
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkSumBackwardWide measures a reduction over a large operand:
|
||||
// the backward fills the operand's shape with one value, and the pass
|
||||
// commits a hundred-thousand-element gradient into the leaf.
|
||||
func BenchmarkSumBackwardWide(b *testing.B) {
|
||||
x := benchLit(b, 8, 100000)
|
||||
s, err := x.Sum()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
x.ZeroGrad()
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkL2NormAxisBackward measures the norm backward over a
|
||||
// (2000×100) tensor reduced along the leading axis: 100 lines of 2000
|
||||
// elements each, long enough for the line sweep to dominate the
|
||||
// allocation of the output.
|
||||
func BenchmarkL2NormAxisBackward(b *testing.B) {
|
||||
x := benchLit(b, 9, 200000)
|
||||
xt := FromArray(mustReshape(b, x.Data(), 2000, 100), true)
|
||||
y, err := xt.L2NormAxis(0)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
s, err := y.Sum()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
xt.ZeroGrad()
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
core "sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Backward benchmarks guard the tape overhead around the kernels: the
|
||||
// two-layer graph is the smallest shape where per-node costs and the
|
||||
// matmul backward both show.
|
||||
|
||||
func benchVals(b *testing.B, seed, n int) []float64 {
|
||||
b.Helper()
|
||||
v := make([]float64, n)
|
||||
for i := range v {
|
||||
v[i] = float64(i%13)*float64(seed%3)*0.25 + float64(i%5) - 2
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func benchTensor(b *testing.B, seed int, shape ...int) *Tensor {
|
||||
b.Helper()
|
||||
n := 1
|
||||
for _, d := range shape {
|
||||
n *= d
|
||||
}
|
||||
a, err := core.FromFloats(benchVals(b, seed, n), shape...)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
return FromArray(a, true)
|
||||
}
|
||||
|
||||
// BenchmarkBackwardTwoLayer runs forward and backward over
|
||||
// (32×64)·(64×32), then tanh, then ·(32×10), then sum.
|
||||
func BenchmarkBackwardTwoLayer(b *testing.B) {
|
||||
x := benchTensor(b, 1, 32, 64)
|
||||
w1 := benchTensor(b, 2, 64, 32)
|
||||
w2 := benchTensor(b, 3, 32, 10)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
h, err := x.MatMul(w1)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
t, err := h.Tanh()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
y, err := t.MatMul(w2)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
s, err := y.Sum()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkForwardOnly isolates the graph construction from the
|
||||
// backward sweep.
|
||||
func BenchmarkForwardOnly(b *testing.B) {
|
||||
x := benchTensor(b, 1, 32, 64)
|
||||
w1 := benchTensor(b, 2, 64, 32)
|
||||
w2 := benchTensor(b, 3, 32, 10)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
h, err := x.MatMul(w1)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
t, err := h.Tanh()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if _, err := t.MatMul(w2); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkGradMatMulBackward isolates one MatMul node's backward
|
||||
// sweep on (128×128) operands.
|
||||
func BenchmarkGradMatMulBackward(b *testing.B) {
|
||||
a := benchTensor(b, 4, 128, 128)
|
||||
c := benchTensor(b, 5, 128, 128)
|
||||
y, err := a.MatMul(c)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
s, err := y.Sum()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
b.ResetTimer()
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
a.ZeroGrad()
|
||||
c.ZeroGrad()
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// Regression pins for the broadcast and shape backwards: a broadcast
|
||||
// gradient must collapse to its source's shape, and every widened op
|
||||
// must hand each operand its own correctly shaped gradient buffer.
|
||||
|
||||
// TestBroadcastRank1Gradient pins the rank-1 broadcast backward: the
|
||||
// gradient of a size-1 source broadcast to length m must collapse back
|
||||
// to a single sum, not arrive with the broadcast shape.
|
||||
func TestBroadcastRank1Gradient(t *testing.T) {
|
||||
x, err := core.FromFloats([]float64{0}, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
xt := FromArray(x, true)
|
||||
y, err := xt.BroadcastTo(3)
|
||||
if err != nil {
|
||||
t.Fatalf("BroadcastTo: %v", err)
|
||||
}
|
||||
w, err := core.FromFloats([]float64{1, 2, 3}, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
wt := FromArray(w, false)
|
||||
prod, err := y.Mul(wt)
|
||||
if err != nil {
|
||||
t.Fatalf("Mul: %v", err)
|
||||
}
|
||||
loss, err := prod.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
g := xt.Grad()
|
||||
if g.NDim() != 1 || g.Shape()[0] != 1 {
|
||||
t.Fatalf("gradient shape %v, want [1]", g.Shape())
|
||||
}
|
||||
if math.Abs(g.FloatAt(0)-6) > 1e-12 {
|
||||
t.Fatalf("gradient = %g, want 6", g.FloatAt(0))
|
||||
}
|
||||
}
|
||||
|
||||
// TestBroadcastRank1ToMatrix pins the (1,) to (m, n) broadcast backward
|
||||
// against central differences.
|
||||
func TestBroadcastRank1ToMatrix(t *testing.T) {
|
||||
x, err := core.FromFloats([]float64{2}, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
xt := FromArray(x, true)
|
||||
y, err := xt.BroadcastTo(2, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("BroadcastTo: %v", err)
|
||||
}
|
||||
w, err := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
wt := FromArray(w, false)
|
||||
prod, err := y.Mul(wt)
|
||||
if err != nil {
|
||||
t.Fatalf("Mul: %v", err)
|
||||
}
|
||||
loss, err := prod.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
g := xt.Grad()
|
||||
if g.NDim() != 1 || g.Shape()[0] != 1 {
|
||||
t.Fatalf("gradient shape %v, want [1]", g.Shape())
|
||||
}
|
||||
if math.Abs(g.FloatAt(0)-21) > 1e-12 {
|
||||
t.Fatalf("gradient = %g, want 21", g.FloatAt(0))
|
||||
}
|
||||
}
|
||||
|
||||
// TestL2NormAxisEmptyDim pins the backward against the integer division
|
||||
// by zero an empty reduced dimension used to hit.
|
||||
func TestL2NormAxisEmptyDim(t *testing.T) {
|
||||
x, err := core.FromFloats([]float64{}, 3, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
xt := FromArray(x, true)
|
||||
n, err := xt.L2NormAxis(1)
|
||||
if err != nil {
|
||||
t.Fatalf("L2NormAxis: %v", err)
|
||||
}
|
||||
loss, err := n.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
g := xt.Grad()
|
||||
if g.Len() != 0 {
|
||||
t.Fatalf("gradient length %d, want 0", g.Len())
|
||||
}
|
||||
}
|
||||
|
||||
// TestShapeOpsRejectNonFloat pins the dtype contract on the shape ops
|
||||
// that used to record graph nodes without validating the dtype.
|
||||
func TestShapeOpsRejectNonFloat(t *testing.T) {
|
||||
i, err := core.FromInts([]int64{1, 2, 3, 4}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromInts: %v", err)
|
||||
}
|
||||
it := FromArray(i, true)
|
||||
if _, err := it.Transpose(); err == nil {
|
||||
t.Error("Transpose accepted an int tensor")
|
||||
}
|
||||
if _, err := it.Squeeze(0); err == nil {
|
||||
t.Error("Squeeze accepted an int tensor")
|
||||
}
|
||||
if _, err := it.Unsqueeze(0); err == nil {
|
||||
t.Error("Unsqueeze accepted an int tensor")
|
||||
}
|
||||
if _, err := it.Reshape(4); err == nil {
|
||||
t.Error("Reshape accepted an int tensor")
|
||||
}
|
||||
if _, err := it.TransposeAxes(1, 0); err == nil {
|
||||
t.Error("TransposeAxes accepted an int tensor")
|
||||
}
|
||||
if _, err := it.Floor(); err == nil {
|
||||
t.Error("Floor accepted an int tensor")
|
||||
}
|
||||
if _, err := it.BroadcastTo(2, 3); err == nil {
|
||||
t.Error("BroadcastTo accepted an int tensor")
|
||||
}
|
||||
}
|
||||
+503
@@ -0,0 +1,503 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Complex differentiation. The graph accepts complex128
|
||||
// tensors alongside float64/float32, with the Wirtinger convention the
|
||||
// optimiser ecosystem settled on: Backward seeds a REAL scalar loss
|
||||
// (a complex output is rejected with an error telling the caller to
|
||||
// reduce first), and the gradient a complex leaf accumulates is
|
||||
// ∂L/∂z̄, the direction gradient descent steps along. Under that
|
||||
// convention the adjoint of a holomorphic op y = f(z) is
|
||||
// dz += g·conj(f′(z)), so every conjugation below sits exactly where
|
||||
// the calculus puts it.
|
||||
//
|
||||
// A real tensor inside a complex graph narrows the incoming complex
|
||||
// gradient by 2·Re: for a real variable x, dL/dx = 2·Re(∂L/∂x̄), and
|
||||
// the factor also cancels the ½ the Real backward contributes, so
|
||||
// mixed graphs compose exactly.
|
||||
|
||||
// checkDiff is checkFloat plus complex: the ops that can differentiate
|
||||
// complex inputs validate with it.
|
||||
func (t *Tensor) checkDiff(name string) error {
|
||||
switch t.data.Dtype() {
|
||||
case core.Float, core.Float32, core.Complex:
|
||||
return nil
|
||||
}
|
||||
return errf("autograd: %s needs a float, float32 or complex tensor, got %s", name, t.data.Dtype())
|
||||
}
|
||||
|
||||
// isComplexArr reports whether a holds complex128 data.
|
||||
func isComplexArr(a *core.Array) bool { return a.Dtype() == core.Complex }
|
||||
|
||||
// eitherComplex reports whether either operand is complex.
|
||||
func eitherComplex(a, b *core.Array) bool { return isComplexArr(a) || isComplexArr(b) }
|
||||
|
||||
// conjArray returns the element-wise conjugate. Real arrays come back
|
||||
// unchanged (their conjugate is themselves), so mixed-dtype adjoints
|
||||
// can call it unconditionally.
|
||||
func conjArray(a *core.Array) *core.Array {
|
||||
if !isComplexArr(a) {
|
||||
return a
|
||||
}
|
||||
out := zeros(core.Complex, a.Shape())
|
||||
cs := out.RawComplexes()
|
||||
if a.Strided() {
|
||||
for i := range cs {
|
||||
cs[i] = conj(a.ComplexAt(i))
|
||||
}
|
||||
return out
|
||||
}
|
||||
as := a.RawComplexes()
|
||||
for i := range cs {
|
||||
z := as[i]
|
||||
cs[i] = complex(real(z), -imag(z))
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func conj(z complex128) complex128 { return complex(real(z), -imag(z)) }
|
||||
|
||||
// copyElem copies one element between gradient arrays of the same
|
||||
// dtype; the callers narrow the incoming gradient to the operand's
|
||||
// dtype with narrowGradient before the copy, so a mixed real/complex
|
||||
// pair never reaches here and a real destination never reads a complex
|
||||
// payload. A complex destination reads through ComplexAt, which serves
|
||||
// a strided source too.
|
||||
func copyElem(dst *core.Array, di int, src *core.Array, si int) {
|
||||
if dst.Dtype() == core.Complex {
|
||||
dst.RawComplexes()[di] = src.ComplexAt(si)
|
||||
return
|
||||
}
|
||||
dst.SetFloatAt(di, src.FloatAt(si))
|
||||
}
|
||||
|
||||
// narrowGradient converts a gradient to the dtype of the tensor it
|
||||
// accumulates into. Complex to real takes 2·Re (the real-tensor rule
|
||||
// above); everything else routes through Astype.
|
||||
func narrowGradient(g gradSlot, dt core.Dtype) (gradSlot, error) {
|
||||
if g.arr.Dtype() == dt {
|
||||
return g, nil
|
||||
}
|
||||
sh := g.sh
|
||||
if sh == nil {
|
||||
sh = g.arr.Shape()
|
||||
}
|
||||
if g.arr.Dtype() == core.Complex && dt != core.Complex {
|
||||
out := zeros(dt, sh)
|
||||
gs := g.arr.RawComplexes()
|
||||
if dt == core.Float32 && !g.arr.Strided() {
|
||||
os := out.RawFloat32s()
|
||||
for i := range os {
|
||||
os[i] = float32(2 * real(gs[i]))
|
||||
}
|
||||
return gradSlot{arr: out, sh: sh}, nil
|
||||
}
|
||||
if dt == core.Float && !g.arr.Strided() {
|
||||
os := out.RawFloats()
|
||||
for i := range os {
|
||||
os[i] = 2 * real(gs[i])
|
||||
}
|
||||
return gradSlot{arr: out, sh: sh}, nil
|
||||
}
|
||||
// out is freshly allocated and dense, so a real destination
|
||||
// takes its payload directly; the complex source keeps the
|
||||
// accessor read that rebases a strided index.
|
||||
switch dt {
|
||||
case core.Float32:
|
||||
os := out.RawFloat32s()
|
||||
for i := range os {
|
||||
os[i] = float32(2 * real(g.arr.ComplexAt(i)))
|
||||
}
|
||||
return gradSlot{arr: out, sh: sh}, nil
|
||||
case core.Float:
|
||||
os := out.RawFloats()
|
||||
for i := range os {
|
||||
os[i] = 2 * real(g.arr.ComplexAt(i))
|
||||
}
|
||||
return gradSlot{arr: out, sh: sh}, nil
|
||||
}
|
||||
for i := range g.arr.Len() {
|
||||
out.SetFloatAt(i, 2*real(g.arr.ComplexAt(i)))
|
||||
}
|
||||
return gradSlot{arr: out, sh: sh}, nil
|
||||
}
|
||||
c, err := core.Astype(g.arr, dt)
|
||||
if err != nil {
|
||||
return gradSlot{}, err
|
||||
}
|
||||
return gradSlot{arr: c, sh: sh}, nil
|
||||
}
|
||||
|
||||
// scalarComplex builds a 1-element complex array holding z.
|
||||
func scalarComplex(z complex128) *core.Array {
|
||||
out := zeros(core.Complex, []int{1})
|
||||
out.RawComplexes()[0] = z
|
||||
return out
|
||||
}
|
||||
|
||||
// fillComplex returns a complex array shaped like a with every element
|
||||
// set to z.
|
||||
func fillComplex(a *core.Array, z complex128) *core.Array {
|
||||
out := zeros(core.Complex, a.Shape())
|
||||
cs := out.RawComplexes()
|
||||
for i := range cs {
|
||||
cs[i] = z
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// Conj returns the element-wise complex conjugate. The conjugate is
|
||||
// anti-holomorphic: its ∂/∂z̄ adjoint conjugates the incoming
|
||||
// gradient (dz = conj(g)), which is what makes ⟨ψ|H|ψ⟩ come out as
|
||||
// Hψ rather than only its real part.
|
||||
func (t *Tensor) Conj() (*Tensor, error) {
|
||||
if err := t.checkDiff("Conj"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
a := t.data
|
||||
out := conjArray(t.data)
|
||||
if !isComplexArr(t.data) {
|
||||
// conj of a real tensor is a copy, so the graph needs its own
|
||||
// node data, not the operand alias.
|
||||
out = cloneReal(t.data)
|
||||
}
|
||||
return t.unaryResult("Conj", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
sh := gradShape(g, a)
|
||||
if !isComplexArr(g.arr) {
|
||||
c, err := copyGradSlot(ar, g, sh)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dst[0] = c
|
||||
return nil
|
||||
}
|
||||
n := g.arr.Len()
|
||||
da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh}
|
||||
cs := da.arr.RawComplexes()[:n]
|
||||
gs := g.arr.RawComplexes()[:n]
|
||||
for i := range cs {
|
||||
z := gs[i]
|
||||
cs[i] = complex(real(z), -imag(z))
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// cloneReal copies a real array (the graph never aliases operands).
|
||||
func cloneReal(a *core.Array) *core.Array {
|
||||
out := zeros(a.Dtype(), a.Shape())
|
||||
switch {
|
||||
case a.Strided():
|
||||
for i := range a.Len() {
|
||||
out.SetFloatAt(i, a.FloatAt(i))
|
||||
}
|
||||
case a.Dtype() == core.Float32:
|
||||
copy(out.RawFloat32s(), a.RawFloat32s())
|
||||
case a.Dtype() == core.Float:
|
||||
copy(out.RawFloats(), a.RawFloats())
|
||||
default:
|
||||
copy(out.RawInts(), a.RawInts())
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// Real returns the real part of each element as a float tensor. The
|
||||
// complex backward halves the gradient (∂Re z/∂z̄ = ½), which the
|
||||
// 2·Re narrowing at any real destination cancels exactly.
|
||||
func (t *Tensor) Real() (*Tensor, error) {
|
||||
if err := t.checkDiff("Real"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !isComplexArr(t.data) {
|
||||
// Real of a real tensor is a copy with its own storage.
|
||||
out := cloneReal(t.data)
|
||||
a := t.data
|
||||
return t.unaryResult("Real", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
c, err := copyGradSlot(ar, g, gradShape(g, a))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dst[0] = c
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
out := zeros(core.Float, t.data.Shape())
|
||||
if t.data.Strided() {
|
||||
for i := range t.data.Len() {
|
||||
out.SetFloatAt(i, real(t.data.ComplexAt(i)))
|
||||
}
|
||||
} else {
|
||||
cs := t.data.RawComplexes()
|
||||
os := out.RawFloats()
|
||||
for i := range os {
|
||||
os[i] = real(cs[i])
|
||||
}
|
||||
}
|
||||
// The shape is captured now: nothing may be read off the input at
|
||||
// backward time, or a ReplaceWith in between would change it.
|
||||
shape := t.data.Shape()
|
||||
return t.unaryResult("Real", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
da := gradSlot{arr: ar.borrowGrad(core.Complex, shape), sh: shape}
|
||||
cs := da.arr.RawComplexes()[:da.arr.Len()]
|
||||
if g.arr.Strided() || g.arr.Dtype() != core.Float {
|
||||
for i := range cs {
|
||||
cs[i] = complex(g.arr.FloatAt(i)/2, 0)
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}
|
||||
gs := g.arr.RawFloats()
|
||||
for i := range cs {
|
||||
cs[i] = complex(gs[i]/2, 0)
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// Imag returns the imaginary part of each element as a float tensor;
|
||||
// the complex backward scales by i/2 (∂Im z/∂z̄ = i/2).
|
||||
func (t *Tensor) Imag() (*Tensor, error) {
|
||||
if !isComplexArr(t.data) {
|
||||
return nil, errf("autograd: Imag needs a complex tensor, got %s", t.data.Dtype())
|
||||
}
|
||||
out := zeros(core.Float, t.data.Shape())
|
||||
if t.data.Strided() {
|
||||
for i := range t.data.Len() {
|
||||
out.SetFloatAt(i, imag(t.data.ComplexAt(i)))
|
||||
}
|
||||
} else {
|
||||
cs := t.data.RawComplexes()
|
||||
os := out.RawFloats()
|
||||
for i := range os {
|
||||
os[i] = imag(cs[i])
|
||||
}
|
||||
}
|
||||
shape := t.data.Shape()
|
||||
return t.unaryResult("Imag", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
da := gradSlot{arr: ar.borrowGrad(core.Complex, shape), sh: shape}
|
||||
cs := da.arr.RawComplexes()[:da.arr.Len()]
|
||||
if g.arr.Strided() || g.arr.Dtype() != core.Float {
|
||||
for i := range cs {
|
||||
cs[i] = complex(0, g.arr.FloatAt(i)/2)
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}
|
||||
gs := g.arr.RawFloats()
|
||||
for i := range cs {
|
||||
cs[i] = complex(0, gs[i]/2)
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// Abs2 returns |z|² of each element, a real tensor. The complex
|
||||
// backward is dz = g·z (∂|z|²/∂z̄ = z); the real input path is the
|
||||
// square with its 2x backward, keeping the operand's width (a float32
|
||||
// input squares in float64 and stays float32, exactly as Pow
|
||||
// does).
|
||||
func (t *Tensor) Abs2() (*Tensor, error) {
|
||||
if err := t.checkDiff("Abs2"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !isComplexArr(t.data) {
|
||||
return t.squareGraph()
|
||||
}
|
||||
a := t.data
|
||||
out := zeros(core.Float, a.Shape())
|
||||
if a.Strided() {
|
||||
for i := range a.Len() {
|
||||
z := a.ComplexAt(i)
|
||||
out.SetFloatAt(i, real(z)*real(z)+imag(z)*imag(z))
|
||||
}
|
||||
} else {
|
||||
// Bound the walk by the destination's length: a rebased view's
|
||||
// payload may run longer than its element count.
|
||||
as := a.RawComplexes()
|
||||
os := out.RawFloats()
|
||||
for i := range os {
|
||||
z := as[i]
|
||||
os[i] = real(z)*real(z) + imag(z)*imag(z)
|
||||
}
|
||||
}
|
||||
return t.unaryResult("Abs2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
sh := gradShape(g, a)
|
||||
da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh}
|
||||
cs := da.arr.RawComplexes()[:da.arr.Len()]
|
||||
if a.Strided() || g.arr.Strided() || g.arr.Dtype() != core.Float {
|
||||
for i := range cs {
|
||||
cs[i] = complex(g.arr.FloatAt(i), 0) * a.ComplexAt(i)
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}
|
||||
as, gs := a.RawComplexes(), g.arr.RawFloats()
|
||||
for i := range cs {
|
||||
cs[i] = complex(gs[i], 0) * as[i]
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// squareGraph is the real-input branch of Abs2: y = x², dx = 2x·g.arr.
|
||||
// The output keeps the operand's width, squared in float64 and rounded
|
||||
// once, exactly as Pow does, so Abs2 and Pow(2) agree on dtype
|
||||
// and value for a float32 operand.
|
||||
func (t *Tensor) squareGraph() (*Tensor, error) {
|
||||
a := t.data
|
||||
out := zeros(a.Dtype(), a.Shape())
|
||||
switch {
|
||||
case a.Dtype() == core.Float32 && !a.Strided():
|
||||
as, os := a.RawFloat32s(), out.RawFloat32s()
|
||||
for i := range os {
|
||||
v := float64(as[i])
|
||||
os[i] = float32(v * v)
|
||||
}
|
||||
case a.Strided():
|
||||
for i := range out.Len() {
|
||||
v := a.FloatAt(i)
|
||||
out.SetFloatAt(i, v*v)
|
||||
}
|
||||
default:
|
||||
as, os := a.RawFloats(), out.RawFloats()
|
||||
for i := range os {
|
||||
v := as[i]
|
||||
os[i] = v * v
|
||||
}
|
||||
}
|
||||
return t.unaryResult("Abs2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
sh := gradShape(g, a)
|
||||
// dx = 2·x·g with the staged chain's rounding: the product
|
||||
// forms first and the doubling multiplies it, per element.
|
||||
if !a.Strided() && !g.arr.Strided() && a.Dtype() == g.arr.Dtype() && a.Len() == g.arr.Len() {
|
||||
n := a.Len()
|
||||
switch a.Dtype() {
|
||||
case core.Float:
|
||||
da := gradSlot{arr: ar.borrowGrad(core.Float, sh), sh: sh}
|
||||
as, gs, ds := a.RawFloats()[:n], g.arr.RawFloats()[:n], da.arr.RawFloats()[:n]
|
||||
for i := range ds {
|
||||
ds[i] = (as[i] * gs[i]) * 2
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
case core.Float32:
|
||||
da := gradSlot{arr: ar.borrowGrad(core.Float32, sh), sh: sh}
|
||||
as, gs, ds := a.RawFloat32s()[:n], g.arr.RawFloat32s()[:n], da.arr.RawFloat32s()[:n]
|
||||
for i := range ds {
|
||||
p := float32(float64(as[i]) * float64(gs[i]))
|
||||
ds[i] = float32(float64(p) * 2)
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}
|
||||
}
|
||||
da, err := core.Mul(a, g.arr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dst[0] = gradSlot{arr: core.MulI(da, 2), sh: sh}
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// Abs returns the absolute value of each element: complex input yields
|
||||
// float magnitudes with dz = g·z/(2|z|) (zero at the origin, the
|
||||
// subgradient). The real branch lives beside the other real kernels in
|
||||
// tensor.go and dispatches here for complex input.
|
||||
func (t *Tensor) absComplex() (*Tensor, error) {
|
||||
a := t.data
|
||||
out := core.Abs(a)
|
||||
return t.unaryResult("Abs", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
sh := gradShape(g, a)
|
||||
da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh}
|
||||
cs := da.arr.RawComplexes()[:da.arr.Len()]
|
||||
if a.Strided() || g.arr.Strided() || out.Strided() ||
|
||||
g.arr.Dtype() != core.Float || out.Dtype() != core.Float {
|
||||
for i := range cs {
|
||||
z := a.ComplexAt(i)
|
||||
m := out.FloatAt(i)
|
||||
if m == 0 {
|
||||
continue
|
||||
}
|
||||
cs[i] = complex(g.arr.FloatAt(i)/(2*m), 0) * z
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}
|
||||
as, gs, os := a.RawComplexes(), g.arr.RawFloats(), out.RawFloats()
|
||||
for i := range cs {
|
||||
m := os[i]
|
||||
if m == 0 {
|
||||
continue
|
||||
}
|
||||
cs[i] = complex(gs[i]/(2*m), 0) * as[i]
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// powComplexGrad builds the Wirtinger backward of y = zⁿ:
|
||||
// dz = g·n·conj(z)ⁿ⁻¹, assembled by repeated conjugate multiplication
|
||||
// (the exponent is a small integer; a loop beats a general power).
|
||||
// sh is the shape the incoming gradient carries, or the operand's own
|
||||
// on the legacy sweep path (gradShape).
|
||||
func powComplexGrad(ar *gradArena, g gradSlot, a *core.Array, n int64, sh []int) *core.Array {
|
||||
da := ar.borrowGrad(core.Complex, sh)
|
||||
cs := da.RawComplexes()[:da.Len()]
|
||||
if a.Strided() || g.arr.Strided() {
|
||||
for i := range cs {
|
||||
term := complex(1, 0)
|
||||
for range n - 1 {
|
||||
term *= conj(a.ComplexAt(i))
|
||||
}
|
||||
cs[i] = complex(float64(n), 0) * g.arr.ComplexAt(i) * term
|
||||
}
|
||||
return da
|
||||
}
|
||||
as, gs := a.RawComplexes(), g.arr.RawComplexes()
|
||||
for i := range cs {
|
||||
term := complex(1, 0)
|
||||
for range n - 1 {
|
||||
term *= conj(as[i])
|
||||
}
|
||||
cs[i] = complex(float64(n), 0) * gs[i] * term
|
||||
}
|
||||
return da
|
||||
}
|
||||
|
||||
// powComplexForward raises each complex element to a non-negative
|
||||
// integer power by repeated multiplication.
|
||||
func powComplexForward(a *core.Array, n int64) *core.Array {
|
||||
out := zeros(core.Complex, a.Shape())
|
||||
cs := out.RawComplexes()
|
||||
if a.Strided() {
|
||||
for i := range cs {
|
||||
p := complex(1, 0)
|
||||
for range n {
|
||||
p *= a.ComplexAt(i)
|
||||
}
|
||||
cs[i] = p
|
||||
}
|
||||
return out
|
||||
}
|
||||
as := a.RawComplexes()
|
||||
for i := range cs {
|
||||
p := complex(1, 0)
|
||||
for range n {
|
||||
p *= as[i]
|
||||
}
|
||||
cs[i] = p
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,539 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// buildComplex wraps a complex array as a leaf tensor; the losses
|
||||
// built on it reduce to a real scalar, so Backward has a real seed.
|
||||
func buildComplex(t *testing.T, vals []complex128, shape ...int) *Tensor {
|
||||
t.Helper()
|
||||
a, err := core.FromComplexes(vals, shape...)
|
||||
if err != nil {
|
||||
t.Fatalf("FromComplexes: %v", err)
|
||||
}
|
||||
return FromArray(a, true)
|
||||
}
|
||||
|
||||
// numericComplexGrad estimates dL/dRe(z) and dL/dIm(z) by central
|
||||
// differences; the Wirtinger gradient the graph reports must satisfy
|
||||
// g = (dL/dRe + i·dL/dIm)/2 element-wise.
|
||||
func numericComplexGrad(f func(*core.Array) float64, a *core.Array) []complex128 {
|
||||
n := a.Len()
|
||||
out := make([]complex128, n)
|
||||
const h = 1e-6
|
||||
up := make([]complex128, n)
|
||||
down := make([]complex128, n)
|
||||
for i := range n {
|
||||
base := make([]complex128, n)
|
||||
for j := range n {
|
||||
base[j] = a.ComplexAt(j)
|
||||
}
|
||||
copy(up, base)
|
||||
copy(down, base)
|
||||
up[i] += complex(h, 0)
|
||||
down[i] -= complex(h, 0)
|
||||
au, _ := core.FromComplexes(up, a.Shape()...)
|
||||
ad, _ := core.FromComplexes(down, a.Shape()...)
|
||||
dRe := (f(au) - f(ad)) / (2 * h)
|
||||
copy(up, base)
|
||||
copy(down, base)
|
||||
up[i] += complex(0, h)
|
||||
down[i] -= complex(0, h)
|
||||
au, _ = core.FromComplexes(up, a.Shape()...)
|
||||
ad, _ = core.FromComplexes(down, a.Shape()...)
|
||||
dIm := (f(au) - f(ad)) / (2 * h)
|
||||
out[i] = complex(dRe/2, dIm/2)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// checkAgainstNumeric compares a Wirtinger gradient with the central-
|
||||
// difference reference.
|
||||
func checkAgainstNumeric(t *testing.T, got *core.Array, want []complex128, tol float64) {
|
||||
t.Helper()
|
||||
for i, w := range want {
|
||||
g := got.ComplexAt(i)
|
||||
if math.Abs(real(g)-real(w)) > tol || math.Abs(imag(g)-imag(w)) > tol {
|
||||
t.Fatalf("grad[%d] = %v, want %v", i, g, w)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestComplexMulGrad pins the Wirtinger adjoint of the element-wise
|
||||
// product: dz = g·w̄.
|
||||
func TestComplexMulGrad(t *testing.T) {
|
||||
z := buildComplex(t, []complex128{1 + 2i, 3 - 1i, -0.5 + 0.25i}, 3)
|
||||
w := buildComplex(t, []complex128{0.5 - 1i, 2 + 2i, 1 - 3i}, 3)
|
||||
prod, err := z.Mul(w)
|
||||
if err != nil {
|
||||
t.Fatalf("Mul: %v", err)
|
||||
}
|
||||
re, err := prod.Real()
|
||||
if err != nil {
|
||||
t.Fatalf("Real: %v", err)
|
||||
}
|
||||
loss, err := re.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
lossOf := func(a *core.Array) float64 {
|
||||
s := 0.0
|
||||
for i := range a.Len() {
|
||||
s += real(a.ComplexAt(i) * w.Data().ComplexAt(i))
|
||||
}
|
||||
return s
|
||||
}
|
||||
want := numericComplexGrad(lossOf, z.Data())
|
||||
checkAgainstNumeric(t, z.Grad(), want, 1e-8)
|
||||
}
|
||||
|
||||
// TestComplexDivGrad pins the division adjoint da = g/b̄,
|
||||
// db = −g·ā/b̄².
|
||||
func TestComplexDivGrad(t *testing.T) {
|
||||
z := buildComplex(t, []complex128{1 + 2i, 3 - 1i}, 2)
|
||||
w := buildComplex(t, []complex128{0.5 - 1i, 2 + 2i}, 2)
|
||||
q, err := z.Div(w)
|
||||
if err != nil {
|
||||
t.Fatalf("Div: %v", err)
|
||||
}
|
||||
im, err := q.Imag()
|
||||
if err != nil {
|
||||
t.Fatalf("Imag: %v", err)
|
||||
}
|
||||
loss, err := im.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
lossOf := func(a *core.Array) float64 {
|
||||
s := 0.0
|
||||
for i := range a.Len() {
|
||||
s += imag(a.ComplexAt(i) / w.Data().ComplexAt(i))
|
||||
}
|
||||
return s
|
||||
}
|
||||
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(lossOf, z.Data()), 1e-8)
|
||||
checkAgainstNumeric(t, w.Grad(), numericComplexGrad(func(a *core.Array) float64 {
|
||||
s := 0.0
|
||||
for i := range a.Len() {
|
||||
s += imag(z.Data().ComplexAt(i) / a.ComplexAt(i))
|
||||
}
|
||||
return s
|
||||
}, w.Data()), 1e-8)
|
||||
}
|
||||
|
||||
// TestComplexQuantumExpectation pins the physics workhorse: the loss
|
||||
// L = Re(ψ̄·(H·ψ)) with a Hermitian H, whose Wirtinger gradient is
|
||||
// ψ̄-independent and equals H·ψ... evaluated against central
|
||||
// differences rather than trust the algebra.
|
||||
func TestComplexQuantumExpectation(t *testing.T) {
|
||||
psi := buildComplex(t, []complex128{1 + 0.5i, -0.3 + 0.8i, 0.2 - 1.1i, 0.9 + 0.4i}, 4)
|
||||
hDense := []complex128{
|
||||
2, 0.5i, 0, -1,
|
||||
-0.5i, 3, 1i, 0,
|
||||
0, -1i, 1.5, 0.5,
|
||||
-1, 0, 0.5, 2.5,
|
||||
}
|
||||
hArr, err := core.FromComplexes(hDense, 4, 4)
|
||||
if err != nil {
|
||||
t.Fatalf("FromComplexes: %v", err)
|
||||
}
|
||||
h := FromArray(hArr, false)
|
||||
|
||||
conj, err := psi.Conj()
|
||||
if err != nil {
|
||||
t.Fatalf("Conj: %v", err)
|
||||
}
|
||||
// Row vector (1,4) times H·ψ (4,) keeps every MatMul shape legal.
|
||||
bra, err := conj.Reshape(1, 4)
|
||||
if err != nil {
|
||||
t.Fatalf("Reshape: %v", err)
|
||||
}
|
||||
hpsi, err := h.MatMul(psi)
|
||||
if err != nil {
|
||||
t.Fatalf("MatMul: %v", err)
|
||||
}
|
||||
prod, err := bra.MatMul(hpsi)
|
||||
if err != nil {
|
||||
t.Fatalf("MatMul: %v", err)
|
||||
}
|
||||
// prod is (1,1); Real then Sum flattens to the scalar loss.
|
||||
r, err := prod.Real()
|
||||
if err != nil {
|
||||
t.Fatalf("Real: %v", err)
|
||||
}
|
||||
loss, err := r.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
expectation := func(a *core.Array) float64 {
|
||||
s := 0.0
|
||||
for i := range 4 {
|
||||
var acc complex128
|
||||
for j := range 4 {
|
||||
acc += hDense[i*4+j] * a.ComplexAt(j)
|
||||
}
|
||||
s += real(complex(real(a.ComplexAt(i)), -imag(a.ComplexAt(i))) * acc)
|
||||
}
|
||||
return s
|
||||
}
|
||||
checkAgainstNumeric(t, psi.Grad(), numericComplexGrad(expectation, psi.Data()), 1e-7)
|
||||
}
|
||||
|
||||
// TestComplexAbs2Grad pins |z|², dz = g·z.
|
||||
func TestComplexAbs2Grad(t *testing.T) {
|
||||
z := buildComplex(t, []complex128{1 + 2i, -3 + 0.5i}, 2)
|
||||
sq, err := z.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)
|
||||
}
|
||||
// d|z|²/dz̄ = z exactly.
|
||||
for i := range 2 {
|
||||
if z.Grad().ComplexAt(i) != z.Data().ComplexAt(i) {
|
||||
t.Fatalf("grad[%d] = %v, want %v", i, z.Grad().ComplexAt(i), z.Data().ComplexAt(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestComplexAbsGrad pins the magnitude gradient dz = g·z/(2|z|).
|
||||
func TestComplexAbsGrad(t *testing.T) {
|
||||
z := buildComplex(t, []complex128{3 + 4i, -1 + 1i}, 2)
|
||||
m, err := z.Abs()
|
||||
if err != nil {
|
||||
t.Fatalf("Abs: %v", err)
|
||||
}
|
||||
loss, err := m.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
expectation := func(a *core.Array) float64 {
|
||||
s := 0.0
|
||||
for i := range a.Len() {
|
||||
s += cmplxAbs(a.ComplexAt(i))
|
||||
}
|
||||
return s
|
||||
}
|
||||
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(expectation, z.Data()), 1e-8)
|
||||
}
|
||||
|
||||
func cmplxAbs(z complex128) float64 { return math.Hypot(real(z), imag(z)) }
|
||||
|
||||
// TestComplexBackwardRejectsComplexLoss pins the real-seed contract.
|
||||
func TestComplexBackwardRejectsComplexLoss(t *testing.T) {
|
||||
z := buildComplex(t, []complex128{1 + 1i}, 1)
|
||||
if err := z.Backward(); err == nil {
|
||||
t.Fatal("Backward accepted a complex output")
|
||||
}
|
||||
}
|
||||
|
||||
// TestComplexMixedRealLeaf pins the 2·Re narrowing: a real tensor
|
||||
// multiplied into a complex chain gets the true real gradient.
|
||||
func TestComplexMixedRealLeaf(t *testing.T) {
|
||||
xArr, err := core.FromFloats([]float64{1.5, -0.5}, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
x := FromArray(xArr, true)
|
||||
w := buildComplex(t, []complex128{0.5 - 1i, 2 + 2i}, 2)
|
||||
prod, err := x.Mul(w)
|
||||
if err != nil {
|
||||
t.Fatalf("Mul: %v", err)
|
||||
}
|
||||
sq, err := prod.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)
|
||||
}
|
||||
// L = Σ x²|w|², dL/dx = 2x|w|².
|
||||
want := []float64{2 * 1.5 * (0.25 + 1), 2 * -0.5 * (4 + 4)}
|
||||
for i, wv := range want {
|
||||
if math.Abs(x.Grad().FloatAt(i)-wv) > 1e-12 {
|
||||
t.Fatalf("x.grad[%d] = %g, want %g", i, x.Grad().FloatAt(i), wv)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestComplexSumAxisGrad pins the axis reduction on a complex leaf:
|
||||
// the backward broadcasts the seed back over the dropped axis, and the
|
||||
// Wirtinger gradient of a weighted real loss matches central
|
||||
// differences.
|
||||
func TestComplexSumAxisGrad(t *testing.T) {
|
||||
z := buildComplex(t, []complex128{1 + 2i, -0.5 + 0.25i, 0.75 - 1.5i, 2 + 0.5i}, 2, 2)
|
||||
out, err := z.SumAxis(1)
|
||||
if err != nil {
|
||||
t.Fatalf("SumAxis: %v", err)
|
||||
}
|
||||
if out.Data().NDim() != 1 || out.Data().Len() != 2 {
|
||||
t.Fatalf("SumAxis shape = %s, want (2)", prettyShape(out.Data().Shape()))
|
||||
}
|
||||
w := buildComplex(t, []complex128{0.3 + 0.4i, -0.6 - 0.1i}, 2)
|
||||
prod, err := out.Mul(w)
|
||||
if err != nil {
|
||||
t.Fatalf("Mul: %v", err)
|
||||
}
|
||||
re, err := prod.Real()
|
||||
if err != nil {
|
||||
t.Fatalf("Real: %v", err)
|
||||
}
|
||||
loss, err := re.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
lossOf := func(a *core.Array) float64 {
|
||||
s := 0.0
|
||||
for j := range 2 {
|
||||
var acc complex128
|
||||
for k := range 2 {
|
||||
acc += a.ComplexAt(j*2 + k)
|
||||
}
|
||||
s += real(acc * w.Data().ComplexAt(j))
|
||||
}
|
||||
return s
|
||||
}
|
||||
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(lossOf, z.Data()), 1e-8)
|
||||
}
|
||||
|
||||
// TestComplexMotionOps pins gradient flow through Slice, Concat,
|
||||
// Reshape and BroadcastTo on complex tensors.
|
||||
func TestComplexMotionOps(t *testing.T) {
|
||||
z := buildComplex(t, []complex128{1 + 1i, 2 - 1i, 3 + 2i, 4 - 3i}, 4)
|
||||
sl, err := z.Slice(0, 1, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("Slice: %v", err)
|
||||
}
|
||||
sq, err := sl.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)
|
||||
}
|
||||
// Only the sliced elements receive gradient, and for L = Σ|z|² it
|
||||
// is exactly z_i.
|
||||
for i := range 4 {
|
||||
want := complex(0, 0)
|
||||
if i == 1 || i == 2 {
|
||||
want = z.Data().ComplexAt(i)
|
||||
}
|
||||
if z.Grad().ComplexAt(i) != want {
|
||||
t.Fatalf("grad[%d] = %v, want %v", i, z.Grad().ComplexAt(i), want)
|
||||
}
|
||||
}
|
||||
|
||||
z2 := buildComplex(t, []complex128{0.5 + 0.5i}, 1)
|
||||
cat, err := z.Concat(z2, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("Concat: %v", err)
|
||||
}
|
||||
sq2, err := cat.Abs2()
|
||||
if err != nil {
|
||||
t.Fatalf("Abs2: %v", err)
|
||||
}
|
||||
loss2, err := sq2.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss2.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
if z2.Grad().ComplexAt(0) != 0.5+0.5i {
|
||||
t.Fatalf("concat gradient = %v, want 0.5+0.5i", z2.Grad().ComplexAt(0))
|
||||
}
|
||||
}
|
||||
|
||||
// TestComplexSumMeanPow pins the complex reducers and integer powers.
|
||||
func TestComplexSumMeanPow(t *testing.T) {
|
||||
z := buildComplex(t, []complex128{1 + 2i, 3 - 1i}, 2)
|
||||
m, err := z.Mean()
|
||||
if err != nil {
|
||||
t.Fatalf("Mean: %v", err)
|
||||
}
|
||||
if m.Data().ComplexAt(0) != 2+0.5i {
|
||||
t.Fatalf("mean = %v, want 2+0.5i", m.Data().ComplexAt(0))
|
||||
}
|
||||
p, err := z.Pow(3)
|
||||
if err != nil {
|
||||
t.Fatalf("Pow: %v", err)
|
||||
}
|
||||
// (1+2i)³ = (1+2i)(1+2i)(1+2i) = -11-2i.
|
||||
if p.Data().ComplexAt(0) != -11-2i {
|
||||
t.Fatalf("pow = %v, want -11-2i", p.Data().ComplexAt(0))
|
||||
}
|
||||
sq, err := p.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)
|
||||
}
|
||||
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(func(a *core.Array) float64 {
|
||||
s := 0.0
|
||||
for i := range a.Len() {
|
||||
pv := complex(1, 0)
|
||||
for range 3 {
|
||||
pv *= a.ComplexAt(i)
|
||||
}
|
||||
s += real(pv)*real(pv) + imag(pv)*imag(pv)
|
||||
}
|
||||
return s
|
||||
}, z.Data()), 1e-6)
|
||||
}
|
||||
|
||||
// TestComplexMatMulGrad pins the 2-D complex matmul adjoint against
|
||||
// central differences.
|
||||
func TestComplexMatMulGrad(t *testing.T) {
|
||||
aVals := []complex128{1 + 1i, 2 - 1i, 0.5 + 0i, -1 + 2i}
|
||||
bVals := []complex128{0.5 - 0.5i, 1 + 1i, -0.5 + 2i, 0.25 - 0.75i}
|
||||
a := buildComplex(t, aVals, 2, 2)
|
||||
b := buildComplex(t, bVals, 2, 2)
|
||||
y, err := a.MatMul(b)
|
||||
if err != nil {
|
||||
t.Fatalf("MatMul: %v", err)
|
||||
}
|
||||
sq, err := y.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)
|
||||
}
|
||||
lossOf := func(av, bv []complex128) float64 {
|
||||
s := 0.0
|
||||
for i := range 2 {
|
||||
for j := range 2 {
|
||||
var acc complex128
|
||||
for k := range 2 {
|
||||
acc += av[i*2+k] * bv[k*2+j]
|
||||
}
|
||||
s += real(acc)*real(acc) + imag(acc)*imag(acc)
|
||||
}
|
||||
}
|
||||
return s
|
||||
}
|
||||
checkAgainstNumeric(t, a.Grad(), numericComplexGrad(func(arr *core.Array) float64 {
|
||||
return lossOf(flatComplex(arr), bVals)
|
||||
}, a.Data()), 1e-7)
|
||||
checkAgainstNumeric(t, b.Grad(), numericComplexGrad(func(arr *core.Array) float64 {
|
||||
return lossOf(aVals, flatComplex(arr))
|
||||
}, b.Data()), 1e-7)
|
||||
}
|
||||
|
||||
func flatComplex(a *core.Array) []complex128 {
|
||||
out := make([]complex128, a.Len())
|
||||
for i := range out {
|
||||
out[i] = a.ComplexAt(i)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestComplexScaleGrad pins Scale on complex tensors.
|
||||
func TestComplexScaleGrad(t *testing.T) {
|
||||
z := buildComplex(t, []complex128{1 + 1i, 2 - 1i}, 2)
|
||||
s, err := z.Scale(2.5)
|
||||
if err != nil {
|
||||
t.Fatalf("Scale: %v", err)
|
||||
}
|
||||
sq, err := s.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)
|
||||
}
|
||||
// d(2.5²|z|²)/dz̄ = 2·2.5²·Re-parts... exact: 6.25·z.
|
||||
for i := range 2 {
|
||||
want := 6.25 * z.Data().ComplexAt(i)
|
||||
got := z.Grad().ComplexAt(i)
|
||||
if math.Abs(real(got)-real(want)) > 1e-10 || math.Abs(imag(got)-imag(want)) > 1e-10 {
|
||||
t.Fatalf("grad[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestComplexExpGrad pins the complex exponential adjoint dz = g·conj(e^z)
|
||||
// against central differences, and e^z against the polar identity
|
||||
// e^{x+iy} = e^x(cos y + i sin y).
|
||||
func TestComplexExpGrad(t *testing.T) {
|
||||
vals := []complex128{0.3 - 0.2i, -1.1 + 0.7i, 0.05 + 0i}
|
||||
z := buildComplex(t, vals, 3)
|
||||
e, err := z.Exp()
|
||||
if err != nil {
|
||||
t.Fatalf("Exp: %v", err)
|
||||
}
|
||||
for i, v := range vals {
|
||||
want := complex(math.Exp(real(v)), 0) * complex(math.Cos(imag(v)), math.Sin(imag(v)))
|
||||
got := e.Data().ComplexAt(i)
|
||||
if cmplxAbs(got-want) > 1e-14 {
|
||||
t.Fatalf("exp[%d] = %v, want %v", i, got, want)
|
||||
}
|
||||
}
|
||||
sq, err := e.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)
|
||||
}
|
||||
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(func(a *core.Array) float64 {
|
||||
s := 0.0
|
||||
for i := range a.Len() {
|
||||
vv := a.ComplexAt(i)
|
||||
ev := complex(math.Exp(real(vv)), 0) * complex(math.Cos(imag(vv)), math.Sin(imag(vv)))
|
||||
s += real(ev)*real(ev) + imag(ev)*imag(ev)
|
||||
}
|
||||
return s
|
||||
}, z.Data()), 1e-7)
|
||||
}
|
||||
+128
@@ -0,0 +1,128 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Concat joins u after t along an existing axis, the differentiable
|
||||
// inverse of Slice, and the building block that lets recurrent layers
|
||||
// assemble per-step outputs into one sequence core.
|
||||
|
||||
// Concat returns the tensors joined along the given existing dimension.
|
||||
// Every other dimension must agree. The backward routes each side its
|
||||
// own span of the incoming gradient along that dimension, narrowed to
|
||||
// the side's dtype first: a real side of a real/complex join receives
|
||||
// 2·Re(g), the rule every other mixed-dtype op applies.
|
||||
func (t *Tensor) Concat(u *Tensor, dim int) (*Tensor, error) {
|
||||
if err := t.checkDiff("Concat"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := u.checkDiff("Concat"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out, err := core.Concat(t.data, u.data, dim)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
at, au := t.data, u.data
|
||||
return binaryResult("Concat", t, u, out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
dt := at.Dtype()
|
||||
du := au.Dtype()
|
||||
// One narrowing per side before either span is copied: the
|
||||
// concat's own dtype promotes along the ladder, so a complex
|
||||
// gradient reaches a real operand in a mixed join and must
|
||||
// narrow by 2·Re exactly as the leaf commit would. The narrowed
|
||||
// arrays also put both copies on the dtype-matched raw path.
|
||||
gA, err := narrowGradient(g, dt)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
gB, err := narrowGradient(g, du)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
shA, shB := at.Shape(), au.Shape()
|
||||
outer, inner, err := outerInner(shA, dim)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
spanA := shA[dim]
|
||||
total := spanA + shB[dim]
|
||||
da := gradSlot{arr: ar.borrowGrad(dt, shA), sh: shA}
|
||||
// Every row contributes one contiguous inner run, so matching
|
||||
// dtypes collapse the triple loop to a raw slice move per row.
|
||||
fastA := !gA.arr.Strided() && gA.arr.Dtype() == dt && dt != core.Int
|
||||
for o := range outer {
|
||||
for i := range spanA {
|
||||
d := o*spanA*inner + i*inner
|
||||
s := o*total*inner + i*inner
|
||||
if fastA {
|
||||
copySegRaw(da.arr, gA.arr, d, s, inner)
|
||||
continue
|
||||
}
|
||||
for j := range inner {
|
||||
copyElem(da.arr, d+j, gA.arr, s+j)
|
||||
}
|
||||
}
|
||||
}
|
||||
db := gradSlot{arr: ar.borrowGrad(du, shB), sh: shB}
|
||||
spanB := shB[dim]
|
||||
fastB := !gB.arr.Strided() && gB.arr.Dtype() == du && du != core.Int
|
||||
for o := range outer {
|
||||
for i := range spanB {
|
||||
d := o*spanB*inner + i*inner
|
||||
s := o*total*inner + (spanA+i)*inner
|
||||
if fastB {
|
||||
copySegRaw(db.arr, gB.arr, d, s, inner)
|
||||
continue
|
||||
}
|
||||
for j := range inner {
|
||||
copyElem(db.arr, d+j, gB.arr, s+j)
|
||||
}
|
||||
}
|
||||
}
|
||||
dst[0], dst[1] = da, db
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// copySegRaw moves n elements from src at sOff to dst at dOff through
|
||||
// the raw payloads. The caller checks dtype equality and contiguity;
|
||||
// the per-element values, and so the bits, are the ones copyElem
|
||||
// writes one accessor call at a time.
|
||||
func copySegRaw(dst, src *core.Array, dOff, sOff, n int) {
|
||||
switch src.Dtype() {
|
||||
case core.Float32:
|
||||
copy(dst.RawFloat32s()[dOff:dOff+n], src.RawFloat32s()[sOff:sOff+n])
|
||||
case core.Float:
|
||||
copy(dst.RawFloats()[dOff:dOff+n], src.RawFloats()[sOff:sOff+n])
|
||||
case core.Complex:
|
||||
copy(dst.RawComplexes()[dOff:dOff+n], src.RawComplexes()[sOff:sOff+n])
|
||||
default:
|
||||
copy(dst.RawInts()[dOff:dOff+n], src.RawInts()[sOff:sOff+n])
|
||||
}
|
||||
}
|
||||
|
||||
// outerInner splits a shape into the products of the dimensions before
|
||||
// and after dim, the strides a flat row-major walk needs when only one
|
||||
// axis is being split or joined.
|
||||
func outerInner(shape []int, dim int) (int, int, error) {
|
||||
if len(shape) == 0 {
|
||||
return 0, 0, errf("Concat: cannot concatenate a scalar")
|
||||
}
|
||||
if dim < 0 || dim >= len(shape) {
|
||||
return 0, 0, errf("Concat: dimension %d is out of range for shape %v", dim, shape)
|
||||
}
|
||||
outer := 1
|
||||
for d := range dim {
|
||||
outer *= shape[d]
|
||||
}
|
||||
inner := 1
|
||||
for d := dim + 1; d < len(shape); d++ {
|
||||
inner *= shape[d]
|
||||
}
|
||||
return outer, inner, nil
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
func TestTensorConcatForward(t *testing.T) {
|
||||
a, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
|
||||
b, _ := core.FromFloats([]float64{5, 6, 7, 8, 9, 10}, 2, 3)
|
||||
|
||||
out, err := FromArray(a, false).Concat(FromArray(b, false), 1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := out.Data().Shape(); got[0] != 2 || got[1] != 5 {
|
||||
t.Fatalf("concat shape: %v", got)
|
||||
}
|
||||
want := []float64{1, 2, 5, 6, 7, 3, 4, 8, 9, 10}
|
||||
for i := range want {
|
||||
if g := out.Data().FloatAt(i); g != want[i] {
|
||||
t.Fatalf("concat[%d] = %v, want %v", i, g, want[i])
|
||||
}
|
||||
}
|
||||
|
||||
// Concatenation along the leading axis stacks the blocks.
|
||||
c, _ := core.FromFloats([]float64{1, 2}, 1, 2)
|
||||
d, _ := core.FromFloats([]float64{3, 4}, 1, 2)
|
||||
vert, err := FromArray(c, false).Concat(FromArray(d, false), 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if vert.Data().Shape()[0] != 2 {
|
||||
t.Fatalf("vertical shape: %v", vert.Data().Shape())
|
||||
}
|
||||
|
||||
// Mismatched ranks and out-of-range axes error.
|
||||
misrank, _ := core.Reshape(c, 2)
|
||||
if _, err := FromArray(a, false).Concat(FromArray(misrank, false), 0); err == nil {
|
||||
t.Fatal("rank mismatch accepted")
|
||||
}
|
||||
if _, err := FromArray(a, false).Concat(FromArray(b, false), 2); err == nil {
|
||||
t.Fatal("out-of-range dimension accepted")
|
||||
}
|
||||
}
|
||||
|
||||
// TestTensorConcatGradients checks both backward spans against central
|
||||
// differences with a weighted loss so every slot gets a distinct weight.
|
||||
func TestTensorConcatGradients(t *testing.T) {
|
||||
cases := []struct {
|
||||
dim int
|
||||
aVal, bVal []float64
|
||||
aShape, bShape []int
|
||||
waVal, wbVal []float64
|
||||
}{
|
||||
{
|
||||
dim: 1,
|
||||
aVal: []float64{0.5, -1, 2, 0.25}, aShape: []int{2, 2},
|
||||
bVal: []float64{1.5, -0.5, 1, 2, -2, 0.75}, bShape: []int{2, 3},
|
||||
waVal: []float64{0.1, -0.4, 0.9, 0.6},
|
||||
wbVal: []float64{0.2, 0.3, -0.7, 0.8, 0.05, -0.6},
|
||||
},
|
||||
{
|
||||
dim: 0,
|
||||
aVal: []float64{0.3, 1, -0.25, 2}, aShape: []int{2, 2},
|
||||
bVal: []float64{-1.5, 0.4, 0.9, 1, -2, 0.7}, bShape: []int{3, 2},
|
||||
waVal: []float64{0.55, -0.35, 0.85, 0.15},
|
||||
wbVal: []float64{0.45, -0.65, 0.95, 0.05, -0.5, 0.75},
|
||||
},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
a, _ := core.FromFloats(tc.aVal, tc.aShape...)
|
||||
b, _ := core.FromFloats(tc.bVal, tc.bShape...)
|
||||
wa, _ := core.FromFloats(tc.waVal, tc.aShape...)
|
||||
wb, _ := core.FromFloats(tc.wbVal, tc.bShape...)
|
||||
|
||||
at := FromArray(a, true)
|
||||
bt := FromArray(b, true)
|
||||
joint, err := at.Concat(bt, tc.dim)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
wc, _ := core.Concat(wa, wb, tc.dim)
|
||||
scaled, err := joint.Mul(FromArray(wc, false))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loss, err := scaled.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
fa := func(v *core.Array) float64 { return weightedConcatSum(v, b, wa, wb, tc.dim) }
|
||||
fb := func(v *core.Array) float64 { return weightedConcatSum(a, v, wa, wb, tc.dim) }
|
||||
checkSpan(t, at.Grad(), numericGrad(fa, a))
|
||||
checkSpan(t, bt.Grad(), numericGrad(fb, b))
|
||||
}
|
||||
}
|
||||
|
||||
// TestTensorConcatGradientDtype keeps each side's gradient in its own
|
||||
// element type: float32 inputs never come back as float64 leaves.
|
||||
func TestTensorConcatGradientDtype(t *testing.T) {
|
||||
gen := core.NewGenerator(5)
|
||||
af, _ := core.Float32s(gen, 6)
|
||||
afArr, _ := core.Reshape(af, 2, 3)
|
||||
bf, _ := core.Float32s(gen, 6)
|
||||
bfArr, _ := core.Reshape(bf, 2, 3)
|
||||
|
||||
at := FromArray(afArr, true)
|
||||
bt := FromArray(bfArr, true)
|
||||
joint, err := at.Concat(bt, 1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loss, err := joint.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if at.Grad().Dtype() != core.Float32 || bt.Grad().Dtype() != core.Float32 {
|
||||
t.Fatalf("gradient dtypes: %v and %v", at.Grad().Dtype(), bt.Grad().Dtype())
|
||||
}
|
||||
for i := range at.Grad().Len() {
|
||||
if at.Grad().FloatAt(i) != 1 {
|
||||
t.Errorf("float32 gradient slot %d: %v, want 1", i, at.Grad().FloatAt(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// weightedConcatSum evaluates Σ w∘Concat(x, y, dim) with fixed weights,
|
||||
// the scalar objective whose gradients the backward is checked against.
|
||||
func weightedConcatSum(x, y *core.Array, wx, wy *core.Array, dim int) float64 {
|
||||
joint, err := FromArray(x, false).Concat(FromArray(y, false), dim)
|
||||
if err != nil {
|
||||
return math.NaN()
|
||||
}
|
||||
wc, _ := core.Concat(wx, wy, dim)
|
||||
total := 0.0
|
||||
for i := range joint.Data().Len() {
|
||||
total += wc.FloatAt(i) * joint.Data().FloatAt(i)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
// checkSpan reports every slot where the analytic gradient drifts from
|
||||
// the central-difference reference.
|
||||
func checkSpan(t *testing.T, got *core.Array, ref []float64) {
|
||||
t.Helper()
|
||||
if got.Len() != len(ref) {
|
||||
t.Fatalf("gradient length %d, reference %d", got.Len(), len(ref))
|
||||
}
|
||||
for i := range ref {
|
||||
if math.Abs(got.FloatAt(i)-ref[i]) > 1e-5 {
|
||||
t.Errorf("gradient[%d] = %v, want ≈%v", i, got.FloatAt(i), ref[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,200 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
func TestTensorSqueezeUnsqueezeClip(t *testing.T) {
|
||||
// Squeeze/Unsqueeze round-trip with gradient.
|
||||
x, _ := core.FromFloats([]float64{1, 2, 3, 4}, 1, 4, 1)
|
||||
xt := FromArray(x, true)
|
||||
sq, err := xt.Squeeze(2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if sq.Data().NDim() != 2 {
|
||||
t.Fatalf("Squeeze ndim: %d", sq.Data().NDim())
|
||||
}
|
||||
back, err := sq.Unsqueeze(2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s, _ := back.Sum()
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range 4 {
|
||||
if g := xt.Grad().FloatAt(i); g != 1 {
|
||||
t.Errorf("Squeeze/Unsqueeze grad[%d]: %v, want 1", i, g)
|
||||
}
|
||||
}
|
||||
|
||||
// Clip gradient: 1 inside [lo, hi], 0 outside.
|
||||
c, _ := core.FromFloats([]float64{-1, 0.5, 2}, 3)
|
||||
ct := FromArray(c, true)
|
||||
cl, err := ct.Clip(0, 1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s2, _ := cl.Sum()
|
||||
if err := s2.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := []float64{0, 1, 0}
|
||||
for i := range 3 {
|
||||
if g := ct.Grad().FloatAt(i); g != want[i] {
|
||||
t.Errorf("Clip grad[%d]: %v, want %v", i, g, want[i])
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
func TestAxisReductionAutograd(t *testing.T) {
|
||||
x, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
||||
xt := FromArray(x, true)
|
||||
|
||||
s, err := xt.SumAxis(1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if s.Data().Len() != 2 {
|
||||
t.Fatalf("SumAxis len: %d", s.Data().Len())
|
||||
}
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range x.Len() {
|
||||
if g := xt.Grad().FloatAt(i); g != 1 {
|
||||
t.Errorf("SumAxis grad[%d]: %v, want 1", i, g)
|
||||
}
|
||||
}
|
||||
|
||||
mean, err := FromArray(x, true).MeanAxis(1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := mean.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestL2NormAxisAutogradGradient(t *testing.T) {
|
||||
xv := []float64{3, 4, 0.5, 0.5}
|
||||
x, _ := core.FromFloats(xv, 1, 1, 2, 2)
|
||||
xt := FromArray(x, true)
|
||||
out, err := xt.L2NormAxis(1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s, _ := out.Sum()
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
analytic := make([]float64, x.Len())
|
||||
for i := range x.Len() {
|
||||
analytic[i] = xt.Grad().FloatAt(i)
|
||||
}
|
||||
ref := numericGrad(func(a *core.Array) float64 {
|
||||
o, err := FromArray(a, false).L2NormAxis(1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ss, _ := o.Sum()
|
||||
return ss.Data().FloatAt(0)
|
||||
}, x)
|
||||
if d := maxAbsDiff(analytic, ref); d > 1e-6 {
|
||||
t.Errorf("L2NormAxis grad: max diff %v", d)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBroadcastToAutograd(t *testing.T) {
|
||||
x, _ := core.FromFloats([]float64{1, 2, 3}, 1, 3)
|
||||
xt := FromArray(x, true)
|
||||
out, err := xt.BroadcastTo(2, 3)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if out.Data().Shape()[0] != 2 {
|
||||
t.Fatalf("BroadcastTo shape: %v", out.Data().Shape())
|
||||
}
|
||||
onesArr, _ := core.Ones(core.Float, 2, 3)
|
||||
loss, _ := out.Mul(FromArray(onesArr, false))
|
||||
s, _ := loss.Sum()
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Gradient sums over replicated rows.
|
||||
for i := range 3 {
|
||||
if g := xt.Grad().FloatAt(i); g != 2 {
|
||||
t.Errorf("BroadcastTo grad[%d]: %v, want 2", i, g)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPowAbsSqrtFloorAutogradGradient(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
vals []float64
|
||||
fn func(*Tensor) (*Tensor, error)
|
||||
}{
|
||||
{"Pow3", []float64{0.5, 1.5}, func(x *Tensor) (*Tensor, error) { return x.Pow(3) }},
|
||||
{"Abs", []float64{0.5, -1.5}, func(x *Tensor) (*Tensor, error) { return x.Abs() }},
|
||||
{"Sqrt", []float64{0.25, 2.25}, func(x *Tensor) (*Tensor, error) { return x.Sqrt() }},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
a, _ := core.FromFloats(tc.vals, 2)
|
||||
at := FromArray(a, true)
|
||||
out, err := tc.fn(at)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s, err := out.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
analytic := make([]float64, a.Len())
|
||||
for i := range a.Len() {
|
||||
analytic[i] = at.Grad().FloatAt(i)
|
||||
}
|
||||
ref := numericGrad(func(v *core.Array) float64 {
|
||||
o, err := tc.fn(FromArray(v, false))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ss, err := o.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return ss.Data().FloatAt(0)
|
||||
}, a)
|
||||
if d := maxAbsDiff(analytic, ref); d > 1e-6 {
|
||||
t.Errorf("%s grad: max diff %v", tc.name, d)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Floor contributes no gradient.
|
||||
a, _ := core.FromFloats([]float64{1.4, 2.6}, 2)
|
||||
at := FromArray(a, true)
|
||||
fl, err := at.Floor()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s, _ := fl.Sum()
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range a.Len() {
|
||||
if g := at.Grad().FloatAt(i); g != 0 {
|
||||
t.Errorf("Floor grad[%d]: %v, want 0", i, g)
|
||||
}
|
||||
}
|
||||
}
|
||||
+82
@@ -0,0 +1,82 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
// Package grad is reverse-mode automatic differentiation over the array
|
||||
// surface: a computation written as ordinary Go calls over [Tensor]
|
||||
// values records a graph, and one call to [Tensor.Backward] propagates
|
||||
// the gradient from the output back to every leaf that requires it.
|
||||
//
|
||||
// # The graph
|
||||
//
|
||||
// Leaves come from [FromFloat64s] or [FromArray]. Every differentiable
|
||||
// method records the operation it performs together with its inputs, and
|
||||
// Backward sweeps the recorded nodes in reverse, applying each node's
|
||||
// adjoint. The method set covers the arithmetic, the matrix products
|
||||
// (single and batched), the element-wise transcendentals, the
|
||||
// reductions, slicing, concatenation, axis permutation and the Fourier
|
||||
// transforms, with [Tensor.Conj], [Tensor.Real], [Tensor.Imag] and
|
||||
// [Tensor.Abs2] carrying complex values into a real loss.
|
||||
//
|
||||
// loss, err := total.Mean() // the forward pass records the nodes
|
||||
// if err != nil {
|
||||
// return err
|
||||
// }
|
||||
// if err := loss.Backward(); err != nil {
|
||||
// return err
|
||||
// }
|
||||
// x.Grad() // the accumulated gradient
|
||||
//
|
||||
// # The contract
|
||||
//
|
||||
// - float, float32 and complex128 tensors differentiate; an int tensor
|
||||
// is refused by every differentiable method.
|
||||
// - The loss must be real: Backward seeds the output with ones, and a
|
||||
// complex output is rejected with an error naming Real, Imag, Abs
|
||||
// and Abs2 as the reducers that turn it into a real scalar.
|
||||
// - Gradients carry the leaf's dtype. A mixed-dtype graph narrows each
|
||||
// gradient to the dtype of the tensor it accumulates into before the
|
||||
// leaf is written.
|
||||
// - Backward accumulates into the gradients already present, so
|
||||
// [Tensor.ZeroGrad] precedes a fresh pass unless accumulation is
|
||||
// wanted.
|
||||
// - The graph is rebuilt on every forward pass. Each operation's
|
||||
// backward closure captures the operands as they were when the
|
||||
// operation ran, so a [Tensor.ReplaceWith] afterwards changes the
|
||||
// next pass and not the recorded one.
|
||||
//
|
||||
// # Complex graphs
|
||||
//
|
||||
// Complex tensors differentiate under the Wirtinger convention: the
|
||||
// gradient a complex leaf accumulates is ∂L/∂z̄, the coefficient g of
|
||||
// dL = 2·Re(g·dz), which is the direction gradient descent steps along.
|
||||
// The adjoint of a holomorphic y = f(z) is therefore dz = g·conj(f′(z)),
|
||||
// and every complex adjoint in this package conjugates exactly where the
|
||||
// calculus puts it. A real tensor inside a complex graph narrows the
|
||||
// incoming gradient by 2·Re, the factor that also cancels the ½ the Real
|
||||
// and Imag backward paths contribute, so a graph mixing the two dtypes
|
||||
// composes exactly.
|
||||
//
|
||||
// # Second-order and solver tools
|
||||
//
|
||||
// On top of the graph sit the second derivative and the methods that
|
||||
// need one: [Hessian] (dense, 2n gradient evaluations),
|
||||
// [HessianVectorProduct] (H·v in two gradient evaluations),
|
||||
// [MinimiseNewtonCG] (truncated conjugate gradients on the Hessian
|
||||
// system with an Armijo line search), [SampleHMC] (Hamiltonian Monte
|
||||
// Carlo on any differentiable unnormalised log density) and [AdjointODE]
|
||||
// (adjoint sensitivities of an ODE solution at the cost of one extra
|
||||
// solve).
|
||||
//
|
||||
// Each of them differentiates through a reverse pass that commits
|
||||
// nothing, so the accumulated gradients of the tensors the caller's
|
||||
// closure holds are left exactly as they were, whether the call succeeds
|
||||
// or fails.
|
||||
//
|
||||
// # What it does not do
|
||||
//
|
||||
// There is no forward-mode differentiation, nothing beyond the second
|
||||
// derivative, no graph serialisation and no parameter registry: the
|
||||
// caller owns the leaves and the graph is a transient record of one
|
||||
// forward pass. The operation set is closed, so a new primitive is a new
|
||||
// method here and never a user-registered op.
|
||||
package grad
|
||||
@@ -0,0 +1,218 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad_test
|
||||
|
||||
// The godoc examples for the autograd package: the flagship workflows
|
||||
// as runnable, checked snippets. Each one pins the numbers it prints,
|
||||
// so a change in the adjoint of an op or in a solver's default shows
|
||||
// up as a failing example rather than as stale prose.
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
|
||||
tensor "sourcedock.dev/petrbalvin/tensor"
|
||||
"sourcedock.dev/petrbalvin/tensor/grad"
|
||||
)
|
||||
|
||||
// A scalar loss by hand on a small graph: z = Σ x² over a two-element
|
||||
// leaf, differentiated in one reverse sweep. The leaf is entered twice
|
||||
// by the product, so the backward adds both contributions and the
|
||||
// answer is the 2x the calculus gives.
|
||||
func ExampleTensor_Backward() {
|
||||
x, err := grad.FromFloat64s([]float64{2, 3}, true, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
sq, err := x.Mul(x)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
loss, err := sq.Sum()
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Println(loss.Data(), x.Grad())
|
||||
// Output: float (1) [13] float (2) [4, 6]
|
||||
}
|
||||
|
||||
// A matrix product and a reduction as graph nodes: the gradient of the
|
||||
// sum of A·B is ones·Bᵀ, one row sum of B per row of A.
|
||||
func ExampleTensor_MatMul() {
|
||||
a, err := grad.FromFloat64s([]float64{1, 2, 3, 4, 5, 6}, true, 2, 3)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
b, err := grad.FromFloat64s([]float64{1, 0, 0, 1, 1, 1}, false, 3, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
prod, err := a.MatMul(b)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
total, err := prod.Sum()
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
if err := total.Backward(); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Println(total.Data(), a.Grad())
|
||||
// Output: float (1) [30] float (2, 3) [1, 1, 2, 1, 1, 2]
|
||||
}
|
||||
|
||||
// A complex graph with a real loss. The leaf is complex, the loss is
|
||||
// Σ|z|², and the gradient a complex leaf accumulates is ∂L/∂z̄, which
|
||||
// for |z|² is z itself.
|
||||
func ExampleTensor_Abs2() {
|
||||
data, err := tensor.FromComplexes([]complex128{1 + 2i, 3 - 1i}, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
z := grad.FromArray(data, true)
|
||||
magnitude, err := z.Abs2()
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
loss, err := magnitude.Sum()
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Println(loss.Data(), z.Grad())
|
||||
// Output: float (1) [15] complex (2) [(1+2i), (3-1i)]
|
||||
}
|
||||
|
||||
// The dense second derivative of Σ x² at (1, 2): the Hessian of a
|
||||
// quadratic form is twice its matrix, here 2·I.
|
||||
func ExampleHessian() {
|
||||
f := func(x *grad.Tensor) (*grad.Tensor, error) {
|
||||
sq, err := x.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sq.Sum()
|
||||
}
|
||||
point, err := grad.FromFloat64s([]float64{1, 2}, true, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
h, err := grad.Hessian(f, point, grad.HessianOptions{})
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Println(h)
|
||||
// Output: float (2, 2) [2, 0, 0, 2]
|
||||
}
|
||||
|
||||
// The same function and point contracted with a direction: H·v in two
|
||||
// gradient evaluations instead of the dense Hessian's four.
|
||||
func ExampleHessianVectorProduct() {
|
||||
f := func(x *grad.Tensor) (*grad.Tensor, error) {
|
||||
sq, err := x.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sq.Sum()
|
||||
}
|
||||
point, err := grad.FromFloat64s([]float64{1, 2}, true, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
direction, err := grad.FromFloat64s([]float64{1, 1}, false, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
hv, err := grad.HessianVectorProduct(f, point, direction, grad.HessianOptions{})
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
// A central difference along the direction, so the answer carries
|
||||
// the rounding of the two gradient evaluations it is built from.
|
||||
fmt.Printf("(%.4f, %.4f)\n", hv.FloatAt(0), hv.FloatAt(1))
|
||||
// Output: (2.0000, 2.0000)
|
||||
}
|
||||
|
||||
// Newton-CG on the quadratic Σ (x − c)², whose minimiser is c and
|
||||
// whose value there is zero. The curvature comes from the
|
||||
// Hessian-vector product, so no dense Hessian is ever formed.
|
||||
func ExampleMinimiseNewtonCG() {
|
||||
centre, err := grad.FromFloat64s([]float64{1.5, -2.5}, false, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
f := func(x *grad.Tensor) (*grad.Tensor, error) {
|
||||
diff, err := x.Sub(centre)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := diff.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sq.Sum()
|
||||
}
|
||||
x0, err := tensor.FromFloats([]float64{0, 0}, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
x, value, err := grad.MinimiseNewtonCG(f, x0, grad.NewtonCGOptions{})
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
fmt.Printf("x = (%.4f, %.4f), f = %.4f\n", x.FloatAt(0), x.FloatAt(1), value)
|
||||
// Output: x = (1.5000, -2.5000), f = 0.0000
|
||||
}
|
||||
|
||||
// Hamiltonian Monte Carlo on the two-dimensional standard normal,
|
||||
// whose log density is −‖q‖²/2. The seed makes the chain reproducible,
|
||||
// so the moments of the first component are fixed numbers and not a
|
||||
// range: the target has mean zero and variance one.
|
||||
func ExampleSampleHMC() {
|
||||
logDensity := func(q *grad.Tensor) (*grad.Tensor, error) {
|
||||
sq, err := q.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
total, err := sq.Sum()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return total.Scale(-0.5)
|
||||
}
|
||||
q0, err := tensor.FromFloats([]float64{2, -2}, 2)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
samples, err := grad.SampleHMC(logDensity, q0, grad.HMCOptions{
|
||||
Step: 0.25,
|
||||
Steps: 16,
|
||||
BurnIn: 500,
|
||||
Thin: 1,
|
||||
Samples: 2000,
|
||||
Seed: 7,
|
||||
})
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
rows := samples.Shape()[0]
|
||||
mean := 0.0
|
||||
for row := range rows {
|
||||
mean += samples.FloatAt(row * 2)
|
||||
}
|
||||
mean /= float64(rows)
|
||||
variance := 0.0
|
||||
for row := range rows {
|
||||
d := samples.FloatAt(row*2) - mean
|
||||
variance += d * d / float64(rows)
|
||||
}
|
||||
fmt.Printf("%v: mean %.3f, variance %.3f\n", samples.Shape(), mean, variance)
|
||||
// Output: [2000 2]: mean -0.012, variance 1.013
|
||||
}
|
||||
@@ -0,0 +1,187 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Property fuzz targets over index-moving operations. `go test` runs
|
||||
// the seed corpus on every commit; longer campaigns run under
|
||||
// -fuzz=Fuzz<Name> when a shape-handling change lands.
|
||||
|
||||
// FuzzTransposeAxesRoundTrip drives random permutations through the
|
||||
// axis move and its inverse: whatever valid permutation arrives, the
|
||||
// double transpose must restore the exact element order.
|
||||
func FuzzTransposeAxesRoundTrip(f *testing.F) {
|
||||
f.Add([]byte{0, 1}, 6)
|
||||
f.Add([]byte{1, 0}, 6)
|
||||
f.Add([]byte{2, 0, 1}, 8)
|
||||
|
||||
f.Fuzz(func(t *testing.T, permBytes []byte, total int) {
|
||||
if total <= 0 || total > 4096 {
|
||||
t.Skip()
|
||||
}
|
||||
rank := len(permBytes)
|
||||
switch rank {
|
||||
case 2:
|
||||
total -= total % 2
|
||||
case 3:
|
||||
total -= total % 4
|
||||
default:
|
||||
t.Skip()
|
||||
}
|
||||
if total == 0 {
|
||||
t.Skip()
|
||||
}
|
||||
vals := make([]float64, total)
|
||||
for i := range vals {
|
||||
vals[i] = float64(i)
|
||||
}
|
||||
var shape []int
|
||||
if rank == 2 {
|
||||
shape = []int{total / 2, 2}
|
||||
} else {
|
||||
shape = []int{total / 4, 2, 2}
|
||||
}
|
||||
a, _ := core.FromFloats(vals, shape...)
|
||||
xt := FromArray(a, false)
|
||||
|
||||
dims := make([]int, rank)
|
||||
for i, pb := range permBytes {
|
||||
dims[i] = int(pb) % rank
|
||||
}
|
||||
moved, err := xt.TransposeAxes(dims...)
|
||||
if err != nil {
|
||||
return // duplicate axes rejected by validation, fine
|
||||
}
|
||||
back, err := moved.TransposeAxes(inversePerm(dims)...)
|
||||
if err != nil {
|
||||
t.Fatalf("inverse of %v failed: %v", dims, err)
|
||||
}
|
||||
for i := range vals {
|
||||
if back.Data().FloatAt(i) != vals[i] {
|
||||
t.Fatalf("round trip lost element %d", i)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// FuzzOneHotContracts checks both sides of the encoder contract for
|
||||
// arbitrary code sets: in-range codes yield exactly one hot cell per
|
||||
// row, any out-of-range code is a loud error. One input byte splits
|
||||
// into a high bit forcing negativity plus a low-bit class selector.
|
||||
func FuzzOneHotContracts(f *testing.F) {
|
||||
f.Add([]byte{0, 1, 2}, uint8(3))
|
||||
f.Add([]byte{5}, uint8(8))
|
||||
f.Add([]byte{200, 201}, uint8(3))
|
||||
|
||||
f.Fuzz(func(t *testing.T, raw []byte, classByte uint8) {
|
||||
classes := int(classByte)%9 + 1
|
||||
codes := make([]int64, len(raw))
|
||||
valid := true
|
||||
for i, b := range raw {
|
||||
c := int64(b)
|
||||
if b >= 128 { // force some negative probes
|
||||
c = -int64(b - 127)
|
||||
} else {
|
||||
c %= int64(classes)
|
||||
}
|
||||
if c < 0 || c >= int64(classes) {
|
||||
valid = false
|
||||
}
|
||||
codes[i] = c
|
||||
}
|
||||
|
||||
arr, _ := core.FromInts(codes, len(codes))
|
||||
hot, err := core.OneHot(arr, classes)
|
||||
if !valid {
|
||||
if err == nil {
|
||||
t.Fatalf("invalid codes accepted for %d classes", classes)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("valid codes rejected: %v", err)
|
||||
}
|
||||
for i := range len(codes) {
|
||||
sum := 0.0
|
||||
for j := range classes {
|
||||
sum += float64(hot.FloatAt(i*classes + j))
|
||||
}
|
||||
if sum != 1 {
|
||||
t.Fatalf("row %d sums to %v, want one hot cell", i, sum)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// FuzzConcatSplitGradientConserves mass: splitting a concatenated
|
||||
// output's gradient must hand every element back to its own side with
|
||||
// coefficient exactly one, for whatever layout the corpus invents.
|
||||
func FuzzConcatSplitGradientConserves(f *testing.F) {
|
||||
f.Add([]byte{1, 2, 3, 4}, uint8(2))
|
||||
f.Add([]byte{9, 7, 5}, uint8(1))
|
||||
f.Add([]byte{10, 20, 30, 40, 50, 60, 70, 80}, uint8(0))
|
||||
f.Add([]byte{11, 21, 31, 41, 51, 61, 71, 81, 91, 101}, uint8(4))
|
||||
|
||||
f.Fuzz(func(t *testing.T, raw []byte, rowsByte uint8) {
|
||||
rows := int(rowsByte)%7 + 1 // left matrix rows, 1..7
|
||||
leftLen := rows * 2
|
||||
if len(raw) <= leftLen {
|
||||
t.Skip()
|
||||
}
|
||||
rightRows := (len(raw) - leftLen) / 2
|
||||
|
||||
leftVals := make([]float64, leftLen)
|
||||
for i := range leftVals {
|
||||
leftVals[i] = float64(raw[i])
|
||||
}
|
||||
rightVals := make([]float64, rightRows*2)
|
||||
for i := range rightVals {
|
||||
rightVals[i] = float64(raw[leftLen+i])
|
||||
}
|
||||
|
||||
a, err := core.FromFloats(leftVals, rows, 2)
|
||||
if err != nil {
|
||||
t.Skip()
|
||||
}
|
||||
b, err := core.FromFloats(rightVals, rightRows, 2)
|
||||
if err != nil {
|
||||
t.Skip()
|
||||
}
|
||||
|
||||
at := FromArray(a, true)
|
||||
bt := FromArray(b, true)
|
||||
joint, cerr := at.Concat(bt, 0)
|
||||
if cerr != nil {
|
||||
t.Fatal(cerr)
|
||||
}
|
||||
loss, serr := joint.Sum()
|
||||
if serr != nil {
|
||||
t.Fatal(serr)
|
||||
}
|
||||
if berr := loss.Backward(); berr != nil {
|
||||
t.Fatal(berr)
|
||||
}
|
||||
|
||||
ga, gb := at.Grad(), bt.Grad()
|
||||
if ga.Len() != a.Len() || gb.Len() != b.Len() {
|
||||
t.Fatalf("gradient spans drifted: %d+%d vs %d+%d",
|
||||
ga.Len(), gb.Len(), a.Len(), b.Len())
|
||||
}
|
||||
for i := range ga.Len() {
|
||||
if ga.FloatAt(i) != 1 {
|
||||
t.Fatalf("left span slot %d = %v", i, ga.FloatAt(i))
|
||||
}
|
||||
}
|
||||
for i := range gb.Len() {
|
||||
if gb.FloatAt(i) != 1 {
|
||||
t.Fatalf("right span slot %d = %v", i, gb.FloatAt(i))
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,230 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
ode "sourcedock.dev/petrbalvin/tensor/integrate"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Regression tests for gradient hygiene: the Add backward
|
||||
// handed both inputs the same gradient instance, AdjointODE polluted
|
||||
// trainable leaves closed over by f but outside params, and
|
||||
// MinimiseNewtonCG panicked on nil inputs where the rest of the
|
||||
// package returns errors.
|
||||
|
||||
// TestAddBackwardIndependentGradients pins the aliasing fix: the two
|
||||
// inputs of an Add receive two independent gradient buffers with
|
||||
// identical values, so a write through one leaf's gradient cannot
|
||||
// corrupt the other's.
|
||||
func TestAddBackwardIndependentGradients(t *testing.T) {
|
||||
x, err := FromFloat64s([]float64{2, 3}, true, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
y, err := FromFloat64s([]float64{4, 5}, true, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
z, err := x.Add(y)
|
||||
if err != nil {
|
||||
t.Fatalf("Add: %v", err)
|
||||
}
|
||||
loss, err := z.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
if x.Grad() == y.Grad() {
|
||||
t.Fatal("Add handed both leaves the same gradient instance")
|
||||
}
|
||||
// d(x+y)/dx = 1 and d(x+y)/dy = 1, element for element.
|
||||
for i := range 2 {
|
||||
if x.Grad().FloatAt(i) != 1 {
|
||||
t.Fatalf("x.grad[%d] = %v, want 1", i, x.Grad().FloatAt(i))
|
||||
}
|
||||
if y.Grad().FloatAt(i) != 1 {
|
||||
t.Fatalf("y.grad[%d] = %v, want 1", i, y.Grad().FloatAt(i))
|
||||
}
|
||||
}
|
||||
// A write through one leaf's buffer must leave the other's intact.
|
||||
x.Grad().SetFloatAt(0, 999)
|
||||
if y.Grad().FloatAt(0) != 1 {
|
||||
t.Fatalf("y.grad[0] = %v after a write through x's gradient, want 1",
|
||||
y.Grad().FloatAt(0))
|
||||
}
|
||||
}
|
||||
|
||||
// TestAddSameLeafAccumulatesTwice pins Add(x, x): the same leaf as both
|
||||
// inputs accumulates both contributions into one gradient of 2.
|
||||
func TestAddSameLeafAccumulatesTwice(t *testing.T) {
|
||||
x, err := FromFloat64s([]float64{0.5, -1.25}, true, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
z, err := x.Add(x)
|
||||
if err != nil {
|
||||
t.Fatalf("Add: %v", err)
|
||||
}
|
||||
loss, err := z.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("Sum: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
// d(x+x)/dx = 2, and the accumulated buffer is one array.
|
||||
for i := range 2 {
|
||||
if x.Grad().FloatAt(i) != 2 {
|
||||
t.Fatalf("x.grad[%d] = %v, want 2", i, x.Grad().FloatAt(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjointODELeavesHiddenLeavesClean pins the vjp fix: a trainable
|
||||
// leaf closed over by f but not listed in params receives no gradient,
|
||||
// because the Jacobian-vector products come from a pass that commits
|
||||
// nothing.
|
||||
func TestAdjointODELeavesHiddenLeavesClean(t *testing.T) {
|
||||
theta, err := FromFloat64s([]float64{0.7}, true, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
hidden, err := FromFloat64s([]float64{1.5}, true, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
f := func(t float64, y *Tensor) (*Tensor, error) {
|
||||
rate, err := theta.Neg()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
hm, err := y.Mul(hidden)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return hm.Mul(rate)
|
||||
}
|
||||
y0, _ := core.FromFloats([]float64{1}, 1)
|
||||
seed, _ := core.FromFloats([]float64{1}, 1)
|
||||
_, blocks, err := AdjointODE(f, []*Tensor{theta}, 0, 1, y0, seed,
|
||||
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||||
if err != nil {
|
||||
t.Fatalf("AdjointODE: %v", err)
|
||||
}
|
||||
// dL/dθ = −1.5·e^{−1.05}, the central-difference answer the exact
|
||||
// dynamics give.
|
||||
want := -1.5 * math.Exp(-1.05)
|
||||
if math.Abs(blocks[0].FloatAt(0)-want) > 1e-6 {
|
||||
t.Fatalf("dL/dθ = %.14g, want %.14g", blocks[0].FloatAt(0), want)
|
||||
}
|
||||
if hidden.Grad() != nil {
|
||||
t.Fatalf("the hidden leaf's gradient = %v, want nil", hidden.Grad())
|
||||
}
|
||||
if theta.Grad() != nil {
|
||||
t.Fatalf("the parameter's gradient = %v, want nil", theta.Grad())
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjointODEThetaAgainstCentralDifferences checks the returned
|
||||
// parameter sensitivity of the same hidden-leaf system against central
|
||||
// differences of the forward solve, so the cleanup left the θ gradient
|
||||
// every bit as accurate as it was.
|
||||
func TestAdjointODEThetaAgainstCentralDifferences(t *testing.T) {
|
||||
theta, _ := FromFloat64s([]float64{0.7}, true, 1)
|
||||
hidden, _ := FromFloat64s([]float64{1.5}, true, 1)
|
||||
f := func(t float64, y *Tensor) (*Tensor, error) {
|
||||
rate, err := theta.Neg()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
hm, err := y.Mul(hidden)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return hm.Mul(rate)
|
||||
}
|
||||
y0, _ := core.FromFloats([]float64{1}, 1)
|
||||
seed, _ := core.FromFloats([]float64{1}, 1)
|
||||
_, blocks, err := AdjointODE(f, []*Tensor{theta}, 0, 1, y0, seed,
|
||||
ode.ODEOptions{RelTol: 1e-9, AbsTol: 1e-12})
|
||||
if err != nil {
|
||||
t.Fatalf("AdjointODE: %v", err)
|
||||
}
|
||||
// Central differences on the loss y(1), with the parameter leaf's
|
||||
// data swapped out for the perturbed values.
|
||||
forwardLoss := func() float64 {
|
||||
end, ferr := ode.IntegrateODE(func(t float64, ya *core.Array) (*core.Array, error) {
|
||||
out, oerr := f(t, FromArray(ya, false))
|
||||
if oerr != nil {
|
||||
return nil, oerr
|
||||
}
|
||||
return out.Data(), nil
|
||||
}, 0, 1, y0, ode.ODEOptions{RelTol: 1e-11, AbsTol: 1e-15})
|
||||
if ferr != nil {
|
||||
t.Fatalf("forward solve: %v", ferr)
|
||||
}
|
||||
return end.FloatAt(0)
|
||||
}
|
||||
const eps = 1e-6
|
||||
orig := theta.Data().FloatAt(0)
|
||||
up, _ := core.FromFloats([]float64{orig + eps}, 1)
|
||||
theta.ReplaceWith(up)
|
||||
hi := forwardLoss()
|
||||
dn, _ := core.FromFloats([]float64{orig - eps}, 1)
|
||||
theta.ReplaceWith(dn)
|
||||
lo := forwardLoss()
|
||||
back, _ := core.FromFloats([]float64{orig}, 1)
|
||||
theta.ReplaceWith(back)
|
||||
fd := (hi - lo) / (2 * eps)
|
||||
if math.Abs(blocks[0].FloatAt(0)-fd) > 1e-5*math.Max(1, math.Abs(fd)) {
|
||||
t.Fatalf("dL/dθ: adjoint %.10g, central difference %.10g",
|
||||
blocks[0].FloatAt(0), fd)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMinimiseNewtonCGNilInputs pins the validation contract: a nil
|
||||
// objective and a nil starting point are errors naming the argument,
|
||||
// in step with SampleHMC, never panics.
|
||||
func TestMinimiseNewtonCGNilInputs(t *testing.T) {
|
||||
f := func(z *Tensor) (*Tensor, error) { return z.Sum() }
|
||||
if _, _, err := MinimiseNewtonCG(nil, nil, NewtonCGOptions{}); err == nil {
|
||||
t.Fatal("expected an error for a nil objective")
|
||||
} else if !strings.Contains(err.Error(), "must not be nil") {
|
||||
t.Fatalf("error = %v, want a must-not-be-nil refusal", err)
|
||||
}
|
||||
x0, err := core.FromFloats([]float64{1}, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
if _, _, err := MinimiseNewtonCG(nil, x0, NewtonCGOptions{}); err == nil {
|
||||
t.Fatal("expected an error for a nil objective")
|
||||
} else if !strings.Contains(err.Error(), "f must not be nil") {
|
||||
t.Fatalf("error = %v, want a refusal naming f", err)
|
||||
}
|
||||
if _, _, err := MinimiseNewtonCG(f, nil, NewtonCGOptions{}); err == nil {
|
||||
t.Fatal("expected an error for a nil starting point")
|
||||
} else if !strings.Contains(err.Error(), "starting point must not be nil") {
|
||||
t.Fatalf("error = %v, want a refusal naming the starting point", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestODETraceSingleNodeRefused pins the interpolation guard: a trace
|
||||
// with fewer than two recorded nodes has no interval to interpolate
|
||||
// over, so the accessor refuses instead of indexing out of range.
|
||||
func TestODETraceSingleNodeRefused(t *testing.T) {
|
||||
tr := &odeTrace{times: []float64{1}, states: [][]float64{{2, 3}},
|
||||
slopes: [][]float64{{0, 0}}, dim: 2}
|
||||
if _, err := tr.at(1); err == nil {
|
||||
t.Fatal("expected an error for a one-node trace")
|
||||
} else if !strings.Contains(err.Error(), "at least two") {
|
||||
t.Fatalf("error = %v, want a refusal naming the node count", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// numericGrad estimates the gradient of a scalar function at a by
|
||||
// central differences, the reference the analytic backward is checked
|
||||
// against.
|
||||
func numericGrad(f func(a *core.Array) float64, a *core.Array) []float64 {
|
||||
n := a.Len()
|
||||
out := make([]float64, n)
|
||||
for i := range n {
|
||||
hi := 1e-6
|
||||
up := cloneFlat(a)
|
||||
down := cloneFlat(a)
|
||||
up[i] += hi
|
||||
down[i] -= hi
|
||||
au, _ := core.FromFloats(up, a.Shape()...)
|
||||
ad, _ := core.FromFloats(down, a.Shape()...)
|
||||
out[i] = (f(au) - f(ad)) / (2 * hi)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func cloneFlat(a *core.Array) []float64 {
|
||||
out := make([]float64, a.Len())
|
||||
for i := range a.Len() {
|
||||
out[i] = a.FloatAt(i)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func maxAbsDiff(a, b []float64) float64 {
|
||||
m := 0.0
|
||||
for i := range a {
|
||||
if d := math.Abs(a[i] - b[i]); d > m {
|
||||
m = d
|
||||
}
|
||||
}
|
||||
return m
|
||||
}
|
||||
+178
@@ -0,0 +1,178 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Second-order differentiation, forward-over-reverse: the
|
||||
// inner derivative is the exact analytic gradient Backward produces,
|
||||
// and only the outer derivative runs by central differences over one
|
||||
// coordinate at a time. The result carries the accuracy of the exact
|
||||
// first derivative with the O(h²) truncation of the outer stencil, the
|
||||
// same trade a hand-written finite-difference Hessian makes but with
|
||||
// none of the first-order error.
|
||||
|
||||
// HessianOptions tunes Hessian and HessianVectorProduct. Step is the
|
||||
// absolute coordinate perturbation (≤ 0 picks sqrt(eps)·max(1, |x_i|)
|
||||
// per coordinate, the stencil that balances truncation against
|
||||
// cancellation at double precision).
|
||||
type HessianOptions struct {
|
||||
Step float64
|
||||
}
|
||||
|
||||
// Hessian returns the Hessian matrix of a scalar function f at x, an
|
||||
// (n, n) float64 array for an n-element x. f receives a tensor that
|
||||
// requires grad and must return a single-element real tensor; complex
|
||||
// outputs are rejected like Backward does. The cost is 2n gradient
|
||||
// evaluations, the price of a dense second derivative by any method
|
||||
// that does not exploit structure; for large n prefer
|
||||
// HessianVectorProduct. The evaluations differentiate the graph
|
||||
// without committing anything, so the accumulated gradients of the
|
||||
// tensors f closes over are left exactly as they were, on the success
|
||||
// and the error path alike.
|
||||
func Hessian(f func(*Tensor) (*Tensor, error), x *Tensor, opts HessianOptions) (*core.Array, error) {
|
||||
if x.Data().Dtype() == core.Complex {
|
||||
return nil, base.Errf("Hessian: complex points are not supported")
|
||||
}
|
||||
n := x.Data().Len()
|
||||
if n == 0 {
|
||||
return nil, base.Errf("Hessian: the point must not be empty")
|
||||
}
|
||||
p := flatFloats(x.Data())
|
||||
out := zeros(core.Float, []int{n, n})
|
||||
h := opts.Step
|
||||
for j := range n {
|
||||
step := h
|
||||
if step <= 0 {
|
||||
step = math.Sqrt(2.220446049250313e-16) * math.Max(1, math.Abs(p[j]))
|
||||
}
|
||||
gp, err := hessianColumn(f, x, p, j, step)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
gm, err := hessianColumn(f, x, p, j, -step)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
inv2h := 1 / (2 * step)
|
||||
for i := range n {
|
||||
out.SetFloatAt(i*n+j, (gp[i]-gm[i])*inv2h)
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// hessianColumn evaluates the analytic gradient of f at the point
|
||||
// perturbed by step along coordinate j, flattened.
|
||||
func hessianColumn(f func(*Tensor) (*Tensor, error), x *Tensor, p []float64, j int, step float64) ([]float64, error) {
|
||||
probe := append([]float64(nil), p...)
|
||||
probe[j] += step
|
||||
pa, err := core.FromFloats(probe, x.Data().Shape()...)
|
||||
if err != nil {
|
||||
return nil, base.Errf("Hessian: %w", err)
|
||||
}
|
||||
xt := FromArray(pa, true)
|
||||
y, err := f(xt)
|
||||
if err != nil {
|
||||
return nil, base.Errf("Hessian: %w", err)
|
||||
}
|
||||
if y.Data().Len() != 1 {
|
||||
return nil, base.Errf("Hessian: f must return a scalar, got %d elements", y.Data().Len())
|
||||
}
|
||||
grads, err := y.reverseGrads()
|
||||
if err != nil {
|
||||
return nil, base.Errf("Hessian: %w", err)
|
||||
}
|
||||
g := grads[xt]
|
||||
if g == nil {
|
||||
return nil, base.Errf("Hessian: the objective does not depend on x, so no gradient exists")
|
||||
}
|
||||
return flatFloats(g), nil
|
||||
}
|
||||
|
||||
// HessianVectorProduct returns H·v, the Hessian of the scalar f at x
|
||||
// contracted with the direction v, by a central difference along the
|
||||
// direction itself, with the step scaled so it never depends on v's
|
||||
// magnitude. Two gradient evaluations
|
||||
// answer for any n, which is what makes Newton-CG tractable where a
|
||||
// dense Hessian is not. As in Hessian, the evaluations leave the
|
||||
// accumulated gradients of every tensor f closes over untouched.
|
||||
func HessianVectorProduct(f func(*Tensor) (*Tensor, error), x, v *Tensor, opts HessianOptions) (*core.Array, error) {
|
||||
if x.Data().Dtype() == core.Complex {
|
||||
return nil, base.Errf("HessianVectorProduct: complex points are not supported")
|
||||
}
|
||||
if v.Data().Dtype() == core.Complex {
|
||||
return nil, base.Errf("HessianVectorProduct: complex directions are not supported")
|
||||
}
|
||||
n := x.Data().Len()
|
||||
if v.Data().Len() != n {
|
||||
return nil, base.Errf("HessianVectorProduct: direction has %d elements for %d variables",
|
||||
v.Data().Len(), n)
|
||||
}
|
||||
vn := 0.0
|
||||
for _, v := range flatFloats(v.Data()) {
|
||||
vn += v * v
|
||||
}
|
||||
vn = math.Sqrt(vn)
|
||||
if vn == 0 {
|
||||
// H·0 = 0 in the shape of the point, the same shape the
|
||||
// quotient below returns: a flat vector here would change the
|
||||
// result's shape with the direction's norm.
|
||||
return zeros(core.Float, x.Data().Shape()), nil
|
||||
}
|
||||
h := opts.Step
|
||||
if h <= 0 {
|
||||
h = 1e-5
|
||||
}
|
||||
p := flatFloats(x.Data())
|
||||
vf := flatFloats(v.Data())
|
||||
eval := func(sign float64) ([]float64, error) {
|
||||
probe := make([]float64, n)
|
||||
for i := range n {
|
||||
probe[i] = p[i] + sign*h*vf[i]/vn
|
||||
}
|
||||
pa, err := core.FromFloats(probe, x.Data().Shape()...)
|
||||
if err != nil {
|
||||
return nil, base.Errf("HessianVectorProduct: %w", err)
|
||||
}
|
||||
xt := FromArray(pa, true)
|
||||
y, err := f(xt)
|
||||
if err != nil {
|
||||
return nil, base.Errf("HessianVectorProduct: %w", err)
|
||||
}
|
||||
if y.Data().Len() != 1 {
|
||||
return nil, base.Errf("HessianVectorProduct: f must return a scalar, got %d elements", y.Data().Len())
|
||||
}
|
||||
grads, err := y.reverseGrads()
|
||||
if err != nil {
|
||||
return nil, base.Errf("HessianVectorProduct: %w", err)
|
||||
}
|
||||
g := grads[xt]
|
||||
if g == nil {
|
||||
return nil, base.Errf("HessianVectorProduct: the objective does not depend on x, so no gradient exists")
|
||||
}
|
||||
return flatFloats(g), nil
|
||||
}
|
||||
gp, err := eval(1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
gm, err := eval(-1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// The step advanced h·v/|v| along v, so the quotient is (H·v/|v|)
|
||||
// and carries the |v| factor back in.
|
||||
out := zeros(core.Float, x.Data().Shape())
|
||||
inv2h := vn / (2 * h)
|
||||
for i := range n {
|
||||
out.SetFloatAt(i, (gp[i]-gm[i])*inv2h)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// Regression tests: an objective whose graph
|
||||
// never reaches x left xt.Grad() nil, and the second-order helpers
|
||||
// dereferenced it.
|
||||
|
||||
// TestHessianDisconnectedObjective pins the error: an objective that
|
||||
// ignores its argument has no gradient to differentiate.
|
||||
func TestHessianDisconnectedObjective(t *testing.T) {
|
||||
x, err := FromFloat64s([]float64{1, 2}, true, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
c, err := FromFloat64s([]float64{3}, true, 1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
objective := func(*Tensor) (*Tensor, error) { return c.Mul(c) }
|
||||
|
||||
if _, err := Hessian(objective, x, HessianOptions{}); err == nil {
|
||||
t.Fatal("expected an error when the objective does not depend on x")
|
||||
} else if !strings.Contains(err.Error(), "does not depend") {
|
||||
t.Fatalf("error = %v, want a disconnected-graph refusal", err)
|
||||
}
|
||||
v, err := FromFloat64s([]float64{1, 0}, true, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := HessianVectorProduct(objective, x, v, HessianOptions{}); err == nil {
|
||||
t.Fatal("expected an error when the objective does not depend on x")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,193 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestHessianQuadratic pins the exact case: for f(x) = ½xᵀAx + bᵀx the
|
||||
// Hessian is A whatever the point.
|
||||
func TestHessianQuadratic(t *testing.T) {
|
||||
a := []float64{4, 1, 1, 3}
|
||||
b := []float64{-1, 2}
|
||||
x0 := []float64{0.5, -1.25}
|
||||
xt, err := FromFloat64s(x0, false, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
f := func(z *Tensor) (*Tensor, error) {
|
||||
az, err := FromFloat64s(a, false, 2, 2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bz, err := FromFloat64s(b, false, 2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
halfAz, err := az.Scale(0.5)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
azx, err := halfAz.MatMul(z)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sum1, err := azx.Add(bz)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// ½xᵀAx + bᵀx = ((½A)x + b)·x
|
||||
return sum1.Mul(z)
|
||||
}
|
||||
// ((½A)x + b)·x is elementwise; the scalar loss needs the sum.
|
||||
fScalar := func(z *Tensor) (*Tensor, error) {
|
||||
p, err := f(z)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return p.Sum()
|
||||
}
|
||||
h, err := Hessian(fScalar, xt, HessianOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("Hessian: %v", err)
|
||||
}
|
||||
for i := range 2 {
|
||||
for j := range 2 {
|
||||
if math.Abs(h.FloatAt(i*2+j)-a[i*2+j]) > 1e-6 {
|
||||
t.Fatalf("H[%d][%d] = %g, want %g", i, j, h.FloatAt(i*2+j), a[i*2+j])
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestHessianRosenbrock pins a nonquadratic landscape against the
|
||||
// analytic Hessian of the 2-D Rosenbrock function.
|
||||
func TestHessianRosenbrock(t *testing.T) {
|
||||
x0 := []float64{-0.5, 1.25}
|
||||
xt, err := FromFloat64s(x0, false, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
f := func(z *Tensor) (*Tensor, error) {
|
||||
els := []int{0, 1}
|
||||
x0t, err := z.Slice(0, els[0], els[0]+1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
x1t, err := z.Slice(0, els[1], els[1]+1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
x0sq, err := x0t.Pow(2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
diff, err := x1t.Sub(x0sq)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
term1v, err := diff.Pow(2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
one, err := FromFloat64s([]float64{1}, false, 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
x0m1, err := x0t.Sub(one)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
term2v, err := x0m1.Pow(2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
term2s, err := term2v.Scale(100)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
total, err := term1v.Add(term2s)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return total.Sum()
|
||||
}
|
||||
h, err := Hessian(f, xt, HessianOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("Hessian: %v", err)
|
||||
}
|
||||
x, y := x0[0], x0[1]
|
||||
// Analytic Hessian of f = (y − x²)² + 100(x − 1)².
|
||||
h00 := 12*x*x - 4*y + 200
|
||||
h01 := -4 * x
|
||||
h11 := 2.0
|
||||
want := [][]float64{{h00, h01}, {h01, h11}}
|
||||
for i := range 2 {
|
||||
for j := range 2 {
|
||||
if math.Abs(h.FloatAt(i*2+j)-want[i][j]) > 1e-4*math.Max(1, math.Abs(want[i][j])) {
|
||||
t.Fatalf("H[%d][%d] = %g, want %g", i, j, h.FloatAt(i*2+j), want[i][j])
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestHessianVectorProduct pins H·v against the dense Hessian.
|
||||
func TestHessianVectorProduct(t *testing.T) {
|
||||
a := []float64{4, 1, 1, 3}
|
||||
x0 := []float64{0.5, -1.25}
|
||||
xt, err := FromFloat64s(x0, false, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
f := func(z *Tensor) (*Tensor, error) {
|
||||
az, err := FromFloat64s(a, false, 2, 2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
halfAz, err := az.Scale(0.5)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
azx, err := halfAz.MatMul(z)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p, err := azx.Mul(z)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return p.Sum()
|
||||
}
|
||||
vArr, err := core.FromFloats([]float64{2, -1}, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
v := FromArray(vArr, false)
|
||||
hv, err := HessianVectorProduct(f, xt, v, HessianOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("HessianVectorProduct: %v", err)
|
||||
}
|
||||
// A·v exactly.
|
||||
want := []float64{4*2 + 1*(-1), 1*2 + 3*(-1)}
|
||||
for i := range 2 {
|
||||
if math.Abs(hv.FloatAt(i)-want[i]) > 1e-5 {
|
||||
t.Fatalf("Hv[%d] = %g, want %g", i, hv.FloatAt(i), want[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestHessianRejectsVectorOutput pins the scalar contract.
|
||||
func TestHessianRejectsVectorOutput(t *testing.T) {
|
||||
xt, err := FromFloat64s([]float64{1, 2}, false, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
f := func(z *Tensor) (*Tensor, error) { return z, nil }
|
||||
if _, err := Hessian(f, xt, HessianOptions{}); err == nil {
|
||||
t.Fatal("Hessian accepted a vector output")
|
||||
}
|
||||
}
|
||||
+194
@@ -0,0 +1,194 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Hamiltonian Monte Carlo on the autograd surface. The target is any
|
||||
// differentiable log density: a trajectory simulates the Hamiltonian
|
||||
// dynamics of a unit-mass particle on that landscape, with the
|
||||
// momentum refreshed from a standard Gaussian each round and the
|
||||
// leapfrog integrator driven by gradients Backward computes, so the
|
||||
// user supplies the density and the chain does the calculus. The
|
||||
// Metropolis correction on the trajectory's energy change makes the
|
||||
// stationary distribution exact despite the integration error.
|
||||
|
||||
// HMCOptions tunes SampleHMC. Step is the leapfrog step size and
|
||||
// Steps the number of leapfrog steps per trajectory, both
|
||||
// problem-dependent with no sensible default; BurnIn trajectories are
|
||||
// discarded before every Thin-th trajectory contributes one sample,
|
||||
// until Samples have been collected, and a Thin of zero or less
|
||||
// quietly normalises to one. Seed feeds the package's own
|
||||
// xoshiro generator, so a run is bit-reproducible.
|
||||
type HMCOptions struct {
|
||||
Step float64
|
||||
Steps int
|
||||
BurnIn int
|
||||
Thin int
|
||||
Samples int
|
||||
Seed int64
|
||||
}
|
||||
|
||||
// SampleHMC draws Samples states from the unnormalised density whose
|
||||
// logarithm logDensity computes, starting at the vector q0 and
|
||||
// returning the kept states as a (Samples × dim) array; rejected
|
||||
// trajectories repeat the current state, as Markov chain sampling
|
||||
// does. logDensity receives a leaf tensor requiring grad and must
|
||||
// return a scalar tensor connected to it; an error it raises at q0 is
|
||||
// fatal, while one raised inside a proposal marks the state as
|
||||
// outside the support and rejects the trajectory, which is how
|
||||
// constrained densities keep the chain away from forbidden regions.
|
||||
// A nil density, a non-vector start, a non-positive step, step count
|
||||
// or sample count, or a density that does not yield a gradient are
|
||||
// errors. The gradient evaluations differentiate the graph without
|
||||
// committing anything, so the accumulated gradients of the tensors
|
||||
// logDensity closes over are left exactly as they were, on the success
|
||||
// and the error path alike.
|
||||
func SampleHMC(logDensity func(q *Tensor) (*Tensor, error),
|
||||
q0 *core.Array, opts HMCOptions) (*core.Array, error) {
|
||||
const name = "SampleHMC"
|
||||
if logDensity == nil {
|
||||
return nil, errf("%s: logDensity must not be nil", name)
|
||||
}
|
||||
if q0 == nil {
|
||||
return nil, errf("%s: the state must not be nil", name)
|
||||
}
|
||||
if q0.NDim() != 1 || q0.Len() == 0 {
|
||||
return nil, errf("%s: the state must be a non-empty vector, got shape %s",
|
||||
name, prettyShape(q0.Shape()))
|
||||
}
|
||||
if q0.Dtype() == core.Complex {
|
||||
return nil, errf("%s: complex states are not supported", name)
|
||||
}
|
||||
if opts.Step <= 0 {
|
||||
return nil, errf("%s: Step must be positive, got %g", name, opts.Step)
|
||||
}
|
||||
if opts.Steps <= 0 {
|
||||
return nil, errf("%s: Steps must be at least 1, got %d", name, opts.Steps)
|
||||
}
|
||||
if opts.Samples <= 0 {
|
||||
return nil, errf("%s: Samples must be at least 1, got %d", name, opts.Samples)
|
||||
}
|
||||
if opts.BurnIn < 0 {
|
||||
return nil, errf("%s: BurnIn must not be negative, got %d", name, opts.BurnIn)
|
||||
}
|
||||
if opts.Thin <= 0 {
|
||||
opts.Thin = 1
|
||||
}
|
||||
dim := q0.Len()
|
||||
rng := core.NewGenerator(opts.Seed)
|
||||
|
||||
// eval computes log π at q together with ∇log π(q): a fresh leaf
|
||||
// per call, one reverse pass that commits nothing, plain floats out.
|
||||
// A leapfrog trajectory calls this once per step, so a pass that
|
||||
// committed would add one contribution to every tensor logDensity
|
||||
// closes over per step; the reverse pass of reverseGrads leaves the
|
||||
// caller's accumulated gradients untouched instead.
|
||||
eval := func(q []float64) (float64, []float64, error) {
|
||||
data, err := core.FromFloats(q, dim)
|
||||
if err != nil {
|
||||
return 0, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
leaf := FromArray(data, true)
|
||||
out, err := logDensity(leaf)
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
if out.Data().Len() != 1 {
|
||||
return 0, nil, errf("%s: logDensity returned shape %s, want a scalar",
|
||||
name, prettyShape(out.Data().Shape()))
|
||||
}
|
||||
grads, err := out.reverseGrads()
|
||||
if err != nil {
|
||||
return 0, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
g := grads[leaf]
|
||||
if g == nil || g.Len() != dim {
|
||||
return 0, nil, errf("%s: logDensity did not yield a gradient of length %d", name, dim)
|
||||
}
|
||||
return out.Data().FloatAt(0), flatFloats(g), nil
|
||||
}
|
||||
|
||||
q := flatFloats(q0)
|
||||
logPi, gradient, err := eval(q)
|
||||
if err != nil {
|
||||
return nil, errf("%s: %w", name, err)
|
||||
}
|
||||
p := make([]float64, dim)
|
||||
qNew := make([]float64, dim)
|
||||
pNew := make([]float64, dim)
|
||||
values := make([]float64, 0, opts.Samples*dim)
|
||||
trajectories := opts.BurnIn + opts.Samples*opts.Thin
|
||||
for traj := 1; traj <= trajectories; traj++ {
|
||||
// Fresh momentum from the standard Gaussian; the kinetic
|
||||
// energy is p·p/2 for unit mass.
|
||||
momentum, err := core.Normal(rng, dim, 0, 1)
|
||||
if err != nil {
|
||||
return nil, errf("%s: %w", name, err)
|
||||
}
|
||||
kinetic0 := 0.0
|
||||
momentumF := flatFloats(momentum)
|
||||
for i := range p {
|
||||
p[i] = momentumF[i]
|
||||
kinetic0 += p[i] * p[i] / 2
|
||||
}
|
||||
copy(qNew, q)
|
||||
copy(pNew, p)
|
||||
// Leapfrog: half kick, drift, full kick per step, one gradient
|
||||
// evaluation each, the last one doubling as the proposal's
|
||||
// log density.
|
||||
diverged := false
|
||||
logPiNew := math.Inf(-1)
|
||||
g := gradient
|
||||
for range opts.Steps {
|
||||
for i := range pNew {
|
||||
pNew[i] += opts.Step / 2 * g[i]
|
||||
}
|
||||
for i := range qNew {
|
||||
qNew[i] += opts.Step * pNew[i]
|
||||
}
|
||||
lp, gNew, eerr := eval(qNew)
|
||||
if eerr != nil {
|
||||
diverged = true
|
||||
break
|
||||
}
|
||||
for i := range pNew {
|
||||
pNew[i] += opts.Step / 2 * gNew[i]
|
||||
}
|
||||
g = gNew
|
||||
logPiNew = lp
|
||||
}
|
||||
if !diverged {
|
||||
kinetic1 := 0.0
|
||||
for i := range pNew {
|
||||
kinetic1 += pNew[i] * pNew[i] / 2
|
||||
}
|
||||
// Metropolis on the energy change; an undefined change
|
||||
// (a NaN crept into the landscape) rejects.
|
||||
logAccept := (-logPi + kinetic0) - (-logPiNew + kinetic1)
|
||||
uniform, uerr := core.Floats(rng, 1)
|
||||
if uerr != nil {
|
||||
return nil, errf("%s: %w", name, uerr)
|
||||
}
|
||||
if !math.IsNaN(logAccept) && math.Log(uniform.FloatAt(0)) < logAccept {
|
||||
copy(q, qNew)
|
||||
logPi = logPiNew
|
||||
gradient = g
|
||||
}
|
||||
}
|
||||
if traj <= opts.BurnIn || (traj-opts.BurnIn)%opts.Thin != 0 {
|
||||
continue
|
||||
}
|
||||
values = append(values, q...)
|
||||
}
|
||||
samples, err := core.FromFloats(values, opts.Samples, dim)
|
||||
if err != nil {
|
||||
return nil, errf("%s: %w", name, err)
|
||||
}
|
||||
return samples, nil
|
||||
}
|
||||
@@ -0,0 +1,217 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// gaussianLogDensity builds the log density of independent standard
|
||||
// normals: log π(q) = −‖q‖²/2.
|
||||
func gaussianLogDensity(q *Tensor) (*Tensor, error) {
|
||||
sq, err := q.Mul(q)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
total, err := sq.Sum()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return total.Scale(-0.5)
|
||||
}
|
||||
|
||||
// gammaLogDensity builds log π(x) = log x − x, the unnormalised log
|
||||
// density of a Gamma(2, 1) distribution.
|
||||
func gammaLogDensity(q *Tensor) (*Tensor, error) {
|
||||
logq, err := q.Log()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
shifted, err := logq.Sub(q)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return shifted.Sum()
|
||||
}
|
||||
|
||||
// TestSampleHMCNormal runs the chain on a two-dimensional standard
|
||||
// normal and pins the sample moments: the target's mean is zero, its
|
||||
// variance one and its components independent. A fixed seed makes the
|
||||
// draw deterministic, so the bounds are checked facts about this run,
|
||||
// not hopes about a random one.
|
||||
func TestSampleHMCNormal(t *testing.T) {
|
||||
q0, err := core.FromFloats([]float64{2, -2}, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
samples, err := SampleHMC(gaussianLogDensity, q0,
|
||||
HMCOptions{Step: 0.3, Steps: 20, BurnIn: 500, Samples: 6000, Thin: 1, Seed: 42})
|
||||
if err != nil {
|
||||
t.Fatalf("SampleHMC: %v", err)
|
||||
}
|
||||
if got := samples.Shape(); got[0] != 6000 || got[1] != 2 {
|
||||
t.Fatalf("samples shape %v, want [6000 2]", got)
|
||||
}
|
||||
rows := samples.Shape()[0]
|
||||
mean := []float64{0, 0}
|
||||
variance := []float64{0, 0}
|
||||
for r := range rows {
|
||||
for c := range 2 {
|
||||
v := samples.FloatAt(r*2 + c)
|
||||
mean[c] += v / float64(rows)
|
||||
}
|
||||
}
|
||||
for r := range rows {
|
||||
for c := range 2 {
|
||||
d := samples.FloatAt(r*2+c) - mean[c]
|
||||
variance[c] += d * d / float64(rows)
|
||||
}
|
||||
}
|
||||
for c := range 2 {
|
||||
if math.Abs(mean[c]) > 0.15 {
|
||||
t.Fatalf("mean[%d] = %.4g, want |mean| ≤ 0.15", c, mean[c])
|
||||
}
|
||||
if math.Abs(variance[c]-1) > 0.2 {
|
||||
t.Fatalf("variance[%d] = %.4g, want 1 ± 0.2", c, variance[c])
|
||||
}
|
||||
}
|
||||
// Cross moment of the independent components.
|
||||
cov := 0.0
|
||||
for r := range rows {
|
||||
cov += (samples.FloatAt(r*2) - mean[0]) * (samples.FloatAt(r*2+1) - mean[1]) / float64(rows)
|
||||
}
|
||||
if math.Abs(cov) > 0.15 {
|
||||
t.Fatalf("cross moment = %.4g, want |cov| ≤ 0.15", cov)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSampleHMCDeterministic checks the seed contract: the same seed
|
||||
// replays bit-identically, a different seed does not.
|
||||
func TestSampleHMCDeterministic(t *testing.T) {
|
||||
q0, _ := core.FromFloats([]float64{1}, 1)
|
||||
run := func(seed int64) *core.Array {
|
||||
samples, err := SampleHMC(gaussianLogDensity, q0,
|
||||
HMCOptions{Step: 0.4, Steps: 16, BurnIn: 100, Samples: 200, Seed: seed})
|
||||
if err != nil {
|
||||
t.Fatalf("SampleHMC: %v", err)
|
||||
}
|
||||
return samples
|
||||
}
|
||||
a, b := run(7), run(7)
|
||||
c := run(8)
|
||||
for i := range a.Len() {
|
||||
if a.FloatAt(i) != b.FloatAt(i) {
|
||||
t.Fatalf("the same seed produced different samples at %d", i)
|
||||
}
|
||||
}
|
||||
same := true
|
||||
for i := range a.Len() {
|
||||
if a.FloatAt(i) != c.FloatAt(i) {
|
||||
same = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if same {
|
||||
t.Fatal("different seeds produced identical samples")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSampleHMCSupportedDensity samples a Gamma(2, 1) target,
|
||||
// log π(x) = log x − x on x > 0. Proposals that overshoot into the
|
||||
// forbidden half-line yield a NaN density and are rejected, so every
|
||||
// kept sample stays positive and the mean approaches 2.
|
||||
func TestSampleHMCSupportedDensity(t *testing.T) {
|
||||
q0, _ := core.FromFloats([]float64{1}, 1)
|
||||
samples, err := SampleHMC(gammaLogDensity, q0,
|
||||
HMCOptions{Step: 0.3, Steps: 10, BurnIn: 500, Samples: 6000, Seed: 3})
|
||||
if err != nil {
|
||||
t.Fatalf("SampleHMC: %v", err)
|
||||
}
|
||||
sum := 0.0
|
||||
for i := range samples.Len() {
|
||||
x := samples.FloatAt(i)
|
||||
if x <= 0 {
|
||||
t.Fatalf("sample %d = %g left the support", i, x)
|
||||
}
|
||||
sum += x
|
||||
}
|
||||
mean := sum / float64(samples.Len())
|
||||
if math.Abs(mean-2) > 0.15 {
|
||||
t.Fatalf("Gamma(2) mean = %.4g, want 2 ± 0.15", mean)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSampleHMCRejectedByGuard exercises the explicit rejection path:
|
||||
// the density errors outside its support instead of returning NaN,
|
||||
// and the chain still stays inside it.
|
||||
func TestSampleHMCRejectedByGuard(t *testing.T) {
|
||||
q0, _ := core.FromFloats([]float64{0.5}, 1)
|
||||
samples, err := SampleHMC(func(q *Tensor) (*Tensor, error) {
|
||||
x := q.Data().FloatAt(0)
|
||||
if x <= 0 {
|
||||
return nil, errf("outside the support")
|
||||
}
|
||||
return gammaLogDensity(q)
|
||||
}, q0, HMCOptions{Step: 0.5, Steps: 20, BurnIn: 300, Samples: 2000, Seed: 5})
|
||||
if err != nil {
|
||||
t.Fatalf("SampleHMC: %v", err)
|
||||
}
|
||||
for i := range samples.Len() {
|
||||
if samples.FloatAt(i) <= 0 {
|
||||
t.Fatalf("sample %d = %g left the support", i, samples.FloatAt(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestSampleHMCErrors pins the validation contract, including a
|
||||
// density that fails at the start, returns a non-scalar, or never
|
||||
// touches the leaf core.
|
||||
func TestSampleHMCErrors(t *testing.T) {
|
||||
q0, _ := core.FromFloats([]float64{1}, 1)
|
||||
if _, err := SampleHMC(nil, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
|
||||
t.Fatal("expected an error for a nil density")
|
||||
}
|
||||
rank2, _ := core.FromFloats([]float64{1, 1}, 1, 2)
|
||||
if _, err := SampleHMC(gaussianLogDensity, rank2, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
|
||||
t.Fatal("expected an error for a rank-2 state")
|
||||
}
|
||||
if _, err := SampleHMC(gaussianLogDensity, nil, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
|
||||
t.Fatal("expected an error for a nil state")
|
||||
}
|
||||
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Steps: 5, Samples: 1}); err == nil {
|
||||
t.Fatal("expected an error for a non-positive step")
|
||||
}
|
||||
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Samples: 1}); err == nil {
|
||||
t.Fatal("expected an error for a non-positive step count")
|
||||
}
|
||||
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Steps: 5}); err == nil {
|
||||
t.Fatal("expected an error for a non-positive sample count")
|
||||
}
|
||||
if _, err := SampleHMC(gaussianLogDensity, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1, BurnIn: -3}); err == nil {
|
||||
t.Fatal("expected an error for a negative BurnIn")
|
||||
}
|
||||
nonScalar := func(q *Tensor) (*Tensor, error) {
|
||||
two, _ := core.FromFloats([]float64{1, 2}, 2)
|
||||
return FromArray(two, false), nil
|
||||
}
|
||||
if _, err := SampleHMC(nonScalar, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
|
||||
t.Fatal("expected an error for a non-scalar density")
|
||||
}
|
||||
detached := func(q *Tensor) (*Tensor, error) {
|
||||
one, _ := core.FromFloats([]float64{1}, 1)
|
||||
return FromArray(one, false), nil
|
||||
}
|
||||
if _, err := SampleHMC(detached, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
|
||||
t.Fatal("expected an error for a density disconnected from the leaf")
|
||||
}
|
||||
failsAtStart := func(q *Tensor) (*Tensor, error) {
|
||||
return nil, errf("no density at the start")
|
||||
}
|
||||
if _, err := SampleHMC(failsAtStart, q0, HMCOptions{Step: 0.1, Steps: 5, Samples: 1}); err == nil {
|
||||
t.Fatal("expected the start-time density error to be fatal")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import "testing"
|
||||
|
||||
// TestHessianVectorProductZeroDirectionShape pins the repair: a zero
|
||||
// direction returned a flat length-n vector while
|
||||
// every other answer carries the point's own shape.
|
||||
func TestHessianVectorProductZeroDirectionShape(t *testing.T) {
|
||||
x0 := []float64{0.5, -1.25, 2, -2}
|
||||
xt, err := FromFloat64s(x0, false, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
f := func(z *Tensor) (*Tensor, error) { return z.Sum() }
|
||||
v, err := FromFloat64s(make([]float64, 4), false, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s: %v", err)
|
||||
}
|
||||
hv, err := HessianVectorProduct(f, xt, v, HessianOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("HessianVectorProduct: %v", err)
|
||||
}
|
||||
want := []int{2, 2}
|
||||
got := hv.Shape()
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("H·0 shape = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
for i := range 4 {
|
||||
if hv.FloatAt(i) != 0 {
|
||||
t.Fatalf("H·0 = %v, want the zero matrix", hv.FloatAt(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Regression pins for the MatMul adjoints: the 1-D × 2-D branch against
|
||||
// central differences, real and complex, and the Newton-CG loop's last
|
||||
// iteration.
|
||||
|
||||
// TestMatMulVectorByMatrixGradient pins the 1-D × 2-D branch of
|
||||
// the MatMul adjoint (da = g·Bᵀ, db = outer(a, g)), real and complex,
|
||||
// against central differences. The branch had no gradient test at all.
|
||||
func TestMatMulVectorByMatrixGradient(t *testing.T) {
|
||||
avec := []float64{1.5, -0.5, 2, 0.25}
|
||||
bvec := []float64{
|
||||
0.5, -1, 2,
|
||||
1.5, 0.25, -0.75,
|
||||
-2, 1, 0.5,
|
||||
1, -0.5, 1.25,
|
||||
}
|
||||
// A weighted linear loss, so every output slot carries its own
|
||||
// coefficient and a wrong routing cannot cancel against another.
|
||||
w := []float64{0.7, -1.3, 2.1}
|
||||
lossOf := func(a, b *core.Array) float64 {
|
||||
out, err := core.MatMul2D(a, b)
|
||||
if err != nil {
|
||||
t.Fatalf("MatMul2D: %v", err)
|
||||
}
|
||||
s := 0.0
|
||||
for i := range out.Len() {
|
||||
s += w[i] * out.FloatAt(i)
|
||||
}
|
||||
return s
|
||||
}
|
||||
cloneWith := func(a *core.Array, i int, v float64) *core.Array {
|
||||
vals := make([]float64, a.Len())
|
||||
for k := range a.Len() {
|
||||
vals[k] = a.FloatAt(k)
|
||||
}
|
||||
vals[i] = v
|
||||
out, err := core.FromFloats(vals, a.Shape()...)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
a, _ := FromFloat64s(avec, true, 4)
|
||||
b, _ := FromFloat64s(bvec, true, 4, 3)
|
||||
if err := backwardWeightedMatMul(t, a, b, w); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
baseA, _ := core.FromFloats(avec, 4)
|
||||
baseB, _ := core.FromFloats(bvec, 4, 3)
|
||||
for i := range 4 {
|
||||
eps := 1e-6
|
||||
want := (lossOf(cloneWith(baseA, i, avec[i]+eps), baseB) -
|
||||
lossOf(cloneWith(baseA, i, avec[i]-eps), baseB)) / (2 * eps)
|
||||
if math.Abs(a.Grad().FloatAt(i)-want) > 1e-5*(1+math.Abs(want)) {
|
||||
t.Fatalf("da[%d] = %g, want %g", i, a.Grad().FloatAt(i), want)
|
||||
}
|
||||
}
|
||||
for i := range 12 {
|
||||
eps := 1e-6
|
||||
want := (lossOf(baseA, cloneWith(baseB, i, bvec[i]+eps)) -
|
||||
lossOf(baseA, cloneWith(baseB, i, bvec[i]-eps))) / (2 * eps)
|
||||
if math.Abs(b.Grad().FloatAt(i)-want) > 1e-5*(1+math.Abs(want)) {
|
||||
t.Fatalf("db[%d] = %g, want %g", i, b.Grad().FloatAt(i), want)
|
||||
}
|
||||
}
|
||||
|
||||
// The complex 1-D × 2-D branch against numericComplexGrad, under
|
||||
// the same weighted fold backwardComplex builds: L = Σ Re(w̄·y) +
|
||||
// Σ|y|²/n with the helper's own deterministic weights.
|
||||
cb := FromArray(mustComplexes([]complex128{
|
||||
0.5 + 0.5i, -1,
|
||||
1.5, 0.25 - 0.75i,
|
||||
-2 + 1i, 1,
|
||||
}, 3, 2), true)
|
||||
op := func(x *Tensor) (*Tensor, error) { return x.MatMul(cb) }
|
||||
vals := []complex128{1 + 0.5i, -0.25 - 1i, 0.75 + 0.25i}
|
||||
xt := backwardComplex(t, op, vals, 3)
|
||||
closs := complexLossOf(t, op, spectralWeights(2, 11))
|
||||
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(closs, xt.Data()), 1e-6)
|
||||
}
|
||||
|
||||
// backwardWeightedMatMul builds L = w·(a·B) over fresh leaves and runs
|
||||
// one backward pass.
|
||||
func backwardWeightedMatMul(t *testing.T, a, b *Tensor, w []float64) error {
|
||||
t.Helper()
|
||||
out, err := a.MatMul(b)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
wv, err := FromFloat64s(w, false, len(w))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
loss, err := out.Mul(wv)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sum, err := loss.Sum()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sum.Backward()
|
||||
return nil
|
||||
}
|
||||
|
||||
// TestNewtonCGConvergesOnTheLastIteration pins that a tolerance
|
||||
// met exactly on the final permitted iteration is a success, not a
|
||||
// budget error whose message prints a gradient already under the
|
||||
// tolerance.
|
||||
func TestNewtonCGConvergesOnTheLastIteration(t *testing.T) {
|
||||
f := func(x *Tensor) (*Tensor, error) {
|
||||
d, err := x.Sub(mustTensorF64(1.5))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return d.Mul(d)
|
||||
}
|
||||
x0, _ := core.FromFloats([]float64{0}, 1)
|
||||
// One truncated-CG step lands within the Hessian-product rounding
|
||||
// of the minimiser (about 1e-11 here); a tolerance of 1e-9 is met
|
||||
// by exactly that step, on the final permitted iteration.
|
||||
out, _, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{MaxIterations: 1, Tolerance: 1e-9})
|
||||
if err != nil {
|
||||
t.Fatalf("MinimiseNewtonCG on a quadratic with one exact step: %v", err)
|
||||
}
|
||||
if math.Abs(out.FloatAt(0)-1.5) > 1e-9 {
|
||||
t.Fatalf("minimum = %g, want 1.5", out.FloatAt(0))
|
||||
}
|
||||
}
|
||||
|
||||
// mustTensorF64 wraps one float as a no-grad tensor.
|
||||
func mustTensorF64(v float64) *Tensor {
|
||||
t, err := FromFloat64s([]float64{v}, false, 1)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return t
|
||||
}
|
||||
|
||||
// mustComplexes builds a complex array or fails the test.
|
||||
func mustComplexes(vals []complex128, shape ...int) *core.Array {
|
||||
a, err := core.FromComplexes(vals, shape...)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Newton-CG minimisation: truncated conjugate gradients on
|
||||
// the Hessian system, driven by autograd. It lives in the grad package
|
||||
// because it is meaningless without the graph: the gradients come from
|
||||
// Backward and the Hessian never forms, each CG iteration buying one
|
||||
// Hessian-vector product for two backward passes. That is the
|
||||
// optimiser large problems want, where a dense second derivative does
|
||||
// not fit memory and the numerical-difference optimisers of the optim
|
||||
// package lose their accuracy.
|
||||
|
||||
// NewtonCGOptions tunes MinimiseNewtonCG. MaxIterations bounds the
|
||||
// outer Newton steps (default 100); Tolerance stops when the gradient
|
||||
// norm falls under it (default 1e-8); MaxCGIterations bounds the inner
|
||||
// CG solve per outer step (default n, the problem dimension).
|
||||
type NewtonCGOptions struct {
|
||||
MaxIterations int
|
||||
Tolerance float64
|
||||
MaxCGIterations int
|
||||
}
|
||||
|
||||
// MinimiseNewtonCG returns the point and value of a local minimum of
|
||||
// the scalar objective f near x0 by the Newton-CG method: each step
|
||||
// solves H·s = −∇f with truncated conjugate gradients (negative
|
||||
// curvature stops the solve and falls back to the first direction),
|
||||
// then an Armijo backtracking line search secures descent. f receives
|
||||
// a leaf tensor and must return a single-element real tensor. A
|
||||
// non-finite objective, an unreachable Armijo condition or an
|
||||
// exhausted iteration budget is an error naming the state it stopped
|
||||
// in; the converged answer is a fresh array the caller owns. The
|
||||
// gradient evaluations differentiate the graph without committing
|
||||
// anything, so the accumulated gradients of the tensors f closes over
|
||||
// are left exactly as they were, on the success and the error path
|
||||
// alike.
|
||||
func MinimiseNewtonCG(f func(*Tensor) (*Tensor, error), x0 *core.Array, opts NewtonCGOptions) (*core.Array, float64, error) {
|
||||
const name = "MinimiseNewtonCG"
|
||||
if f == nil {
|
||||
return nil, 0, errf("%s: f must not be nil", name)
|
||||
}
|
||||
if x0 == nil {
|
||||
return nil, 0, errf("%s: the starting point must not be nil", name)
|
||||
}
|
||||
n := x0.Len()
|
||||
if n == 0 {
|
||||
return nil, 0, errf("%s: the starting point must have at least one element", name)
|
||||
}
|
||||
if x0.Dtype() == core.Complex {
|
||||
return nil, 0, errf("%s: complex starting points are not supported", name)
|
||||
}
|
||||
maxIter := opts.MaxIterations
|
||||
if maxIter <= 0 {
|
||||
maxIter = 100
|
||||
}
|
||||
tol := opts.Tolerance
|
||||
if tol <= 0 {
|
||||
tol = 1e-8
|
||||
}
|
||||
maxCG := opts.MaxCGIterations
|
||||
if maxCG <= 0 {
|
||||
maxCG = n
|
||||
}
|
||||
|
||||
// eval runs the objective and its reverse pass at point p, returning
|
||||
// the loss and the flattened gradient. The pass commits nothing, so
|
||||
// the caller's own gradients survive the evaluation untouched.
|
||||
eval := func(p *core.Array) (float64, *core.Array, error) {
|
||||
xt := FromArray(p, true)
|
||||
y, err := f(xt)
|
||||
if err != nil {
|
||||
return 0, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
if y.Data().Len() != 1 {
|
||||
return 0, nil, errf("%s: the objective must return a scalar, got %d elements", name, y.Data().Len())
|
||||
}
|
||||
v := y.Data().FloatAt(0)
|
||||
if math.IsNaN(v) || math.IsInf(v, 0) {
|
||||
return 0, nil, errf("%s: the objective is non-finite (%g)", name, v)
|
||||
}
|
||||
grads, err := y.reverseGrads()
|
||||
if err != nil {
|
||||
return 0, nil, errf("%s: %w", name, err)
|
||||
}
|
||||
g := grads[xt]
|
||||
if g == nil {
|
||||
return 0, nil, errf("%s: the objective does not depend on the starting point", name)
|
||||
}
|
||||
return v, g, nil
|
||||
}
|
||||
|
||||
x := x0
|
||||
f0, g, err := eval(x)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
for iter := 1; iter <= maxIter; iter++ {
|
||||
gnorm := flatNorm(g)
|
||||
if gnorm <= tol {
|
||||
return clonePoint(x), f0, nil
|
||||
}
|
||||
|
||||
// Truncated CG on H·s = −g. The Hessian acts through the
|
||||
// Hessian-vector product, two backward passes per iteration.
|
||||
xt := FromArray(x, true)
|
||||
s := make([]float64, n)
|
||||
r := make([]float64, n)
|
||||
p := make([]float64, n)
|
||||
gs := make([]float64, n)
|
||||
gFloats := flatFloats(g)
|
||||
for i := range n {
|
||||
r[i] = -gFloats[i]
|
||||
p[i] = r[i]
|
||||
gs[i] = gFloats[i]
|
||||
}
|
||||
rr := 0.0
|
||||
for i := range n {
|
||||
rr += r[i] * r[i]
|
||||
}
|
||||
for cg := 0; cg < maxCG; cg++ {
|
||||
pArr, herr := core.FromFloats(p, n)
|
||||
if herr != nil {
|
||||
return nil, 0, errf("%s: %w", name, herr)
|
||||
}
|
||||
hp, herr2 := HessianVectorProduct(f, xt, FromArray(pArr, false), HessianOptions{})
|
||||
if herr2 != nil {
|
||||
return nil, 0, errf("%s: %w", name, herr2)
|
||||
}
|
||||
hpF := flatFloats(hp)
|
||||
pHp := 0.0
|
||||
for i := range n {
|
||||
pHp += p[i] * hpF[i]
|
||||
}
|
||||
if pHp <= 0 {
|
||||
// Negative or vanishing curvature: the quadratic model
|
||||
// is not convex here. The first iteration falls back to
|
||||
// the steepest descent direction; later ones keep what
|
||||
// the solve has accumulated.
|
||||
if cg == 0 {
|
||||
copy(s, p)
|
||||
}
|
||||
break
|
||||
}
|
||||
alpha := rr / pHp
|
||||
for i := range n {
|
||||
s[i] += alpha * p[i]
|
||||
r[i] -= alpha * hpF[i]
|
||||
}
|
||||
rrNew := 0.0
|
||||
for i := range n {
|
||||
rrNew += r[i] * r[i]
|
||||
}
|
||||
if math.Sqrt(rrNew) <= 0.1*gnorm {
|
||||
break
|
||||
}
|
||||
beta := rrNew / rr
|
||||
for i := range n {
|
||||
p[i] = r[i] + beta*p[i]
|
||||
}
|
||||
rr = rrNew
|
||||
}
|
||||
|
||||
// Armijo backtracking along s; gᵀs is negative by construction.
|
||||
gsDot := 0.0
|
||||
for i := range n {
|
||||
gsDot += gs[i] * s[i]
|
||||
}
|
||||
if gsDot >= 0 {
|
||||
return nil, 0, errf("%s: the CG direction does not descend at step %d", name, iter)
|
||||
}
|
||||
step := 1.0
|
||||
xFloats := flatFloats(x)
|
||||
var xNew *core.Array
|
||||
var fNew float64
|
||||
accepted := false
|
||||
for range 40 {
|
||||
vals := make([]float64, n)
|
||||
for i := range n {
|
||||
vals[i] = xFloats[i] + step*s[i]
|
||||
}
|
||||
cand, cerr := core.FromFloats(vals, x.Shape()...)
|
||||
if cerr != nil {
|
||||
return nil, 0, errf("%s: %w", name, cerr)
|
||||
}
|
||||
// cand assigns to the outer xNew; a := here would shadow
|
||||
// it and hand the post-loop update a nil.
|
||||
fNew, g, err = eval(cand)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if fNew <= f0+1e-4*step*gsDot {
|
||||
accepted = true
|
||||
xNew = cand
|
||||
break
|
||||
}
|
||||
step /= 2
|
||||
}
|
||||
if !accepted {
|
||||
return nil, 0, errf("%s: the line search found no descent at step %d (f = %.6g)", name, iter, f0)
|
||||
}
|
||||
x = xNew
|
||||
f0 = fNew
|
||||
}
|
||||
// The last accepted step updated g after the loop-top test, so a
|
||||
// run whose tolerance was met exactly on the final iteration must
|
||||
// re-test before the budget refusal reports it; the message below
|
||||
// would otherwise print a gradient already under the tolerance.
|
||||
if flatNorm(g) <= tol {
|
||||
return clonePoint(x), f0, nil
|
||||
}
|
||||
return nil, 0, errf("%s: no convergence in %d steps (gradient norm %.3g)", name, maxIter, flatNorm(g))
|
||||
}
|
||||
|
||||
// flatNorm returns the Euclidean norm of a flattened gradient. The sum
|
||||
// runs in ascending element order on the raw payload when it can; the
|
||||
// walk is bounded by the element count, not the payload, because a
|
||||
// rebased view's storage may run longer than its own elements.
|
||||
func flatNorm(g *core.Array) float64 {
|
||||
s := 0.0
|
||||
if !g.Strided() && g.Dtype() == core.Float {
|
||||
gs := g.RawFloats()
|
||||
for i := range g.Len() {
|
||||
s += gs[i] * gs[i]
|
||||
}
|
||||
return math.Sqrt(s)
|
||||
}
|
||||
for i := range g.Len() {
|
||||
s += g.FloatAt(i) * g.FloatAt(i)
|
||||
}
|
||||
return math.Sqrt(s)
|
||||
}
|
||||
|
||||
// clonePoint copies the converged point so the caller owns it.
|
||||
func clonePoint(x *core.Array) *core.Array {
|
||||
vals := make([]float64, x.Len())
|
||||
for i := range vals {
|
||||
vals[i] = x.FloatAt(i)
|
||||
}
|
||||
out, _ := core.FromFloats(vals, x.Shape()...)
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// TestNewtonCGQuadratic pins the exactly-Newtonian case: a quadratic
|
||||
// with SPD Hessian converges to the analytic minimiser in a couple of
|
||||
// steps.
|
||||
func TestNewtonCGQuadratic(t *testing.T) {
|
||||
a := []float64{4, 1, 1, 3}
|
||||
b := []float64{-1, 2}
|
||||
x0, err := core.FromFloats([]float64{0.5, -1.25}, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
f := func(z *Tensor) (*Tensor, error) {
|
||||
az, err := FromFloat64s(a, false, 2, 2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bz, err := FromFloat64s(b, false, 2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
halfA, err := az.Scale(0.5)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
azx, err := halfA.MatMul(z)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
lin, err := azx.Add(bz)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
prod, err := lin.Mul(z)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return prod.Sum()
|
||||
}
|
||||
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{Tolerance: 1e-12})
|
||||
if err != nil {
|
||||
t.Fatalf("MinimiseNewtonCG: %v", err)
|
||||
}
|
||||
// x* = −A⁻¹b: solve 4x+y = 1, x+3y = −2 so x = 5/11, y = −9/11.
|
||||
if math.Abs(x.FloatAt(0)-5.0/11) > 1e-9 || math.Abs(x.FloatAt(1)+9.0/11) > 1e-9 {
|
||||
t.Fatalf("minimiser = (%g, %g), want (5/11, -9/11)", x.FloatAt(0), x.FloatAt(1))
|
||||
}
|
||||
// f* = ½x*ᵀAx* + bᵀx* = 253/242 − 23/11 = −253/242.
|
||||
const want = -253.0 / 242.0
|
||||
if math.Abs(fv-want) > 1e-10 {
|
||||
t.Fatalf("value = %.12g, want %.12g", fv, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewtonCGRosenbrock pins a nonquadratic valley: the classic
|
||||
// Rosenbrock minimum at (1, 1) from the far side.
|
||||
func TestNewtonCGRosenbrock(t *testing.T) {
|
||||
x0, err := core.FromFloats([]float64{-1.5, 2}, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
f := func(z *Tensor) (*Tensor, error) {
|
||||
x0t, err := z.Slice(0, 0, 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
x1t, err := z.Slice(0, 1, 2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
x0sq, err := x0t.Pow(2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
diff, err := x1t.Sub(x0sq)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
term1, err := diff.Pow(2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
one, err := FromFloat64s([]float64{1}, false, 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
x0m1, err := x0t.Sub(one)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
term2, err := x0m1.Pow(2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
term2s, err := term2.Scale(100)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
total, err := term1.Add(term2s)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return total.Sum()
|
||||
}
|
||||
x, _, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{Tolerance: 1e-7, MaxIterations: 200})
|
||||
if err != nil {
|
||||
t.Fatalf("MinimiseNewtonCG: %v", err)
|
||||
}
|
||||
if math.Abs(x.FloatAt(0)-1) > 1e-4 || math.Abs(x.FloatAt(1)-1) > 1e-4 {
|
||||
t.Fatalf("minimiser = (%.6f, %.6f), want (1, 1)", x.FloatAt(0), x.FloatAt(1))
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewtonCGNegativeCurvature pins the fallback: a double well
|
||||
// whose start sits in the concave region between the minima. The CG
|
||||
// must take its steepest-descent fallback there and still land in a
|
||||
// well.
|
||||
func TestNewtonCGNegativeCurvature(t *testing.T) {
|
||||
x0, err := core.FromFloats([]float64{0.1, 0.2}, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
// f = Σ(x⁴ − x²): Hessian 12x² − 2 is negative for |x| < 1/√6,
|
||||
// so the start is concave; the wells sit at ±1/√2 per coordinate.
|
||||
f := func(z *Tensor) (*Tensor, error) {
|
||||
q, err := z.Pow(4)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := z.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
d, err := q.Sub(sq)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return d.Sum()
|
||||
}
|
||||
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{Tolerance: 1e-9})
|
||||
if err != nil {
|
||||
t.Fatalf("MinimiseNewtonCG: %v", err)
|
||||
}
|
||||
const well = 1.0 / math.Sqrt2
|
||||
for i := range 2 {
|
||||
if math.Abs(math.Abs(x.FloatAt(i))-well) > 1e-6 {
|
||||
t.Fatalf("coordinate %d = %g, want magnitude %g", i, x.FloatAt(i), well)
|
||||
}
|
||||
}
|
||||
// f at a well: Σ(1/4 − 1/2) = −1/2.
|
||||
if math.Abs(fv+0.5) > 1e-9 {
|
||||
t.Fatalf("value = %.12g, want -0.5", fv)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewtonCGScalarInput pins the n = 1 path and the error contract.
|
||||
func TestNewtonCGScalarInput(t *testing.T) {
|
||||
x0, err := core.FromFloats([]float64{3}, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
f := func(z *Tensor) (*Tensor, error) {
|
||||
sq, err := z.Pow(2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
four, err := sq.Scale(4)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return four.Sum()
|
||||
}
|
||||
x, fv, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{})
|
||||
if err != nil {
|
||||
t.Fatalf("MinimiseNewtonCG: %v", err)
|
||||
}
|
||||
if math.Abs(x.FloatAt(0)) > 1e-7 || math.Abs(fv) > 1e-12 {
|
||||
t.Fatalf("minimiser = %g, value = %g", x.FloatAt(0), fv)
|
||||
}
|
||||
// A slice of the point itself: the objective is disconnected from
|
||||
// the minimised point, so the run exhausts its iterations on a
|
||||
// constant value and errors loudly.
|
||||
c, _ := core.FromFloats([]float64{1}, 1)
|
||||
if _, _, err := MinimiseNewtonCG(func(z *Tensor) (*Tensor, error) { return z.Slice(0, 0, 1) }, c, NewtonCGOptions{}); err == nil {
|
||||
t.Fatal("a slice of the point itself minimised without error")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// TransposeAxes reorders the axes of a tensor by dims. The backward
|
||||
// applies the inverse permutation to the incoming gradient: axis moves
|
||||
// are invertible data motion, so no element mixing occurs and the
|
||||
// gradient is exactly the same move played backwards.
|
||||
func (t *Tensor) TransposeAxes(dims ...int) (*Tensor, error) {
|
||||
if err := t.checkDiff("TransposeAxes"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out, err := core.TransposeAxes(t.data, dims...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
perm := append([]int(nil), dims...)
|
||||
orig := t.data.Shape()
|
||||
return t.unaryResult("TransposeAxes", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
dx, err := core.TransposeAxes(g.arr, inversePerm(perm)...)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dst[0] = gradSlot{arr: dx, sh: orig}
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// inversePerm flips an axis permutation: if out = permute(x, p), then
|
||||
// permute(out, p⁻¹) restores x's axis order.
|
||||
func inversePerm(perm []int) []int {
|
||||
inv := make([]int, len(perm))
|
||||
for i, p := range perm {
|
||||
inv[p] = i
|
||||
}
|
||||
return inv
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
func TestTensorTransposeAxesValues(t *testing.T) {
|
||||
x, _ := core.FromFloats([]float64{
|
||||
1, 2, 3,
|
||||
4, 5, 6,
|
||||
}, 2, 3)
|
||||
|
||||
out, err := FromArray(x, false).TransposeAxes(1, 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := out.Data().Shape(); got[0] != 3 || got[1] != 2 {
|
||||
t.Fatalf("shape: %v", got)
|
||||
}
|
||||
want := []float64{1, 4, 2, 5, 3, 6}
|
||||
for i := range want {
|
||||
if g := out.Data().FloatAt(i); g != want[i] {
|
||||
t.Fatalf("[%d] = %v, want %v", i, g, want[i])
|
||||
}
|
||||
}
|
||||
|
||||
// A rank-3 rotation moves the trailing axis to the front.
|
||||
y, _ := core.FromFloats([]float64{
|
||||
1, 2, 3, 4,
|
||||
5, 6, 7, 8,
|
||||
9, 10, 11, 12,
|
||||
}, 2, 3, 2)
|
||||
rotated, err := FromArray(y, false).TransposeAxes(2, 0, 1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := rotated.Data().Shape(); got[0] != 2 || got[1] != 2 || got[2] != 3 {
|
||||
t.Fatalf("rank-3 shape: %v", got)
|
||||
}
|
||||
|
||||
// Invalid permutations error before any graph work.
|
||||
if _, err := FromArray(x, false).TransposeAxes(0, 0); err == nil {
|
||||
t.Fatal("duplicate axis accepted")
|
||||
}
|
||||
if _, err := FromArray(x, false).TransposeAxes(0); err == nil {
|
||||
t.Fatal("short permutation accepted")
|
||||
}
|
||||
}
|
||||
|
||||
// TestTensorTransposeAxesGradient routes a weighted sum through the
|
||||
// permutation: the analytic input gradient is exactly the weight tensor
|
||||
// played back through the inverse permutation.
|
||||
func TestTensorTransposeAxesGradient(t *testing.T) {
|
||||
x, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
||||
w, _ := core.FromFloats([]float64{0.5, -1, 2, 0.25, -0.75, 1.5}, 3, 2)
|
||||
|
||||
xt := FromArray(x, true)
|
||||
joint, err := xt.TransposeAxes(1, 0) // gives (3, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
scaled, err := joint.Mul(FromArray(w, false))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loss, err := scaled.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
g := xt.Grad()
|
||||
if g == nil || g.Dtype() != core.Float {
|
||||
t.Fatalf("gradient missing or wrong dtype: %v", g)
|
||||
}
|
||||
for i := range 6 {
|
||||
row, col := i/3, i%3
|
||||
if got := g.FloatAt(i); got != w.FloatAt(col*2+row) {
|
||||
t.Errorf("grad[%d] = %v, want %v", i, got, w.FloatAt(col*2+row))
|
||||
}
|
||||
}
|
||||
|
||||
// Round-trip: permuting by (1,0) then back restores the values.
|
||||
back, err := joint.TransposeAxes(1, 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range 6 {
|
||||
if back.Data().FloatAt(i) != x.FloatAt(i) {
|
||||
t.Fatalf("round-trip[%d] = %v, want %v", i, back.Data().FloatAt(i), x.FloatAt(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
+352
@@ -0,0 +1,352 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Gradient buffer recycling for the backward sweep. A sweep allocates
|
||||
// one gradient array per node output and per folded contribution; the
|
||||
// arrays die within the sweep that made them, except the ones committed
|
||||
// to leaves, which escape to the caller. The pool reclaims the
|
||||
// intermediates: a borrowed array arrives with a fully zeroed payload,
|
||||
// the pool retains a bounded number of elements, and nothing is
|
||||
// recycled while a live reference to it exists. That last rule rests on
|
||||
// an invariant the closures maintain: a backward closure never returns
|
||||
// the incoming gradient buffer itself, never returns one buffer in two
|
||||
// slots and never returns a view of another live array, so an entry the
|
||||
// sweep releases is unreachable from the graph, from the returned map
|
||||
// and from every other entry.
|
||||
//
|
||||
// An array the pool declines (an exotic dtype, a rank above six, a
|
||||
// payload above the cap, a full bucket) falls back to ordinary
|
||||
// allocation; correctness never depends on a hit.
|
||||
|
||||
// gradSlot is one gradient array together with the shape it was built
|
||||
// for, which is the pool's reuse key. Slots travel instead of bare
|
||||
// arrays so the sweep can release a buffer without re-deriving its
|
||||
// shape, which would cost an allocation of its own.
|
||||
type gradSlot struct {
|
||||
arr *core.Array
|
||||
sh []int
|
||||
}
|
||||
|
||||
// poolKey identifies a gradient buffer exactly: dtype, rank, element
|
||||
// count and dimensions. The fixed dimension array keeps the key
|
||||
// comparable, so buckets need no stored shape for matching.
|
||||
type poolKey struct {
|
||||
dt core.Dtype
|
||||
nd int8
|
||||
n int
|
||||
d [6]int
|
||||
}
|
||||
|
||||
// poolKeyOf builds the key for dt and shape, reporting false for
|
||||
// anything the pool does not accept.
|
||||
func poolKeyOf(dt core.Dtype, shape []int) (poolKey, bool) {
|
||||
var k poolKey
|
||||
if len(shape) == 0 || len(shape) > len(k.d) {
|
||||
return k, false
|
||||
}
|
||||
switch dt {
|
||||
case core.Float, core.Float32, core.Complex:
|
||||
default:
|
||||
return k, false
|
||||
}
|
||||
n := 1
|
||||
for i, d := range shape {
|
||||
n *= d
|
||||
k.d[i] = d
|
||||
}
|
||||
if n > gradPoolMaxArrayElems {
|
||||
return k, false
|
||||
}
|
||||
k.dt, k.nd, k.n = dt, int8(len(shape)), n
|
||||
return k, true
|
||||
}
|
||||
|
||||
// The retention caps: no more than gradPoolMaxArrayElems elements in
|
||||
// one array, gradPoolMaxElems retained across all buckets and
|
||||
// gradPoolPerBucket arrays of one exact shape. A buffer outside the
|
||||
// caps is dropped to the garbage collector instead of retained, so the
|
||||
// pool cannot pin memory beyond these bounds however hard one workload
|
||||
// pushes it.
|
||||
const (
|
||||
gradPoolMaxArrayElems = 1 << 20
|
||||
gradPoolMaxElems = 1 << 20
|
||||
gradPoolPerBucket = 32
|
||||
)
|
||||
|
||||
var gradPool = struct {
|
||||
sync.Mutex
|
||||
buckets map[poolKey][]*core.Array
|
||||
elems int
|
||||
}{buckets: make(map[poolKey][]*core.Array)}
|
||||
|
||||
// freshGrad allocates a zeroed array of k's dtype and shape, taking
|
||||
// ownership of the payload the way the constructor documents: grad
|
||||
// writes that payload only through the array's own raw accessor.
|
||||
func freshGrad(k poolKey, shape []int) *core.Array {
|
||||
var a *core.Array
|
||||
switch k.dt {
|
||||
case core.Float:
|
||||
a, _ = core.FloatsFromArray(make([]float64, k.n), shape...)
|
||||
case core.Float32:
|
||||
a, _ = core.FromFloat32Slice(make([]float32, k.n), shape...)
|
||||
default:
|
||||
a, _ = core.ComplexFromArray(make([]complex128, k.n), shape...)
|
||||
}
|
||||
if a == nil {
|
||||
a, _ = core.Zeros(k.dt, shape...)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
// clearGradPayload zeroes every slot of a's payload, the borrow-side
|
||||
// rule: a recycled buffer must never carry the previous sweep's values
|
||||
// into a reader.
|
||||
func clearGradPayload(a *core.Array) {
|
||||
switch a.Dtype() {
|
||||
case core.Float:
|
||||
clear(a.RawFloats())
|
||||
case core.Float32:
|
||||
clear(a.RawFloat32s())
|
||||
case core.Complex:
|
||||
clear(a.RawComplexes())
|
||||
}
|
||||
}
|
||||
|
||||
// releaseGrad offers a dead gradient array to the pool. The caller must
|
||||
// have proven the array unreachable: releasing one a tape, a graph or a
|
||||
// caller still holds would let the next borrower corrupt it. The shape
|
||||
// must be the shape the array was built with; a mismatch is caught by
|
||||
// the element-count check and drops the array instead of pooling it.
|
||||
func releaseGrad(a *core.Array, shape []int) {
|
||||
if a == nil || a.Strided() {
|
||||
return
|
||||
}
|
||||
k, ok := poolKeyOf(a.Dtype(), shape)
|
||||
if !ok || k.n != a.Len() {
|
||||
return
|
||||
}
|
||||
gradPool.Lock()
|
||||
b := gradPool.buckets[k]
|
||||
if len(b) >= gradPoolPerBucket || gradPool.elems+k.n > gradPoolMaxElems {
|
||||
gradPool.Unlock()
|
||||
return
|
||||
}
|
||||
gradPool.buckets[k] = append(b, a)
|
||||
gradPool.elems += k.n
|
||||
gradPool.Unlock()
|
||||
}
|
||||
|
||||
// gradArena is one sweep's private free list. A sweep borrows and
|
||||
// releases in near-LIFO order, so most round trips stay on the calling
|
||||
// goroutine under no lock; the global pool absorbs overflow and
|
||||
// supplies misses, and a sweep-end flush returns the leftovers under a
|
||||
// single lock. Every sweep owns its arena, so concurrent sweeps on
|
||||
// different graphs never share one.
|
||||
type gradArena struct {
|
||||
free []gradFree
|
||||
elems int
|
||||
pooled bool
|
||||
}
|
||||
|
||||
type gradFree struct {
|
||||
arr *core.Array
|
||||
sh []int
|
||||
k poolKey
|
||||
}
|
||||
|
||||
// gradArenaMaxFree bounds what one arena carries between flushes; a
|
||||
// sweep that releases beyond it hands the surplus to the global pool.
|
||||
const gradArenaMaxFree = 256
|
||||
|
||||
var gradArenaPool = sync.Pool{New: func() any { return &gradArena{} }}
|
||||
|
||||
func borrowArena() *gradArena {
|
||||
ar := gradArenaPool.Get().(*gradArena)
|
||||
ar.pooled = true
|
||||
// A recycled arena comes back with its previous free list: the
|
||||
// slice is emptied here, or stale entries would both starve the
|
||||
// scan and pin the buffers they still name.
|
||||
ar.free = ar.free[:0]
|
||||
ar.elems = 0
|
||||
return ar
|
||||
}
|
||||
|
||||
// borrowGrad returns a zeroed array of dt and shape: the arena's own
|
||||
// free list first, then the global pool, then fresh allocation. A nil
|
||||
// arena means the legacy sweep path, which allocates exactly what it
|
||||
// allocated before and takes no part in the pool.
|
||||
func (ar *gradArena) borrowGrad(dt core.Dtype, shape []int) *core.Array {
|
||||
if ar == nil {
|
||||
a, _ := core.Zeros(dt, shape...)
|
||||
return a
|
||||
}
|
||||
k, ok := poolKeyOf(dt, shape)
|
||||
if !ok {
|
||||
a, _ := core.Zeros(dt, shape...)
|
||||
return a
|
||||
}
|
||||
for i := len(ar.free) - 1; i >= 0; i-- {
|
||||
e := ar.free[i]
|
||||
if e.k != k {
|
||||
continue
|
||||
}
|
||||
ar.free[i] = ar.free[len(ar.free)-1]
|
||||
ar.free = ar.free[:len(ar.free)-1]
|
||||
ar.elems -= k.n
|
||||
clearGradPayload(e.arr)
|
||||
return e.arr
|
||||
}
|
||||
gradPool.Lock()
|
||||
b := gradPool.buckets[k]
|
||||
if len(b) > 0 {
|
||||
a := b[len(b)-1]
|
||||
gradPool.buckets[k] = b[:len(b)-1]
|
||||
gradPool.elems -= k.n
|
||||
gradPool.Unlock()
|
||||
clearGradPayload(a)
|
||||
return a
|
||||
}
|
||||
gradPool.Unlock()
|
||||
return freshGrad(k, shape)
|
||||
}
|
||||
|
||||
// releaseGrad returns a dead gradient array to the arena, falling
|
||||
// through to the global pool when the arena is full. A nil arena means
|
||||
// the caller owns the lifetime, so the array is left to the collector.
|
||||
func (ar *gradArena) releaseGrad(a *core.Array, shape []int) {
|
||||
if ar == nil || a == nil || a.Strided() {
|
||||
return
|
||||
}
|
||||
k, ok := poolKeyOf(a.Dtype(), shape)
|
||||
if !ok || k.n != a.Len() {
|
||||
return
|
||||
}
|
||||
if len(ar.free) >= gradArenaMaxFree || ar.elems+k.n > gradPoolMaxElems {
|
||||
releaseGrad(a, shape)
|
||||
return
|
||||
}
|
||||
ar.free = append(ar.free, gradFree{arr: a, sh: shape, k: k})
|
||||
ar.elems += k.n
|
||||
}
|
||||
|
||||
// flush returns everything the arena still holds to the global pool.
|
||||
// Buffers the pool declines are dropped to the collector; the arena
|
||||
// itself returns to the sync.Pool for the next sweep.
|
||||
func (ar *gradArena) flush() {
|
||||
if ar == nil {
|
||||
return
|
||||
}
|
||||
gradPool.Lock()
|
||||
for _, e := range ar.free {
|
||||
b := gradPool.buckets[e.k]
|
||||
if len(b) >= gradPoolPerBucket || gradPool.elems+e.k.n > gradPoolMaxElems {
|
||||
continue
|
||||
}
|
||||
gradPool.buckets[e.k] = append(b, e.arr)
|
||||
gradPool.elems += e.k.n
|
||||
}
|
||||
gradPool.Unlock()
|
||||
clear(ar.free)
|
||||
ar.free = ar.free[:0]
|
||||
ar.elems = 0
|
||||
if ar.pooled {
|
||||
gradArenaPool.Put(ar)
|
||||
}
|
||||
}
|
||||
|
||||
// fillGradSlotC writes z into every complex element of s.
|
||||
func fillGradSlotC(s gradSlot, z complex128) {
|
||||
if s.arr == nil {
|
||||
return
|
||||
}
|
||||
cs := s.arr.RawComplexes()[:s.arr.Len()]
|
||||
for i := range cs {
|
||||
cs[i] = z
|
||||
}
|
||||
}
|
||||
|
||||
// fillGradSlot writes v into every element of s, the seed and fill
|
||||
// helper. Each dtype takes the same spelling fillConst writes: a
|
||||
// float32 destination narrows the constant once and stores it, a
|
||||
// complex one stores complex(v, 0).
|
||||
func fillGradSlot(s gradSlot, v float64) {
|
||||
if s.arr == nil {
|
||||
return
|
||||
}
|
||||
n := s.arr.Len()
|
||||
switch s.arr.Dtype() {
|
||||
case core.Float:
|
||||
fs := s.arr.RawFloats()[:n]
|
||||
for i := range fs {
|
||||
fs[i] = v
|
||||
}
|
||||
case core.Float32:
|
||||
fs := s.arr.RawFloat32s()[:n]
|
||||
fv := float32(v)
|
||||
for i := range fs {
|
||||
fs[i] = fv
|
||||
}
|
||||
case core.Complex:
|
||||
cs := s.arr.RawComplexes()[:n]
|
||||
z := complex(v, 0)
|
||||
for i := range cs {
|
||||
cs[i] = z
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// tapeFrame is one node's position in the reverse sweep's explicit
|
||||
// walk: the node being expanded and the next operand index to visit.
|
||||
type tapeFrame struct {
|
||||
node *gradNode
|
||||
next int
|
||||
}
|
||||
|
||||
// tapeWork is the sweep's traversal scratch: the topological order, the
|
||||
// walk stack and the seen set. It is borrowed per sweep and returned
|
||||
// with its references cleared, so a pooled copy never pins a dead
|
||||
// graph; a workload whose graph exceeds the retention cap drops the
|
||||
// buffers to the collector instead of pinning them.
|
||||
type tapeWork struct {
|
||||
order []*gradNode
|
||||
stack []tapeFrame
|
||||
seen map[*Tensor]bool
|
||||
}
|
||||
|
||||
const gradPoolMaxTapeNodes = 1 << 16
|
||||
|
||||
var tapeWorkPool = sync.Pool{New: func() any {
|
||||
return &tapeWork{seen: make(map[*Tensor]bool)}
|
||||
}}
|
||||
|
||||
func borrowTapeWork() *tapeWork {
|
||||
w := tapeWorkPool.Get().(*tapeWork)
|
||||
w.order = w.order[:0]
|
||||
w.stack = w.stack[:0]
|
||||
clear(w.seen)
|
||||
if cap(w.order) > gradPoolMaxTapeNodes || cap(w.stack) > gradPoolMaxTapeNodes {
|
||||
return &tapeWork{seen: make(map[*Tensor]bool)}
|
||||
}
|
||||
return w
|
||||
}
|
||||
|
||||
func releaseTapeWork(w *tapeWork) {
|
||||
if w == nil {
|
||||
return
|
||||
}
|
||||
clear(w.order[:cap(w.order)])
|
||||
clear(w.stack[:cap(w.stack)])
|
||||
clear(w.seen)
|
||||
if cap(w.order) > gradPoolMaxTapeNodes || cap(w.stack) > gradPoolMaxTapeNodes {
|
||||
return
|
||||
}
|
||||
tapeWorkPool.Put(w)
|
||||
}
|
||||
@@ -0,0 +1,200 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"runtime"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// The gradient pool's contract, measured: a repeated backward sweep on
|
||||
// one process must keep the heap flat rather than growing with the
|
||||
// iteration count, the sweep's arithmetic must be bit-for-bit
|
||||
// reproducible across runs that share the pool, and the per-sweep cost
|
||||
// itself is pinned by benchmarks that separate graph construction from
|
||||
// the reverse pass.
|
||||
|
||||
// tapeChain builds a chain of n element-wise nodes over x and w and
|
||||
// reduces it to a scalar, the fixture the sweep benchmarks repeat.
|
||||
func tapeChain(t testing.TB, x, w *Tensor, n int) *Tensor {
|
||||
t.Helper()
|
||||
h := x
|
||||
for i := range n {
|
||||
var err error
|
||||
switch i % 4 {
|
||||
case 0:
|
||||
h, err = h.Add(w)
|
||||
case 1:
|
||||
h, err = h.Mul(w)
|
||||
case 2:
|
||||
h, err = h.Tanh()
|
||||
default:
|
||||
h, err = h.Scale(0.25)
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
s, err := h.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func tapeLeaf(t testing.TB, seed, n int) *Tensor {
|
||||
t.Helper()
|
||||
v := make([]float64, n)
|
||||
for i := range v {
|
||||
v[i] = 0.25 + float64((i*7+seed)%13)*0.125
|
||||
}
|
||||
a, err := core.FromFloats(v, n)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return FromArray(a, true)
|
||||
}
|
||||
|
||||
// TestGradPoolFlatHeap runs one repeated backward workload and checks
|
||||
// the live heap stops growing: the pool's retention caps must bound
|
||||
// what one process holds, however many sweeps it serves.
|
||||
func TestGradPoolFlatHeap(t *testing.T) {
|
||||
x, w := tapeLeaf(t, 1, 8), tapeLeaf(t, 2, 8)
|
||||
s := tapeChain(t, x, w, 128)
|
||||
run := func() {
|
||||
for range 500 {
|
||||
x.ZeroGrad()
|
||||
w.ZeroGrad()
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
var early, late runtime.MemStats
|
||||
runtime.GC()
|
||||
run()
|
||||
runtime.GC()
|
||||
runtime.ReadMemStats(&early)
|
||||
run()
|
||||
run()
|
||||
runtime.GC()
|
||||
runtime.ReadMemStats(&late)
|
||||
// Two more batches of a thousand sweeps may add pool slack but not
|
||||
// a growth trend: the second reading stays within a small factor of
|
||||
// the first, which a leaking pool would break.
|
||||
if late.HeapInuse > early.HeapInuse*2+1<<20 {
|
||||
t.Fatalf("heap grew across repeated sweeps: %d then %d bytes in use", early.HeapInuse, late.HeapInuse)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGradBackwardDeterminismBits runs the same program twice through
|
||||
// the pooled sweep and demands identical gradient bits: recycling a
|
||||
// buffer must never leak a previous sweep's values into a result.
|
||||
func TestGradBackwardDeterminismBits(t *testing.T) {
|
||||
gradOf := func() []float64 {
|
||||
x, w := tapeLeaf(t, 3, 8), tapeLeaf(t, 4, 8)
|
||||
s := tapeChain(t, x, w, 64)
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
gx, gw := x.Grad(), w.Grad()
|
||||
if gx == nil || gw == nil {
|
||||
t.Fatal("missing leaf gradient")
|
||||
}
|
||||
out := make([]float64, 0, gx.Len()+gw.Len())
|
||||
out = append(out, gx.RawFloats()[:gx.Len()]...)
|
||||
out = append(out, gw.RawFloats()[:gw.Len()]...)
|
||||
return out
|
||||
}
|
||||
// Warm the pool with unrelated sweeps, so the measured runs borrow
|
||||
// recycled buffers carrying other work's values.
|
||||
for range 64 {
|
||||
a, b := tapeLeaf(t, 9, 8), tapeLeaf(t, 10, 8)
|
||||
s := tapeChain(t, a, b, 32)
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
first, second := gradOf(), gradOf()
|
||||
if len(first) != len(second) {
|
||||
t.Fatalf("gradient lengths differ: %d and %d", len(first), len(second))
|
||||
}
|
||||
for i := range first {
|
||||
if math.Float64bits(first[i]) != math.Float64bits(second[i]) {
|
||||
t.Fatalf("gradient bit %d differs: %x and %x", i,
|
||||
math.Float64bits(first[i]), math.Float64bits(second[i]))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkTapeChainSweepBackward measures the reverse sweep alone
|
||||
// on a 129-node chain that is built once: every allocation here is the
|
||||
// sweep's own, not the graph's.
|
||||
func BenchmarkTapeChainSweepBackward(b *testing.B) {
|
||||
x, w := tapeLeaf(b, 1, 8), tapeLeaf(b, 2, 8)
|
||||
s := tapeChain(b, x, w, 128)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
x.ZeroGrad()
|
||||
w.ZeroGrad()
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkWaveTapeChainForwardRebuild measures building the same chain
|
||||
// afresh with no backward pass, the per-node graph-construction cost
|
||||
// the sweep benchmarks otherwise carry inside their loop.
|
||||
func BenchmarkWaveTapeChainForwardRebuild(b *testing.B) {
|
||||
x, w := tapeLeaf(b, 1, 8), tapeLeaf(b, 2, 8)
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
s := tapeChain(b, x, w, 128)
|
||||
if s.Data().Len() != 1 {
|
||||
b.Fatal("unexpected shape")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkWaveWideFanSweepBackward measures the reverse sweep of a
|
||||
// 64-way fan over one shared leaf: the fold-heavy edge pattern, built
|
||||
// once.
|
||||
func BenchmarkWaveWideFanSweepBackward(b *testing.B) {
|
||||
x := tapeLeaf(b, 3, 16)
|
||||
leaves := make([]*Tensor, 64)
|
||||
for i := range leaves {
|
||||
leaves[i] = tapeLeaf(b, 10+i, 16)
|
||||
}
|
||||
acc, err := x.Mul(leaves[0])
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
for _, l := range leaves[1:] {
|
||||
p, err := x.Mul(l)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if acc, err = acc.Add(p); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
s, err := acc.Sum()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
x.ZeroGrad()
|
||||
for _, l := range leaves {
|
||||
l.ZeroGrad()
|
||||
}
|
||||
if err := s.Backward(); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,164 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// Pins for the pooled reverse sweep: the pooled Backward must answer
|
||||
// bit-identically to the legacy map-returning sweep on the same graph,
|
||||
// concurrent sweeps on separate graphs must answer the serial reference
|
||||
// bits, and the pool's retention caps must refuse releases past them.
|
||||
|
||||
// pinLit builds a leaf of n elements from fixed literals, the fixture
|
||||
// shape the tape benchmarks use.
|
||||
func pinLit(seed, n int) *Tensor {
|
||||
v := make([]float64, n)
|
||||
for i := range v {
|
||||
v[i] = 0.25 + float64((i*7+seed)%13)*0.125
|
||||
}
|
||||
a, err := core.FromFloats(v, n)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return FromArray(a, true)
|
||||
}
|
||||
|
||||
// chainLeafBits builds the same deep chain the tape benchmarks build,
|
||||
// runs one pooled Backward and returns the two leaf gradients' raw
|
||||
// bits. The chain fans both leaves into every node, so every fold
|
||||
// accumulates multiple contributions.
|
||||
func chainLeafBits(seedA, seedB, nodes int) ([]float64, error) {
|
||||
x, w := pinLit(seedA, 8), pinLit(seedB, 8)
|
||||
s, err := deepChain(x, w, nodes)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
x.ZeroGrad()
|
||||
w.ZeroGrad()
|
||||
if err := s.Backward(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
xg, wg := x.Grad(), w.Grad()
|
||||
if xg == nil || wg == nil {
|
||||
return nil, errf("pinned chain: missing leaf gradient")
|
||||
}
|
||||
out := make([]float64, 0, 16)
|
||||
out = append(out, xg.RawFloats()[:8]...)
|
||||
out = append(out, wg.RawFloats()[:8]...)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func TestPooledBackwardBitsMatchLegacySweep(t *testing.T) {
|
||||
pooled, err := chainLeafBits(1, 2, 64)
|
||||
if err != nil {
|
||||
t.Fatalf("pooled sweep: %v", err)
|
||||
}
|
||||
// The same graph through the legacy sweep: reverseGrads commits
|
||||
// nothing, so its map carries the leaves' gradients from this pass
|
||||
// alone, which is what the pooled sweep commits on a fresh leaf.
|
||||
x, w := pinLit(1, 8), pinLit(2, 8)
|
||||
s, err := deepChain(x, w, 64)
|
||||
if err != nil {
|
||||
t.Fatalf("legacy chain: %v", err)
|
||||
}
|
||||
grads, err := s.reverseGrads()
|
||||
if err != nil {
|
||||
t.Fatalf("legacy sweep: %v", err)
|
||||
}
|
||||
gx, gw := grads[x], grads[w]
|
||||
if gx == nil || gw == nil {
|
||||
t.Fatal("legacy sweep returned no leaf gradient")
|
||||
}
|
||||
legacy := append(append([]float64{}, gx.RawFloats()[:8]...), gw.RawFloats()[:8]...)
|
||||
if len(legacy) != len(pooled) {
|
||||
t.Fatalf("length %d, want %d", len(pooled), len(legacy))
|
||||
}
|
||||
for i := range pooled {
|
||||
if math.Float64bits(pooled[i]) != math.Float64bits(legacy[i]) {
|
||||
t.Fatalf("leaf gradient %d: pooled %#x, legacy %#x",
|
||||
i, math.Float64bits(pooled[i]), math.Float64bits(legacy[i]))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcurrentBackwardDeterminism(t *testing.T) {
|
||||
ref, err := chainLeafBits(3, 4, 48)
|
||||
if err != nil {
|
||||
t.Fatalf("serial reference: %v", err)
|
||||
}
|
||||
const sweeps = 40
|
||||
outs := make([][]float64, 2)
|
||||
errs := make([]error, 2)
|
||||
var wg sync.WaitGroup
|
||||
for g := range 2 {
|
||||
wg.Go(func() {
|
||||
for range sweeps {
|
||||
bits, err := chainLeafBits(3, 4, 48)
|
||||
if err != nil {
|
||||
errs[g] = err
|
||||
return
|
||||
}
|
||||
outs[g] = bits
|
||||
}
|
||||
})
|
||||
}
|
||||
wg.Wait()
|
||||
for g := range 2 {
|
||||
if errs[g] != nil {
|
||||
t.Fatalf("goroutine %d: %v", g, errs[g])
|
||||
}
|
||||
for i := range ref {
|
||||
if math.Float64bits(outs[g][i]) != math.Float64bits(ref[i]) {
|
||||
t.Fatalf("goroutine %d element %d: %#x, want %#x",
|
||||
g, i, math.Float64bits(outs[g][i]), math.Float64bits(ref[i]))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGradPoolCapsAreEnforced(t *testing.T) {
|
||||
k, ok := poolKeyOf(core.Float, []int{8})
|
||||
if !ok {
|
||||
t.Fatal("poolKeyOf refused a float shape of 8")
|
||||
}
|
||||
keep, err := core.Zeros(core.Float, 8)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Bucket cap: a full bucket refuses the next release even when the
|
||||
// pool's element budget has room.
|
||||
gradPool.Lock()
|
||||
savedB, savedE := gradPool.buckets[k], gradPool.elems
|
||||
gradPool.buckets[k] = make([]*core.Array, gradPoolPerBucket)
|
||||
gradPool.elems = gradPoolMaxElems - 16
|
||||
gradPool.Unlock()
|
||||
releaseGrad(keep, []int{8})
|
||||
gradPool.Lock()
|
||||
gotLen := len(gradPool.buckets[k])
|
||||
gradPool.buckets[k], gradPool.elems = savedB, savedE
|
||||
gradPool.Unlock()
|
||||
if gotLen != gradPoolPerBucket {
|
||||
t.Fatalf("bucket accepted a release past its cap: %d entries, cap %d", gotLen, gradPoolPerBucket)
|
||||
}
|
||||
// Element cap: a full pool refuses the next release however empty
|
||||
// the bucket is.
|
||||
gradPool.Lock()
|
||||
gradPool.buckets[k] = nil
|
||||
gradPool.elems = gradPoolMaxElems
|
||||
gradPool.Unlock()
|
||||
releaseGrad(keep, []int{8})
|
||||
gradPool.Lock()
|
||||
gotElems := gradPool.elems
|
||||
gradPool.buckets[k], gradPool.elems = savedB, savedE
|
||||
gradPool.Unlock()
|
||||
if gotElems != gradPoolMaxElems {
|
||||
t.Fatalf("pool accepted a release past its element cap: %d, cap %d", gotElems, gradPoolMaxElems)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import "sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
|
||||
// Slice extracts a range along the given dimension as a new tensor;
|
||||
// the backward writes the incoming gradient into the corresponding
|
||||
// region of the original shape.
|
||||
func (t *Tensor) Slice(dim, start, stop int) (*Tensor, error) {
|
||||
if err := t.checkDiff("Slice"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out, err := core.Slice(t.data, dim, start, stop)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
orig := t.data.Shape()
|
||||
dt := t.data.Dtype()
|
||||
return t.unaryResult("Slice", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
// The narrowing Concat applies: a complex gradient reaching a
|
||||
// real slice narrows by 2·Re before the span is copied. A
|
||||
// slice's output dtype equals its input's, so only a complex
|
||||
// gradient on a real tensor can differ here.
|
||||
gn, err := narrowGradient(g, dt)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
da := gradSlot{arr: ar.borrowGrad(dt, orig), sh: orig}
|
||||
outer := 1
|
||||
for d := range dim {
|
||||
outer *= orig[d]
|
||||
}
|
||||
inner := 1
|
||||
for d := dim + 1; d < len(orig); d++ {
|
||||
inner *= orig[d]
|
||||
}
|
||||
nS := stop - start
|
||||
// Each kept row is one contiguous inner run, so matching dtypes
|
||||
// ride raw slice moves instead of per-element accessor calls.
|
||||
fast := !gn.arr.Strided() && gn.arr.Dtype() == dt && dt != core.Int
|
||||
for o := range outer {
|
||||
for si := range nS {
|
||||
d := o*orig[dim]*inner + (start+si)*inner
|
||||
s := o*nS*inner + si*inner
|
||||
if fast {
|
||||
copySegRaw(da.arr, gn.arr, d, s, inner)
|
||||
continue
|
||||
}
|
||||
for j := range inner {
|
||||
copyElem(da.arr, d+j, gn.arr, s+j)
|
||||
}
|
||||
}
|
||||
}
|
||||
dst[0] = da
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// Reshape returns a new view-equivalent tensor of the given shape; the
|
||||
// backward simply reshapes the incoming gradient back.
|
||||
func (t *Tensor) Reshape(shape ...int) (*Tensor, error) {
|
||||
if err := t.checkDiff("Reshape"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out, err := core.Reshape(t.data, shape...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
orig := append([]int{}, t.data.Shape()...)
|
||||
return t.unaryResult("Reshape", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
gr, err := core.Reshape(g.arr, orig...)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dst[0] = gradSlot{arr: gr, sh: orig}
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
@@ -0,0 +1,246 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"sourcedock.dev/petrbalvin/tensor/signal"
|
||||
)
|
||||
|
||||
// Spectral autograd: the Fourier transforms as graph nodes,
|
||||
// so deconvolution, spectral de-noising and frequency-domain fitting
|
||||
// differentiate end to end. The adjoint of the unnormalised forward
|
||||
// DFT y = F·z is dz = Fᴴ·g = n·IFFT(g) in the Wirtinger convention
|
||||
// (the conjugate transpose falls out of dz = 2Re[ḡᵀ·dy] exactly the
|
||||
// way the MatMul adjoint does); the inverse transform is its mirror.
|
||||
// Real inputs flow through unchanged: signal.FFT widens them to
|
||||
// complex, and the engine's complex-to-real narrowing (2·Re) is
|
||||
// precisely the adjoint of that widening.
|
||||
|
||||
// FFT is the forward discrete Fourier transform of a rank-1 tensor;
|
||||
// the backward multiplies the incoming gradient by Fᴴ, which is the
|
||||
// inverse transform scaled by n.
|
||||
func (t *Tensor) FFT() (*Tensor, error) {
|
||||
if err := t.checkDiff("FFT"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if t.data.NDim() != 1 {
|
||||
return nil, errf("autograd FFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
|
||||
}
|
||||
out, err := signal.FFT(t.data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scale := complex(float64(t.data.Len()), 0)
|
||||
sh := t.data.Shape()
|
||||
return t.unaryResult("FFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
inv, err := signal.IFFT(g.arr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cs := inv.RawComplexes()
|
||||
for i := range cs {
|
||||
cs[i] *= scale
|
||||
}
|
||||
dst[0] = gradSlot{arr: inv, sh: sh}
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// IFFT is the inverse transform of a rank-1 tensor; the backward runs
|
||||
// the forward transform scaled by 1/n.
|
||||
func (t *Tensor) IFFT() (*Tensor, error) {
|
||||
if err := t.checkDiff("IFFT"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if t.data.NDim() != 1 {
|
||||
return nil, errf("autograd IFFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
|
||||
}
|
||||
out, err := signal.IFFT(t.data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scale := complex(1/float64(t.data.Len()), 0)
|
||||
sh := t.data.Shape()
|
||||
return t.unaryResult("IFFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
fwd, err := signal.FFT(g.arr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cs := fwd.RawComplexes()
|
||||
for i := range cs {
|
||||
cs[i] *= scale
|
||||
}
|
||||
dst[0] = gradSlot{arr: fwd, sh: sh}
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// FFT2 is the 2-D forward transform; the backward is the 2-D inverse
|
||||
// scaled by H·W, the total element count.
|
||||
func (t *Tensor) FFT2() (*Tensor, error) {
|
||||
if err := t.checkDiff("FFT2"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if t.data.NDim() != 2 {
|
||||
return nil, errf("autograd FFT2: needs a rank-2 tensor, got shape %s", prettyShape(t.data.Shape()))
|
||||
}
|
||||
out, err := signal.FFT2(t.data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scale := complex(float64(t.data.Len()), 0)
|
||||
sh := t.data.Shape()
|
||||
return t.unaryResult("FFT2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
inv, err := signal.IFFT2(g.arr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cs := inv.RawComplexes()
|
||||
for i := range cs {
|
||||
cs[i] *= scale
|
||||
}
|
||||
dst[0] = gradSlot{arr: inv, sh: sh}
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// IFFT2 is the 2-D inverse transform; the backward is the 2-D forward
|
||||
// scaled by 1/(H·W).
|
||||
func (t *Tensor) IFFT2() (*Tensor, error) {
|
||||
if err := t.checkDiff("IFFT2"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if t.data.NDim() != 2 {
|
||||
return nil, errf("autograd IFFT2: needs a rank-2 tensor, got shape %s", prettyShape(t.data.Shape()))
|
||||
}
|
||||
out, err := signal.IFFT2(t.data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scale := complex(1/float64(t.data.Len()), 0)
|
||||
sh := t.data.Shape()
|
||||
return t.unaryResult("IFFT2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
fwd, err := signal.FFT2(g.arr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cs := fwd.RawComplexes()
|
||||
for i := range cs {
|
||||
cs[i] *= scale
|
||||
}
|
||||
dst[0] = gradSlot{arr: fwd, sh: sh}
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// RFFT is the real-input half-spectrum transform. The input must be a
|
||||
// rank-1 real tensor; the backward folds the incoming half-spectrum
|
||||
// gradient into dx = 2·Re(F_halfᴴ·g), evaluated by one padded forward
|
||||
// FFT so the cost matches the forward transform. The factor 2 lands
|
||||
// only on the mirrored bins through the zero padding, exactly the
|
||||
// combinatorics the derivation gives.
|
||||
func (t *Tensor) RFFT() (*Tensor, error) {
|
||||
if err := t.checkDiff("RFFT"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if isComplexArr(t.data) {
|
||||
return nil, errf("autograd RFFT: needs a real tensor, got %s", t.data.Dtype())
|
||||
}
|
||||
if t.data.NDim() != 1 {
|
||||
return nil, errf("autograd RFFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
|
||||
}
|
||||
out, err := signal.RFFT(t.data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
in := t.data
|
||||
return t.unaryResult("RFFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
n := in.Len()
|
||||
half := n/2 + 1
|
||||
// The adjoint needs Σ_{k<half} g_k·e^{+2πijk/n}, a +sign DFT
|
||||
// of the zero-padded gradient: conj(FFT(conj(·))).
|
||||
pad := make([]complex128, n)
|
||||
for k := range half {
|
||||
pad[k] = conj(g.arr.ComplexAt(k))
|
||||
}
|
||||
padArr, err := core.ComplexFromArray(pad, n)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
spec, err := signal.FFT(padArr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sh := in.Shape()
|
||||
dx := gradSlot{arr: ar.borrowGrad(in.Dtype(), sh), sh: sh}
|
||||
// The transform's output and dx are both freshly allocated and
|
||||
// dense, so the doubled real part is taken from the payload.
|
||||
ss := spec.RawComplexes()
|
||||
if in.Dtype() == core.Float32 {
|
||||
ds := dx.arr.RawFloat32s()[:dx.arr.Len()]
|
||||
for j := range n {
|
||||
ds[j] = float32(2 * real(ss[j]))
|
||||
}
|
||||
} else {
|
||||
ds := dx.arr.RawFloats()[:dx.arr.Len()]
|
||||
for j := range n {
|
||||
ds[j] = 2 * real(ss[j])
|
||||
}
|
||||
}
|
||||
dst[0] = dx
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// IRFFT is the inverse half-spectrum transform: a rank-1 complex
|
||||
// tensor of n/2+1 bins into a real signal of length n. The backward
|
||||
// widens the real gradient to a full forward FFT and halves the
|
||||
// self-mirrored bins (DC, and Nyquist when n is even): dIn_k =
|
||||
// FFT(g)_k/n on the ordinary bins and half of that on the mirrored
|
||||
// bins, the transpose of the Hermitian extension the forward performs.
|
||||
func (t *Tensor) IRFFT(n int) (*Tensor, error) {
|
||||
if !isComplexArr(t.data) {
|
||||
return nil, errf("autograd IRFFT: needs a complex tensor, got %s", t.data.Dtype())
|
||||
}
|
||||
if t.data.NDim() != 1 {
|
||||
return nil, errf("autograd IRFFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
|
||||
}
|
||||
out, err := signal.IRFFT(t.data, n)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
in := t.data
|
||||
return t.unaryResult("IRFFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
gs := make([]complex128, n)
|
||||
for j := range n {
|
||||
gs[j] = complex(g.arr.FloatAt(j), 0)
|
||||
}
|
||||
gArr, err := core.ComplexFromArray(gs, n)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
spec, err := signal.FFT(gArr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
half := in.Len()
|
||||
full := complex(float64(n), 0)
|
||||
sh := []int{half}
|
||||
dx := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh}
|
||||
ds := dx.arr.RawComplexes()[:half]
|
||||
ss := spec.RawComplexes()
|
||||
ds[0] = ss[0] / (2 * full)
|
||||
for k := 1; k < half; k++ {
|
||||
if n%2 == 0 && k == half-1 {
|
||||
// Nyquist mirrors itself.
|
||||
ds[k] = ss[k] / (2 * full)
|
||||
continue
|
||||
}
|
||||
ds[k] = ss[k] / full
|
||||
}
|
||||
dst[0] = dx
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
@@ -0,0 +1,370 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// spectralWeights builds a deterministic complex weight vector used to
|
||||
// fold a spectrum into a real scalar loss.
|
||||
func spectralWeights(n int, seed int) []complex128 {
|
||||
g := core.NewGenerator(int64(seed))
|
||||
w := make([]complex128, n)
|
||||
for i := range n {
|
||||
w[i] = complex(g.NormalUnit(), g.NormalUnit())
|
||||
}
|
||||
return w
|
||||
}
|
||||
|
||||
// complexLossOf runs the op chain and folds the result into the real
|
||||
// scalar Σ Re(w·y) + Σ|y|²/len, the same fold the graph's foldReal
|
||||
// builds from Mul and Real, so oracle and graph define one loss.
|
||||
func complexLossOf(t *testing.T, op func(*Tensor) (*Tensor, error), w []complex128) func(*core.Array) float64 {
|
||||
t.Helper()
|
||||
return func(a *core.Array) float64 {
|
||||
y, err := op(FromArray(a, false))
|
||||
if err != nil {
|
||||
t.Fatalf("forward: %v", err)
|
||||
}
|
||||
s := 0.0
|
||||
for i := range y.Data().Len() {
|
||||
z := y.Data().ComplexAt(i)
|
||||
s += real(w[i%len(w)] * z)
|
||||
s += (real(z)*real(z) + imag(z)*imag(z)) / float64(y.Data().Len())
|
||||
}
|
||||
return s
|
||||
}
|
||||
}
|
||||
|
||||
// realLossOf is complexLossOf for chains that end in a real tensor.
|
||||
func realLossOf(t *testing.T, op func(*Tensor) (*Tensor, error), w []complex128) func(*core.Array) float64 {
|
||||
t.Helper()
|
||||
return func(a *core.Array) float64 {
|
||||
y, err := op(FromArray(a, false))
|
||||
if err != nil {
|
||||
t.Fatalf("forward: %v", err)
|
||||
}
|
||||
s := 0.0
|
||||
for i := range y.Data().Len() {
|
||||
v := y.Data().FloatAt(i)
|
||||
s += real(w[i%len(w)])*v + v*v/float64(y.Data().Len())
|
||||
}
|
||||
return s
|
||||
}
|
||||
}
|
||||
|
||||
// backwardComplex runs op on a fresh tensor over vals and returns the
|
||||
// leaf after Backward.
|
||||
func backwardComplex(t *testing.T, op func(*Tensor) (*Tensor, error), vals []complex128, shape ...int) *Tensor {
|
||||
t.Helper()
|
||||
a, err := core.FromComplexes(vals, shape...)
|
||||
if err != nil {
|
||||
t.Fatalf("FromComplexes: %v", err)
|
||||
}
|
||||
xt := FromArray(a, true)
|
||||
y, err := op(xt)
|
||||
if err != nil {
|
||||
t.Fatalf("forward: %v", err)
|
||||
}
|
||||
// Fold to a real scalar so Backward has its seed.
|
||||
w := spectralWeights(y.Data().Len(), 11)
|
||||
var loss *Tensor
|
||||
loss, err = foldReal(t, y, w)
|
||||
if err != nil {
|
||||
t.Fatalf("fold: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
return xt
|
||||
}
|
||||
|
||||
// foldReal reduces a complex tensor to Σ Re(w̄·y) + Σ|y|²/len through
|
||||
// the graph ops, and a real tensor to the analogous real fold.
|
||||
func foldReal(t *testing.T, y *Tensor, w []complex128) (*Tensor, error) {
|
||||
t.Helper()
|
||||
n := y.Data().Len()
|
||||
wa, err := core.FromComplexes(w, y.Data().Shape()...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
wt := FromArray(wa, false)
|
||||
if isComplexArr(y.Data()) {
|
||||
prod, err := y.Mul(wt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
re, err := prod.Real()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s1, err := re.Sum()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := y.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s2, err := sq.Sum()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s2s, err := s2.Scale(1 / float64(n))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s1.Add(s2s)
|
||||
}
|
||||
prod, err := y.Mul(wt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
re, err := prod.Real()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s1, err := re.Sum()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sq, err := y.Abs2()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s2, err := sq.Sum()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s2s, err := s2.Scale(1 / float64(n))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s1.Add(s2s)
|
||||
}
|
||||
|
||||
// TestGradFFTWirtinger pins the FFT adjoint against central
|
||||
// differences on a power-of-two and a Bluestein length.
|
||||
func TestGradFFTWirtinger(t *testing.T) {
|
||||
for _, n := range []int{8, 12} {
|
||||
g := core.NewGenerator(int64(n))
|
||||
vals := make([]complex128, n)
|
||||
for i := range n {
|
||||
vals[i] = complex(g.NormalUnit(), g.NormalUnit())
|
||||
}
|
||||
op := func(x *Tensor) (*Tensor, error) { return x.FFT() }
|
||||
xt := backwardComplex(t, op, vals, n)
|
||||
loss := complexLossOf(t, op, spectralWeights(n, 11))
|
||||
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGradIFFTWirtinger pins the IFFT adjoint.
|
||||
func TestGradIFFTWirtinger(t *testing.T) {
|
||||
n := 10
|
||||
g := core.NewGenerator(3)
|
||||
vals := make([]complex128, n)
|
||||
for i := range n {
|
||||
vals[i] = complex(g.NormalUnit(), g.NormalUnit())
|
||||
}
|
||||
op := func(x *Tensor) (*Tensor, error) { return x.IFFT() }
|
||||
xt := backwardComplex(t, op, vals, n)
|
||||
loss := complexLossOf(t, op, spectralWeights(n, 11))
|
||||
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7)
|
||||
}
|
||||
|
||||
// TestGradFFT2Wirtinger pins the 2-D adjoint.
|
||||
func TestGradFFT2Wirtinger(t *testing.T) {
|
||||
rows, cols := 3, 4
|
||||
g := core.NewGenerator(5)
|
||||
vals := make([]complex128, rows*cols)
|
||||
for i := range vals {
|
||||
vals[i] = complex(g.NormalUnit(), g.NormalUnit())
|
||||
}
|
||||
op := func(x *Tensor) (*Tensor, error) { return x.FFT2() }
|
||||
xt := backwardComplex(t, op, vals, rows, cols)
|
||||
loss := complexLossOf(t, op, spectralWeights(rows*cols, 11))
|
||||
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7)
|
||||
}
|
||||
|
||||
// TestGradFFTRealInput pins the 2·Re narrowing path: a real leaf under
|
||||
// a complex FFT node.
|
||||
func TestGradFFTRealInput(t *testing.T) {
|
||||
n := 8
|
||||
g := core.NewGenerator(9)
|
||||
vals := make([]float64, n)
|
||||
for i := range n {
|
||||
vals[i] = g.NormalUnit()
|
||||
}
|
||||
op := func(x *Tensor) (*Tensor, error) { return x.FFT() }
|
||||
a, err := core.FromFloats(vals, n)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
xt := FromArray(a, true)
|
||||
y, err := op(xt)
|
||||
if err != nil {
|
||||
t.Fatalf("forward: %v", err)
|
||||
}
|
||||
loss, err := foldReal(t, y, spectralWeights(n, 11))
|
||||
if err != nil {
|
||||
t.Fatalf("fold: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
lossOf := realToComplexLoss(t, op, n)
|
||||
want := numericGrad(lossOf, a)
|
||||
for i := range n {
|
||||
if math.Abs(xt.Grad().FloatAt(i)-want[i]) > 1e-6 {
|
||||
t.Fatalf("grad[%d] = %g, want %g", i, xt.Grad().FloatAt(i), want[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// realToComplexLoss adapts a chain over real input for numericGrad: it
|
||||
// re-runs the forward and folds the (complex) output into a scalar.
|
||||
func realToComplexLoss(t *testing.T, op func(*Tensor) (*Tensor, error), n int) func(*core.Array) float64 {
|
||||
t.Helper()
|
||||
w := spectralWeights(n, 11)
|
||||
return func(a *core.Array) float64 {
|
||||
y, err := op(FromArray(a, false))
|
||||
if err != nil {
|
||||
t.Fatalf("forward: %v", err)
|
||||
}
|
||||
s := 0.0
|
||||
for i := range n {
|
||||
z := y.Data().ComplexAt(i)
|
||||
s += real(w[i] * z)
|
||||
s += (real(z)*real(z) + imag(z)*imag(z)) / float64(n)
|
||||
}
|
||||
return s
|
||||
}
|
||||
}
|
||||
|
||||
// TestGradRFFT pins the half-spectrum adjoint, even and odd lengths.
|
||||
func TestGradRFFT(t *testing.T) {
|
||||
for _, n := range []int{8, 9} {
|
||||
g := core.NewGenerator(int64(n * 2))
|
||||
vals := make([]float64, n)
|
||||
for i := range n {
|
||||
vals[i] = g.NormalUnit()
|
||||
}
|
||||
a, err := core.FromFloats(vals, n)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloats: %v", err)
|
||||
}
|
||||
xt := FromArray(a, true)
|
||||
y, err := xt.RFFT()
|
||||
if err != nil {
|
||||
t.Fatalf("RFFT: %v", err)
|
||||
}
|
||||
w := spectralWeights(y.Data().Len(), 11)
|
||||
var loss *Tensor
|
||||
loss, err = foldReal(t, y, w)
|
||||
if err != nil {
|
||||
t.Fatalf("fold: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
lossOf := func(a *core.Array) float64 {
|
||||
yp, err := FromArray(a, false).RFFT()
|
||||
if err != nil {
|
||||
t.Fatalf("RFFT: %v", err)
|
||||
}
|
||||
s := 0.0
|
||||
for i := range yp.Data().Len() {
|
||||
z := yp.Data().ComplexAt(i)
|
||||
s += real(w[i] * z)
|
||||
s += (real(z)*real(z) + imag(z)*imag(z)) / float64(yp.Data().Len())
|
||||
}
|
||||
return s
|
||||
}
|
||||
want := numericGrad(lossOf, a)
|
||||
for i := range n {
|
||||
if math.Abs(xt.Grad().FloatAt(i)-want[i]) > 1e-6 {
|
||||
t.Fatalf("n=%d grad[%d] = %g, want %g", n, i, xt.Grad().FloatAt(i), want[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestGradIRFFT pins the half-spectrum inverse adjoint, on an even
|
||||
// length (Nyquist half-weight bin) and an odd one (the last bin is an
|
||||
// ordinary mirrored bin with the full 1/n weight).
|
||||
func TestGradIRFFT(t *testing.T) {
|
||||
for _, n := range []int{8, 9} {
|
||||
half := n/2 + 1
|
||||
g := core.NewGenerator(21)
|
||||
vals := make([]complex128, half)
|
||||
for i := range half {
|
||||
vals[i] = complex(g.NormalUnit(), g.NormalUnit())
|
||||
}
|
||||
op := func(x *Tensor) (*Tensor, error) { return x.IRFFT(n) }
|
||||
xt := backwardComplex(t, op, vals, half)
|
||||
// numericGrad over complex perturbations, folded through the real
|
||||
// output.
|
||||
w := spectralWeights(n, 11)
|
||||
loss := func(a *core.Array) float64 {
|
||||
y, err := FromArray(a, false).IRFFT(n)
|
||||
if err != nil {
|
||||
t.Fatalf("IRFFT: %v", err)
|
||||
}
|
||||
s := 0.0
|
||||
for i := range n {
|
||||
v := y.Data().FloatAt(i)
|
||||
s += real(w[i])*v + v*v/float64(n)
|
||||
}
|
||||
return s
|
||||
}
|
||||
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGradFFTRoundtripIdentity pins the composition: gradient through
|
||||
// IFFT∘FFT must arrive unchanged (Fᴴ·(1/n)F = I).
|
||||
func TestGradFFTRoundtripIdentity(t *testing.T) {
|
||||
n := 8
|
||||
vals := make([]complex128, n)
|
||||
g := core.NewGenerator(4)
|
||||
for i := range n {
|
||||
vals[i] = complex(g.NormalUnit(), g.NormalUnit())
|
||||
}
|
||||
a, err := core.FromComplexes(vals, n)
|
||||
if err != nil {
|
||||
t.Fatalf("FromComplexes: %v", err)
|
||||
}
|
||||
xt := FromArray(a, true)
|
||||
f, err := xt.FFT()
|
||||
if err != nil {
|
||||
t.Fatalf("FFT: %v", err)
|
||||
}
|
||||
fi, err := f.IFFT()
|
||||
if err != nil {
|
||||
t.Fatalf("IFFT: %v", err)
|
||||
}
|
||||
w := spectralWeights(n, 7)
|
||||
loss, err := foldReal(t, fi, w)
|
||||
if err != nil {
|
||||
t.Fatalf("fold: %v", err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
// dL/dy at the roundtrip output, propagated through both adjoints,
|
||||
// must equal dL/dy itself: the fold's Wirtinger gradient is
|
||||
// w̄/2 + y/n (Mul+Real contributes w̄/2, Abs2/n contributes y/n).
|
||||
for i := range n {
|
||||
dy := conj(w[i])/2 + fi.Data().ComplexAt(i)/complex(float64(n), 0)
|
||||
got := xt.Grad().ComplexAt(i)
|
||||
if cmplxAbs(got-dy) > 1e-8 {
|
||||
t.Fatalf("grad[%d] = %v, want %v", i, got, dy)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,286 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
)
|
||||
|
||||
// The cross-rank gradient sweep (the machine that catches regressions
|
||||
// hiding in untested shapes, like the LayerNorm affine reduction):
|
||||
// every listed differentiable op runs through a finite-difference
|
||||
// check on several input ranks and both float element types.
|
||||
|
||||
type sweepCase struct {
|
||||
name string
|
||||
ranks [][]int // the shapes the op must answer for
|
||||
run func(x *Tensor) (*Tensor, error)
|
||||
}
|
||||
|
||||
// sweepValue is deterministic, sign-varying and comfortably away from
|
||||
// kinks and saturation boundaries.
|
||||
func sweepValue(i int) float64 {
|
||||
return math.Sin(float64(i%17)*0.7)*2 + 0.25
|
||||
}
|
||||
|
||||
func sweepPattern(n int) []float64 {
|
||||
out := make([]float64, n)
|
||||
for i := range out {
|
||||
out[i] = 0.5*float64(i%4) - 0.75
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func sweepMatrixPattern(width int) *core.Array {
|
||||
vals := sweepPattern(width * width)
|
||||
arr, _ := core.FromFloats(vals, width, width)
|
||||
return arr
|
||||
}
|
||||
|
||||
// sweepPositive keeps logs and divisions inside their real domains no
|
||||
// matter how the signed sweep values land, shaped like the input.
|
||||
func sweepPositive(shape []int) *core.Array {
|
||||
n := numEl(shape)
|
||||
out := make([]float64, n)
|
||||
for i := range out {
|
||||
out[i] = 3 + float64(i%3)
|
||||
}
|
||||
arr, _ := core.FromFloats(out, shape...)
|
||||
return arr
|
||||
}
|
||||
|
||||
// sweepSecond derives a second operand from an independent pattern:
|
||||
// paired cases need two leaves but stay deterministic.
|
||||
func sweepSecond(shape []int) (*Tensor, error) {
|
||||
total := 1
|
||||
for _, d := range shape {
|
||||
total *= d
|
||||
}
|
||||
vals := make([]float64, total)
|
||||
for i := range vals {
|
||||
vals[i] = math.Cos(float64(i%11))*1.5 - 0.5
|
||||
}
|
||||
a, err := core.FromFloats(vals, shape...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return FromArray(a, true), nil
|
||||
}
|
||||
|
||||
func numEl(shape []int) int {
|
||||
n := 1
|
||||
for _, d := range shape {
|
||||
n *= d
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// checkOp runs one case per shape and dtype: forward on a gradient
|
||||
// leaf, weighted-sum loss with a fixed mask so every slot gets its own
|
||||
// coefficient, then analytic-vs-central-difference compare.
|
||||
func checkOp(t *testing.T, tc sweepCase, dt core.Dtype) {
|
||||
t.Helper()
|
||||
for _, dims := range tc.ranks {
|
||||
n := numEl(dims)
|
||||
build := func(reqGrad bool) *Tensor {
|
||||
vals := make([]float64, n)
|
||||
for i := range vals {
|
||||
vals[i] = sweepValue(i + len(dims))
|
||||
}
|
||||
a, err := core.FromFloats(vals, dims...)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if dt == core.Float32 {
|
||||
a32, cerr := core.Astype(a, core.Float32)
|
||||
if cerr != nil {
|
||||
t.Fatal(cerr)
|
||||
}
|
||||
a = a32
|
||||
}
|
||||
return FromArray(a, reqGrad)
|
||||
}
|
||||
|
||||
x := build(true)
|
||||
out, err := tc.run(x)
|
||||
if err != nil {
|
||||
t.Fatalf("%s %v %v: forward: %v", tc.name, dims, dt, err)
|
||||
}
|
||||
|
||||
// The mask matches the OUTPUT shape: reducing ops return fewer
|
||||
// slots than their input carries.
|
||||
outN := out.Data().Len()
|
||||
maskVals := sweepPattern(outN)
|
||||
mArr, _ := core.FromFloats(maskVals, out.Data().Shape()...)
|
||||
scaled, err := out.Mul(FromArray(mArr, false))
|
||||
if err != nil {
|
||||
t.Fatalf("%s %v %v: loss mul: %v", tc.name, dims, dt, err)
|
||||
}
|
||||
loss, err := scaled.Sum()
|
||||
if err != nil {
|
||||
t.Fatalf("%s %v %v: loss sum: %v", tc.name, dims, dt, err)
|
||||
}
|
||||
if err := loss.Backward(); err != nil {
|
||||
t.Fatalf("%s %v %v: backward: %v", tc.name, dims, dt, err)
|
||||
}
|
||||
|
||||
ref := numericGrad(func(v *core.Array) float64 {
|
||||
o, rerr := tc.run(FromArray(v, false))
|
||||
if rerr != nil {
|
||||
return math.NaN()
|
||||
}
|
||||
total := 0.0
|
||||
for i := range maskVals {
|
||||
total += mArr.FloatAt(i) * o.Data().FloatAt(i)
|
||||
}
|
||||
return total
|
||||
}, x.Data())
|
||||
|
||||
got := x.Grad()
|
||||
if got.Len() != len(ref) {
|
||||
t.Fatalf("%s %v %v: gradient length %d, reference %d",
|
||||
tc.name, dims, dt, got.Len(), len(ref))
|
||||
}
|
||||
scale := 1.0
|
||||
for _, r := range ref {
|
||||
if s := math.Abs(r); s > scale {
|
||||
scale = s
|
||||
}
|
||||
}
|
||||
tol := 1e-4
|
||||
if dt == core.Float32 {
|
||||
tol = 8e-2
|
||||
}
|
||||
for i := range ref {
|
||||
if math.Abs(got.FloatAt(i)-ref[i]) > tol*scale {
|
||||
t.Errorf("%s %v %v: grad[%d] = %v, want ≈%v",
|
||||
tc.name, dims, dt, i, got.FloatAt(i), ref[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGradientSweepAcrossRanksAndDtypes(t *testing.T) {
|
||||
shapes234 := [][]int{{4}, {2, 3}, {2, 2, 2}}
|
||||
|
||||
simple := []sweepCase{
|
||||
{name: "Neg", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Neg() }},
|
||||
{name: "Exp", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Exp() }},
|
||||
{name: "Sigmoid", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Sigmoid() }},
|
||||
{name: "Tanh", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Tanh() }},
|
||||
{name: "Abs", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Abs() }},
|
||||
{name: "Pow3", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Pow(3) }},
|
||||
{name: "Scale", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Scale(-1.75) }},
|
||||
{name: "ClipInterior", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Clip(-2.75, 2.75) }},
|
||||
{name: "LogShifted", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) {
|
||||
up, err := x.Add(FromArray(sweepPositive(x.Data().Shape()), false))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return up.Log()
|
||||
}},
|
||||
{name: "TransposeAxesReverse", ranks: [][]int{{2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
|
||||
d := x.Data().NDim()
|
||||
perm := make([]int, d)
|
||||
for i := range perm {
|
||||
perm[i] = d - 1 - i
|
||||
}
|
||||
return x.TransposeAxes(perm...)
|
||||
}},
|
||||
{name: "ReshapeFlatten", ranks: [][]int{{2, 3}, {2, 2, 2}}, run: func(x *Tensor) (*Tensor, error) {
|
||||
return x.Reshape(x.Data().Len())
|
||||
}},
|
||||
{name: "SumAxisZero", ranks: [][]int{{2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
|
||||
return x.SumAxis(0)
|
||||
}},
|
||||
{name: "MeanAxisLast", ranks: [][]int{{2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
|
||||
return x.MeanAxis(x.Data().NDim() - 1)
|
||||
}},
|
||||
// The L2 norm's backward runs a dedicated float32 sweep beside
|
||||
// the float64 one; the two dtype legs below drive both.
|
||||
{name: "L2NormAxisLast", ranks: [][]int{{4}, {2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
|
||||
return x.L2NormAxis(x.Data().NDim() - 1)
|
||||
}},
|
||||
}
|
||||
|
||||
elementPairs := []struct {
|
||||
name string
|
||||
op func(a, b *Tensor) (*Tensor, error)
|
||||
}{
|
||||
{"Add", func(a, b *Tensor) (*Tensor, error) { return a.Add(b) }},
|
||||
{"Sub", func(a, b *Tensor) (*Tensor, error) { return a.Sub(b) }},
|
||||
{"Mul", func(a, b *Tensor) (*Tensor, error) { return a.Mul(b) }},
|
||||
}
|
||||
for _, ep := range elementPairs {
|
||||
simple = append(simple, sweepCase{
|
||||
name: ep.name,
|
||||
ranks: shapes234,
|
||||
run: func(x *Tensor) (*Tensor, error) {
|
||||
other, err := sweepSecond(x.Data().Shape())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return ep.op(x, other)
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Division keeps both operands positive via the shared shift.
|
||||
simple = append(simple, sweepCase{
|
||||
name: "DivShifted",
|
||||
ranks: shapes234,
|
||||
run: func(x *Tensor) (*Tensor, error) {
|
||||
other, err := sweepSecond(x.Data().Shape())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
lift, err := other.Add(FromArray(sweepPositive(x.Data().Shape()), false))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return x.Div(lift)
|
||||
},
|
||||
})
|
||||
|
||||
// Column concatenation against half of a second leaf.
|
||||
simple = append(simple, sweepCase{
|
||||
name: "ConcatColumns",
|
||||
ranks: [][]int{{2, 4}},
|
||||
run: func(x *Tensor) (*Tensor, error) {
|
||||
extra, err := sweepSecond(x.Data().Shape())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
halves, err := extra.Slice(1, 0, x.Data().Shape()[1]/2)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return x.Concat(halves, 1)
|
||||
},
|
||||
})
|
||||
|
||||
// A matmul product collapsed by an axis sum, the inference
|
||||
// backbone's gradient path.
|
||||
simple = append(simple, sweepCase{
|
||||
name: "MatMulSumRows",
|
||||
ranks: [][]int{{3, 4}},
|
||||
run: func(x *Tensor) (*Tensor, error) {
|
||||
cols := x.Data().Shape()[1]
|
||||
product, err := x.MatMul(FromArray(sweepMatrixPattern(cols), false))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return product.SumAxis(0)
|
||||
},
|
||||
})
|
||||
|
||||
for _, tc := range simple {
|
||||
for _, dt := range []core.Dtype{core.Float, core.Float32} {
|
||||
checkOp(t, tc, dt)
|
||||
}
|
||||
}
|
||||
}
|
||||
+2303
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,343 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package grad
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func mustTensor(t *testing.T, vals []float64, shape ...int) *Tensor {
|
||||
t.Helper()
|
||||
tt, err := FromFloat64s(vals, true, shape...)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat64s(%v, %v): %v", vals, shape, err)
|
||||
}
|
||||
return tt
|
||||
}
|
||||
|
||||
func TestAutogradBasicChain(t *testing.T) {
|
||||
x := mustTensor(t, []float64{2, 3}, 2)
|
||||
y := mustTensor(t, []float64{4, 5}, 2)
|
||||
z, err := x.Add(y)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sq, err := z.Mul(z)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s, err := sq.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// d/dx (x+y)^2 summed = 2(x+y); at x=2: 12, at x=3: 16.
|
||||
if gx := x.Grad().FloatAt(0); math.Abs(gx-12) > 1e-9 {
|
||||
t.Errorf("grad x[0]: got %v, want 12", gx)
|
||||
}
|
||||
if gx := x.Grad().FloatAt(1); math.Abs(gx-16) > 1e-9 {
|
||||
t.Errorf("grad x[1]: got %v, want 16", gx)
|
||||
}
|
||||
// y's gradient matches x's, symmetric in the sum.
|
||||
if gy := y.Grad().FloatAt(0); math.Abs(gy-12) > 1e-9 {
|
||||
t.Errorf("grad y[0]: got %v, want 12", gy)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutogradMatMul(t *testing.T) {
|
||||
a := mustTensor(t, []float64{1, 2, 3, 4}, 2, 2)
|
||||
b := mustTensor(t, []float64{5, 6, 7, 8}, 2, 2)
|
||||
p, err := a.MatMul(b)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s, err := p.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// d/dA sum(A·B) = J·Bᵀ; Bᵀ = [[5,7],[6,8]], so J·Bᵀ =
|
||||
// [[11,15],[11,15]] (each row is the column sums of Bᵀ).
|
||||
wantA := []float64{11, 15, 11, 15}
|
||||
for i := range 4 {
|
||||
if g := a.Grad().FloatAt(i); math.Abs(g-wantA[i]) > 1e-9 {
|
||||
t.Errorf("grad A[%d]: got %v, want %v", i, g, wantA[i])
|
||||
}
|
||||
}
|
||||
// d/dB sum(A·B) = Aᵀ·J; Aᵀ = [[1,3],[2,4]], row sums: 4, 6, so
|
||||
// Aᵀ·J = [[4,4],[6,6]].
|
||||
wantB := []float64{4, 4, 6, 6}
|
||||
for i := range 4 {
|
||||
if g := b.Grad().FloatAt(i); math.Abs(g-wantB[i]) > 1e-9 {
|
||||
t.Errorf("grad B[%d]: got %v, want %v", i, g, wantB[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutogradMatMulVector(t *testing.T) {
|
||||
// 2-D × 1-D: y = A·x.
|
||||
a := mustTensor(t, []float64{1, 2, 3, 4}, 2, 2)
|
||||
x := mustTensor(t, []float64{2, 3}, 2)
|
||||
y, err := a.MatMul(x)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s, err := y.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// grad x = Aᵀ·1 = column sums of A: 4, 6.
|
||||
if g := x.Grad().FloatAt(0); math.Abs(g-4) > 1e-9 {
|
||||
t.Errorf("grad x[0]: got %v, want 4", g)
|
||||
}
|
||||
if g := x.Grad().FloatAt(1); math.Abs(g-6) > 1e-9 {
|
||||
t.Errorf("grad x[1]: got %v, want 6", g)
|
||||
}
|
||||
// grad A = outer(1, x): [[2,3],[2,3]].
|
||||
if g := a.Grad().FloatAt(2); math.Abs(g-2) > 1e-9 {
|
||||
t.Errorf("grad A[2]: got %v, want 2", g)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutogradActivations(t *testing.T) {
|
||||
// Sigmoid at 0: σ(0)=0.5, σ' = 0.25.
|
||||
x3 := mustTensor(t, []float64{0}, 1)
|
||||
sg, _ := x3.Sigmoid()
|
||||
if err := sg.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if g := x3.Grad().FloatAt(0); math.Abs(g-0.25) > 1e-9 {
|
||||
t.Errorf("Sigmoid grad at 0: got %v, want 0.25", g)
|
||||
}
|
||||
|
||||
// Exp and Log compose to identity: grad log(exp(x)) = 1.
|
||||
x4 := mustTensor(t, []float64{2}, 1)
|
||||
e, _ := x4.Exp()
|
||||
l, _ := e.Log()
|
||||
if err := l.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if g := x4.Grad().FloatAt(0); math.Abs(g-1) > 1e-9 {
|
||||
t.Errorf("log(exp) grad: got %v, want 1", g)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutogradGradientAccumulation(t *testing.T) {
|
||||
x := mustTensor(t, []float64{1}, 1)
|
||||
a, _ := x.Mul(x)
|
||||
b, _ := x.Mul(x)
|
||||
s, err := a.Add(b)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// d/dx (x² + x²) at x=1 = 4.
|
||||
if g := x.Grad().FloatAt(0); math.Abs(g-4) > 1e-9 {
|
||||
t.Errorf("accumulated grad: got %v, want 4", g)
|
||||
}
|
||||
// A second Backward without ZeroGrad accumulates into the leaf.
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if g := x.Grad().FloatAt(0); math.Abs(g-8) > 1e-9 {
|
||||
t.Errorf("accumulated grad after 2nd pass: got %v, want 8", g)
|
||||
}
|
||||
x.ZeroGrad()
|
||||
if x.Grad() != nil {
|
||||
t.Error("ZeroGrad did not clear the gradient")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutogradRejectsNonFloat(t *testing.T) {
|
||||
i, err := core.FromInts([]int64{1, 2}, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
it := FromArray(i, true)
|
||||
if _, err := it.Sum(); err == nil || !strings.Contains(err.Error(), "float") {
|
||||
t.Errorf("int Sum: %v", err)
|
||||
}
|
||||
c, _ := core.FromComplexes([]complex128{1 + 2i}, 1)
|
||||
ct := FromArray(c, true)
|
||||
// Complex Exp is differentiable (the Wirtinger graph); the
|
||||
// real-only kernels are the ones that must still refuse it.
|
||||
if _, err := ct.Exp(); err != nil {
|
||||
t.Errorf("complex Exp must differentiate: %v", err)
|
||||
}
|
||||
if _, err := ct.Log(); err == nil {
|
||||
t.Error("complex Log must error")
|
||||
}
|
||||
if _, err := ct.Tanh(); err == nil {
|
||||
t.Error("complex Tanh must error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutogradDivTanhNeg(t *testing.T) {
|
||||
// d/dx (x/y) at x=4, y=2 = 1/2.
|
||||
x := mustTensor(t, []float64{4}, 1)
|
||||
y := mustTensor(t, []float64{2}, 1)
|
||||
q, err := x.Div(y)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := q.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if g := x.Grad().FloatAt(0); math.Abs(g-0.5) > 1e-9 {
|
||||
t.Errorf("Div grad x: got %v, want 0.5", g)
|
||||
}
|
||||
// d/dy (x/y) at y=2 = -x/y² = -1.
|
||||
if g := y.Grad().FloatAt(0); math.Abs(g+1) > 1e-9 {
|
||||
t.Errorf("Div grad y: got %v, want -1", g)
|
||||
}
|
||||
|
||||
// tanh'(0) = 1.
|
||||
t0 := mustTensor(t, []float64{0}, 1)
|
||||
th, _ := t0.Tanh()
|
||||
if err := th.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if g := t0.Grad().FloatAt(0); math.Abs(g-1) > 1e-9 {
|
||||
t.Errorf("Tanh grad at 0: got %v, want 1", g)
|
||||
}
|
||||
|
||||
// d/dx (-x) = -1.
|
||||
n := mustTensor(t, []float64{3}, 1)
|
||||
neg, _ := n.Neg()
|
||||
if err := neg.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if g := n.Grad().FloatAt(0); g != -1 {
|
||||
t.Errorf("Neg grad: got %v, want -1", g)
|
||||
}
|
||||
|
||||
// Accessors and Mean grad.
|
||||
m := mustTensor(t, []float64{1, 2, 3, 4}, 2, 2)
|
||||
if m.Data() != m.Data() || m.RequiresGrad() != true {
|
||||
t.Error("accessors wrong")
|
||||
}
|
||||
mean, err := m.Mean()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := mean.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// d/dx mean(x) = 1/n = 1/4.
|
||||
if g := m.Grad().FloatAt(0); math.Abs(g-0.25) > 1e-9 {
|
||||
t.Errorf("Mean grad: got %v, want 0.25", g)
|
||||
}
|
||||
if m.Grad() == nil {
|
||||
t.Error("Grad() must be non-nil after Backward")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutogradReshape(t *testing.T) {
|
||||
x, err := FromFloat64s([]float64{1, 2, 3, 4}, true, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
r, err := x.Reshape(4)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sumT, err := r.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := sumT.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
g, err := x.Grad().Elements[float64]()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range g {
|
||||
if g[i] != 1 {
|
||||
t.Fatalf("Reshape grad[%d]: %v, want 1", i, g[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutogradTransposeBackward(t *testing.T) {
|
||||
x, err := FromFloat64s([]float64{1, 2, 3, 4}, true, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tr, err := x.Transpose()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s, err := tr.Sum()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := s.Backward(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
g, err := x.Grad().Elements[float64]()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range g {
|
||||
if g[i] != 1 {
|
||||
t.Errorf("Transpose grad[%d]: %v, want 1", i, g[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestAutogradPowZeroGradient pins the exponent-0 backward: d/dx x⁰
|
||||
// is the zero gradient everywhere, including at x = 0 where the
|
||||
// chain rule would evaluate 0·∞ and produce NaN.
|
||||
func TestAutogradPowZeroGradient(t *testing.T) {
|
||||
x := mustTensor(t, []float64{0, 2}, 2)
|
||||
y, err := x.Pow(0)
|
||||
if err != nil {
|
||||
t.Fatalf("Pow(0): %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 2 {
|
||||
g := x.Grad().FloatAt(i)
|
||||
if math.IsNaN(g) || g != 0 {
|
||||
t.Errorf("d/dx x⁰ at %g = %v, want exactly 0", x.Data().FloatAt(i), g)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestAutogradLeafBackwardAccumulates pins that Backward on a leaf
|
||||
// accumulates into the existing gradient like any other backward
|
||||
// pass, instead of overwriting it.
|
||||
func TestAutogradLeafBackwardAccumulates(t *testing.T) {
|
||||
x := mustTensor(t, []float64{3}, 1)
|
||||
if err := x.Backward(); err != nil {
|
||||
t.Fatalf("Backward: %v", err)
|
||||
}
|
||||
if got := x.Grad().FloatAt(0); got != 1 {
|
||||
t.Fatalf("first leaf Backward: grad %v, want 1", got)
|
||||
}
|
||||
if err := x.Backward(); err != nil {
|
||||
t.Fatalf("second Backward: %v", err)
|
||||
}
|
||||
if got := x.Grad().FloatAt(0); got != 2 {
|
||||
t.Fatalf("second leaf Backward: grad %v, want 2 (accumulated)", got)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user