Files
tensor/grad/hmc.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

195 lines
6.4 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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
}