Thermal Hamiltonian Monte Carlo

Our implementation of Hamiltonian Monte Carlo (HMC) is a light wrapper around the AdvancedHMC.jl package. If you want to learn about the HMC theory, refer to the references and documentation provided with AdvancedHMC.jl.

Currently, our implementation works for systems with classical nuclei only (i.e. Simulation but not RingPolymerSimulation).

Example

In this example we use Hamiltonian Monte Carlo to sample the canonical distribution of a 3 dimensional harmonic oscillator potential containing 4 atoms.

using NQCDynamics
using Unitful
using UnitfulAtomic

sim = Simulation(Atoms([:H, :H, :C, :C]), Harmonic(dofs=3); temperature=300u"K")
r0 = randn(size(sim))
chain, stats = InitialConditions.ThermalMonteCarlo.run_advancedhmc_sampling(sim, r0, 1e4)
┌ Info: Finished 1000 adapation steps
  adaptor =
   StanHMCAdaptor(
       pc=WelfordVar,
       ssa=NesterovDualAveraging(γ=0.05, t_0=10.0, κ=0.75, δ=0.5, state.ϵ=1.1564875835736586),
       init_buffer=75, term_buffer=50, window_size=25,
       state=window(76, 950), window_splits(100, 150, 250, 450, 950)
   )
  κ.τ.integrator = Leapfrog(ϵ=1.16)
  h.metric = DiagEuclideanMetric([0.0008475424992935758, 0.0 ...])
┌ Info: Finished 10000 sampling steps for 1 chains in 1.125623147 (s)
  h = Hamiltonian(metric=DiagEuclideanMetric([0.0008475424992935758, 0.0 ...]), kinetic=AdvancedHMC.GaussianKinetic())
  κ = AdvancedHMC.HMCKernel{AdvancedHMC.FullMomentumRefreshment, AdvancedHMC.Trajectory{AdvancedHMC.MultinomialTS, AdvancedHMC.Leapfrog{Float64}, AdvancedHMC.GeneralisedNoUTurn{Float64}}}(AdvancedHMC.FullMomentumRefreshment(), Trajectory{AdvancedHMC.MultinomialTS}(integrator=Leapfrog(ϵ=1.16), tc=AdvancedHMC.GeneralisedNoUTurn{Float64}(10, 1000.0)))
  EBFMI_est = 0.4690394806133129
  average_acceptance_rate = 0.6115926394187681

The Monte Carlo chain contains the nuclear configurations that we have sampled:

