Files
tensor/grad/hmc.go
T

195 lines
6.4 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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
}