Source code for pytc.vmc.test.test_adaptive_burn_in

import unittest

import jax
jax.config.update("jax_enable_x64", True)
import jax.numpy as jnp
import numpy as np
from jax import random
from pyscf import gto, scf

from pytc.vmc.sampling import adaptive_burn_in
from pytc.vmc.walker import initialize_walkers
from pytc.ansatz.sj import SlaterJastrow
from pytc.ansatz.det import SlaterDet
from pytc.jastrow import NuclearCusp, CompositeJastrow
from pytc.jastrow.bha import BoysHandyAnalytical


[docs] class TestAdaptiveBurnIn(unittest.TestCase): """Task #12 tier-2: burn_in() should terminate on ensemble-stability evidence rather than a fixed, easy-to-under-provision step count."""
[docs] @classmethod def setUpClass(cls): mol = gto.M(atom="H 0 0 0; H 0 0 1.8; H 0 0 3.6; H 0 0 5.4", basis="sto-3g", unit="Bohr", verbose=0) mf = scf.RHF(mol).run() cls.det = SlaterDet.create(mol, mf.mo_coeff) ncusp = NuclearCusp.create(mol, name="ncusp") bha = BoysHandyAnalytical.create(mol) jastrow = CompositeJastrow.create([ncusp, bha]) cls.sj = SlaterJastrow.create(mol, jastrow, [cls.det]) cls.params = [jastrow.init_params(), jnp.ones(1)]
def _fresh_walkers(self, n_walkers=64, seed=0): key = random.PRNGKey(seed) key, subkey = random.split(key) return initialize_walkers(self.det, n_walkers, key=subkey), key
[docs] def test_terminates_early_when_stable(self): """With loose tolerances, should stop well short of max_steps and report the E/Var that triggered termination.""" walkers, key = self._fresh_walkers() _, history, _, _, total_steps = adaptive_burn_in( self.det, self.sj, walkers, self.params, step_size=0.02, key=key, chunk_size=20, max_steps=400, acceptance_tol=0.5, stability_window=2, energy_stability_atol=100.0, variance_stability_rtol=1.0, ) self.assertLess(total_steps, 400) self.assertGreater(len(history), 0) last = history[-1] self.assertIsNotNone(last["mean_energy"]) self.assertIsNotNone(last["variance"])
[docs] def test_hits_hard_cap_without_hanging(self): """An impossible stability tolerance must still terminate at max_steps rather than loop forever.""" walkers, key = self._fresh_walkers(seed=1) _, history, _, _, total_steps = adaptive_burn_in( self.det, self.sj, walkers, self.params, step_size=0.02, key=key, chunk_size=20, max_steps=60, acceptance_tol=0.5, stability_window=2, energy_stability_atol=0.0, variance_stability_rtol=0.0, ) self.assertEqual(total_steps, 60) self.assertEqual(len(history), 3) # 60 / chunk_size(20)
[docs] def test_pregate_skips_evaluation_when_acceptance_out_of_band(self): """A near-zero acceptance tolerance on freshly-initialized (highly accepted) walkers should never pass the pre-gate, so every chunk's mean_energy/variance stay None.""" walkers, key = self._fresh_walkers(seed=2) _, history, _, _, total_steps = adaptive_burn_in( self.det, self.sj, walkers, self.params, step_size=0.02, key=key, chunk_size=20, max_steps=60, acceptance_target=0.5, acceptance_tol=1e-6, stability_window=2, energy_stability_atol=100.0, variance_stability_rtol=1.0, ) self.assertEqual(total_steps, 60) self.assertTrue(all(h["mean_energy"] is None for h in history))
[docs] def test_step_size_adapts_between_chunks_on_chunk_mean(self): """The between-chunk controller must pull step_size toward the acceptance target from both sides: starting far apart, two runs with the same seed must end closer together than they started.""" walkers, key = self._fresh_walkers(seed=3) _, _, _, size_small, _ = adaptive_burn_in( self.det, self.sj, walkers, self.params, step_size=0.01, key=key, chunk_size=20, max_steps=60, acceptance_tol=0.5, stability_window=2, energy_stability_atol=0.0, variance_stability_rtol=0.0, ) walkers, key = self._fresh_walkers(seed=3) _, _, _, size_large, _ = adaptive_burn_in( self.det, self.sj, walkers, self.params, step_size=0.5, key=key, chunk_size=20, max_steps=60, acceptance_tol=0.5, stability_window=2, energy_stability_atol=0.0, variance_stability_rtol=0.0, ) self.assertLess(size_large / size_small, 0.5 / 0.01)
[docs] def test_oversized_step_never_zeros_step_size(self): """Regression: with a wildly oversized initial step (acceptance ~0) the old unclipped single-step controller could drive step_size to exactly 0, killing the chain permanently. The clipped chunk-mean controller must keep it positive and shrinking.""" walkers, key = self._fresh_walkers(seed=4) _, _, _, step_size, _ = adaptive_burn_in( self.det, self.sj, walkers, self.params, step_size=100.0, key=key, chunk_size=20, max_steps=60, acceptance_tol=0.5, stability_window=2, energy_stability_atol=0.0, variance_stability_rtol=0.0, ) self.assertGreater(step_size, 0.0) self.assertLess(step_size, 100.0)
[docs] def test_no_adaptation_before_first_complete_interval(self): """With n_steps < report_interval, no adaptation may fire at all -- in particular not at step 0 off a single acceptance sample.""" from pytc.vmc.sampling import burn_in walkers, key = self._fresh_walkers(seed=6) _, _, _, step_size = burn_in( self.det, walkers, n_steps=5, step_size=0.02, key=key, params=self.params, report_interval=10, ) self.assertEqual(step_size, 0.02)
[docs] def test_first_complete_interval_uses_all_interval_samples(self): """After exactly one full interval, the returned step_size must equal initial * clip(mean(all interval acceptances) / 0.5).""" from pytc.vmc.sampling import burn_in walkers, key = self._fresh_walkers(seed=7) _, acc_hist, _, step_size = burn_in( self.det, walkers, n_steps=10, step_size=0.02, key=key, params=self.params, report_interval=10, ) self.assertEqual(len(acc_hist), 10) expected = 0.02 * float(np.clip(np.mean(acc_hist) / 0.5, 0.5, 2.0)) self.assertAlmostEqual(step_size, expected, places=12)
[docs] def test_acceptance_target_zero_rejected(self): """acceptance_target=0 would divide by zero in the step-size controller; it must raise ValueError instead.""" walkers, key = self._fresh_walkers(seed=8) with self.assertRaises(ValueError): adaptive_burn_in( self.det, self.sj, walkers, self.params, step_size=0.02, key=key, chunk_size=20, max_steps=40, acceptance_target=0.0, )
[docs] def test_nonpositive_chunk_and_max_steps_rejected(self): """chunk_size<=0 never advances total_steps (infinite loop) and max_steps<=0 is meaningless; both must raise ValueError.""" walkers, key = self._fresh_walkers(seed=9) with self.assertRaises(ValueError): adaptive_burn_in( self.det, self.sj, walkers, self.params, step_size=0.02, key=key, chunk_size=0, max_steps=40, ) with self.assertRaises(ValueError): adaptive_burn_in( self.det, self.sj, walkers, self.params, step_size=0.02, key=key, chunk_size=20, max_steps=0, )
[docs] def test_burn_in_honours_adapt_step_size_false(self): """burn_in with adapt_step_size=False must return the initial step_size unchanged (controller owned by the caller).""" from pytc.vmc.sampling import burn_in walkers, key = self._fresh_walkers(seed=5) _, _, _, step_size = burn_in( self.det, walkers, n_steps=10, step_size=0.02, key=key, params=self.params, report_interval=5, adapt_step_size=False, ) self.assertEqual(step_size, 0.02)
if __name__ == "__main__": unittest.main()