chain
10000-element Vector{Matrix{Float64}}:
 [0.021137392966815577 0.45424685356055217 -0.12276871012376711 0.04197388713965336; 0.6465718982760524 -0.6552103222311132 -0.7552305991977999 0.23464288073822792; -0.2339847569256659 0.006697194409789969 -0.09203333211382086 0.009408149741854999]
 [0.021137392966815577 0.45424685356055217 -0.12276871012376711 0.04197388713965336; 0.6465718982760524 -0.6552103222311132 -0.7552305991977999 0.23464288073822792; -0.2339847569256659 0.006697194409789969 -0.09203333211382086 0.009408149741854999]
 [0.021137392966815577 0.45424685356055217 -0.12276871012376711 0.04197388713965336; 0.6465718982760524 -0.6552103222311132 -0.7552305991977999 0.23464288073822792; -0.2339847569256659 0.006697194409789969 -0.09203333211382086 0.009408149741854999]
 [0.021137392966815577 0.45424685356055217 -0.12276871012376711 0.04197388713965336; 0.6465718982760524 -0.6552103222311132 -0.7552305991977999 0.23464288073822792; -0.2339847569256659 0.006697194409789969 -0.09203333211382086 0.009408149741854999]
 [0.046223184583950944 0.29264935671243475 -0.10577651730531897 -0.012488406292710301; 0.40568549388863334 -0.4115165943512975 -0.45798226095241845 0.11916975016324503; -0.13571744097499944 -0.04880896918392102 -0.07735145475553382 -0.00691714945774162]
 [0.046223184583950944 0.29264935671243475 -0.10577651730531897 -0.012488406292710301; 0.40568549388863334 -0.4115165943512975 -0.45798226095241845 0.11916975016324503; -0.13571744097499944 -0.04880896918392102 -0.07735145475553382 -0.00691714945774162]
 [-0.03304915235325874 0.011181683182931085 -0.02913497236729891 -0.028813116118793693; -0.023554123250907666 -0.06310749307420316 0.010166372835559156 0.031924185727135956; -0.050246441266183606 -0.04033763510302371 0.06630310745847041 -0.0421291773795052]
 [-0.03304915235325874 0.011181683182931085 -0.02913497236729891 -0.028813116118793693; -0.023554123250907666 -0.06310749307420316 0.010166372835559156 0.031924185727135956; -0.050246441266183606 -0.04033763510302371 0.06630310745847041 -0.0421291773795052]
 [0.03793706769837372 0.01504326101750994 0.011359776145786614 0.007041813392134437; -0.0007653992134919324 0.03566786905949798 -0.020427298831124895 -0.016280607008062963; 0.011894782038020489 0.020180307057579443 -0.01530232685488462 0.03662467171795188]
 [0.03793706769837372 0.01504326101750994 0.011359776145786614 0.007041813392134437; -0.0007653992134919324 0.03566786905949798 -0.020427298831124895 -0.016280607008062963; 0.011894782038020489 0.020180307057579443 -0.01530232685488462 0.03662467171795188]
 ⋮
 [0.031392537944778275 -0.004520427516859646 -0.012947670604272989 0.003534792852914536; 0.005682280925026935 -0.08366175508446787 0.032586996623206056 -0.015527567357262947; -0.04184711905580277 0.024518131329394476 -0.037257941771158 0.04594967629025397]
 [0.031392537944778275 -0.004520427516859646 -0.012947670604272989 0.003534792852914536; 0.005682280925026935 -0.08366175508446787 0.032586996623206056 -0.015527567357262947; -0.04184711905580277 0.024518131329394476 -0.037257941771158 0.04594967629025397]
 [0.04451391958710227 0.03936061983853474 -0.012995396515471914 -0.027299070759101898; -0.01216730502525622 -0.0664921533428957 -0.032079729463257954 0.029389140453073274; 0.017331123939777746 -0.02543184921892523 0.020176724990877706 0.012333047709898966]
 [-0.05009581840319889 -0.028622216667255167 0.024561019146884314 0.03034074722411897; 0.03322331111177525 0.06708599431221038 0.008022438777915433 -0.005575483229007715; 0.025831082800538724 0.027497593875579155 -0.013591568205059266 -0.015628214541162956]
 [0.0299990874617227 0.00631299497438248 -0.0326260792229387 -0.017562864724275568; -0.0014172025135839456 -0.005919929844034083 -0.009901563335671656 0.03971188159298362; -0.03200637887661158 0.010168772215746916 0.019924533127650815 -0.03129618312077602]
 [-0.018650647626983593 0.012385104093074965 -0.011369179510421053 -0.03950454779732125; 0.03411802072209781 -0.01688058617635888 -0.023685806731113764 0.022481154518063706; -0.007803857013381446 -0.0014725060985075998 0.05015218162352235 0.05254518979145145]
 [0.0072947006695383565 0.01812132235079997 -0.012696222842270588 0.026757092615177294; 0.02864280182976996 0.04577785210519422 0.03570372026644773 0.02706459997432535; 0.030126342526724796 0.015119339181175171 0.03348533766793986 0.001239847170297792]
 [0.0037146822936602528 -0.020597928583314058 0.012235492197190175 0.007466843623591429; -0.026795042624345723 -0.03517663703715502 -0.027927728367784115 0.006275320285172026; -0.040247726354546855 -0.01779066605745973 -0.04693787568193716 0.0024853817300053323]
 [-0.002855126475563805 -0.0007911399154692834 -0.0122417427357388 -0.02909306766746504; 0.01804767217944634 0.03946802908498357 0.02568326113815293 -0.008305411505533027; 0.011117609984288221 0.009292030173460468 0.024879609010275756 -0.0025059010135745413]

