// Copyright (c) 2026 Petr Balvín (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 }