195 lines
6.4 KiB
Go
195 lines
6.4 KiB
Go
// 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
|
|||
|
|
}
|