and stats contains extra information about the sampling procedure:

stats
10000-element Vector{NamedTuple}:
 (n_steps = 3, is_accept = true, acceptance_rate = 1.0, log_density = -926.1469111310249, hamiltonian_energy = 3656.8030411344926, hamiltonian_energy_error = -5237.265158455463, max_hamiltonian_energy_error = -5266.260287386818, tree_depth = 1, numerical_error = false, step_size = 0.05, nom_step_size = 0.05, is_adapt = true)
 (n_steps = 1, is_accept = true, acceptance_rate = 0.0, log_density = -926.1469111310249, hamiltonian_energy = 930.9545513063861, hamiltonian_energy_error = 0.0, max_hamiltonian_energy_error = 2.468959118525787e11, tree_depth = 0, numerical_error = true, step_size = 1.241032542311506, nom_step_size = 1.241032542311506, is_adapt = true)
 (n_steps = 1, is_accept = true, acceptance_rate = 0.0, log_density = -926.1469111310249, hamiltonian_energy = 929.8466838330029, hamiltonian_energy_error = 0.0, max_hamiltonian_energy_error = 1.0349547185291501e9, tree_depth = 0, numerical_error = true, step_size = 0.5, nom_step_size = 0.5, is_adapt = true)
 (n_steps = 1, is_accept = true, acceptance_rate = 0.0, log_density = -926.1469111310249, hamiltonian_energy = 930.4599161414881, hamiltonian_energy_error = 0.0, max_hamiltonian_energy_error = 279257.3232540031, tree_depth = 0, numerical_error = true, step_size = 0.13192866018819532, nom_step_size = 0.13192866018819532, is_adapt = true)
 (n_steps = 1, is_accept = true, acceptance_rate = 1.0, log_density = -359.89579678096993, hamiltonian_energy = 811.1434014978433, hamiltonian_energy_error = -122.87480921079555, max_hamiltonian_energy_error = -122.87480921079555, tree_depth = 1, numerical_error = false, step_size = 0.028716309633808685, nom_step_size = 0.028716309633808685, is_adapt = true)
 (n_steps = 1, is_accept = true, acceptance_rate = 0.0, log_density = -359.89579678096993, hamiltonian_energy = 364.21576133275755, hamiltonian_energy_error = 0.0, max_hamiltonian_energy_error = 34402.921684373505, tree_depth = 0, numerical_error = true, step_size = 0.11260612534950616, nom_step_size = 0.11260612534950616, is_adapt = true)
 (n_steps = 3, is_accept = true, acceptance_rate = 1.0, log_density = -9.93585024627839, hamiltonian_energy = 317.7329322811849, hamiltonian_energy_error = -50.42607728939538, max_hamiltonian_energy_error = -50.42607728939538, tree_depth = 2, numerical_error = false, step_size = 0.02340023160450795, nom_step_size = 0.02340023160450795, is_adapt = true)
 (n_steps = 1, is_accept = true, acceptance_rate = 0.0, log_density = -9.93585024627839, hamiltonian_energy = 14.902085687018479, hamiltonian_energy_error = 0.0, max_hamiltonian_energy_error = 1148.7701191025144, tree_depth = 0, numerical_error = true, step_size = 0.10545494475829686, nom_step_size = 0.10545494475829686, is_adapt = true)
 (n_steps = 3, is_accept = true, acceptance_rate = 0.8806726585236526, log_density = -3.1175027723376, hamiltonian_energy = 13.186666147670458, hamiltonian_energy_error = -0.8358029673808769, max_hamiltonian_energy_error = -0.999863841590594, tree_depth = 2, numerical_error = false, step_size = 0.021583114937805906, nom_step_size = 0.021583114937805906, is_adapt = true)
 (n_steps = 1, is_accept = true, acceptance_rate = 2.5315117494475362e-9, log_density = -3.1175027723376, hamiltonian_energy = 9.246903853306685, hamiltonian_energy_error = 0.0, max_hamiltonian_energy_error = 19.794449183230704, tree_depth = 1, numerical_error = false, step_size = 0.07072771795767636, nom_step_size = 0.07072771795767636, is_adapt = true)
 ⋮
 (n_steps = 3, is_accept = true, acceptance_rate = 0.27752957458412325, log_density = -8.09042192796784, hamiltonian_energy = 17.47375424136587, hamiltonian_energy_error = 0.48407207869221835, max_hamiltonian_energy_error = 3.355615735338233, tree_depth = 2, numerical_error = false, step_size = 1.1564875835736586, nom_step_size = 1.1564875835736586, is_adapt = false)
 (n_steps = 3, is_accept = true, acceptance_rate = 0.21923052151811764, log_density = -8.09042192796784, hamiltonian_energy = 20.355673751059665, hamiltonian_energy_error = 0.0, max_hamiltonian_energy_error = 2.4130151859888613, tree_depth = 2, numerical_error = false, step_size = 1.1564875835736586, nom_step_size = 1.1564875835736586, is_adapt = false)
 (n_steps = 3, is_accept = true, acceptance_rate = 0.9551879567620682, log_density = -6.533016080969199, hamiltonian_energy = 12.328319499777326, hamiltonian_energy_error = -0.38030855508632655, max_hamiltonian_energy_error = -0.38030855508632655, tree_depth = 2, numerical_error = false, step_size = 1.1564875835736586, nom_step_size = 1.1564875835736586, is_adapt = false)
 (n_steps = 3, is_accept = true, acceptance_rate = 0.8294288232346029, log_density = -6.52850433232871, hamiltonian_energy = 11.703559301420665, hamiltonian_energy_error = -0.03967359111634394, max_hamiltonian_energy_error = 0.43781923889361707, tree_depth = 2, numerical_error = false, step_size = 1.1564875835736586, nom_step_size = 1.1564875835736586, is_adapt = false)
 (n_steps = 3, is_accept = true, acceptance_rate = 0.8202793282767097, log_density = -3.436207643968704, hamiltonian_energy = 10.200663546935413, hamiltonian_energy_error = -0.8665433993034668, max_hamiltonian_energy_error = -0.8665433993034668, tree_depth = 2, numerical_error = false, step_size = 1.1564875835736586, nom_step_size = 1.1564875835736586, is_adapt = false)
 (n_steps = 3, is_accept = true, acceptance_rate = 0.669111074968893, log_density = -5.287034411630798, hamiltonian_energy = 8.607072950328785, hamiltonian_energy_error = 0.5274289528341871, max_hamiltonian_energy_error = 0.6982281801296146, tree_depth = 2, numerical_error = false, step_size = 1.1564875835736586, nom_step_size = 1.1564875835736586, is_adapt = false)
 (n_steps = 3, is_accept = true, acceptance_rate = 0.6840693891961115, log_density = -4.442426783682812, hamiltonian_energy = 9.322924532427294, hamiltonian_energy_error = -0.2282613725100333, max_hamiltonian_energy_error = 0.659354338470866, tree_depth = 2, numerical_error = false, step_size = 1.1564875835736586, nom_step_size = 1.1564875835736586, is_adapt = false)
 (n_steps = 3, is_accept = true, acceptance_rate = 0.6089455133452514, log_density = -3.980855153647008, hamiltonian_energy = 10.335479941879214, hamiltonian_energy_error = -0.10706856912929652, max_hamiltonian_energy_error = 1.3226002887302215, tree_depth = 2, numerical_error = false, step_size = 1.1564875835736586, nom_step_size = 1.1564875835736586, is_adapt = false)
 (n_steps = 3, is_accept = true, acceptance_rate = 0.8329962035341089, log_density = -2.343216790079096, hamiltonian_energy = 6.724705514562206, hamiltonian_energy_error = -0.5763610289725234, max_hamiltonian_energy_error = 0.6951720079353345, tree_depth = 2, numerical_error = false, step_size = 1.1564875835736586, nom_step_size = 1.1564875835736586, is_adapt = false)

Here we should see that the energy expectation for the generated ensemble matches with the equipartition theorem:

julia> Estimators.@estimate potential_energy(sim, chain)0.006096305581940325
julia> austrip(sim.temperature) * 3 * 4 / 20.005700260814198944