Source code for pytc.vmc.test.test_mcmc_utils

import tempfile
import unittest
import numpy as np
import jax
import jax.numpy as jnp

from pytc.vmc.mcmc_utils import (
    save_optimization_history, load_optimization_history,
    save_walkers, load_walkers, resample_walkers,
)
from pytc.vmc.walker import Walker, initialize_walker_state


[docs] class _FakeDetAnsatz: """Minimal ansatz stub -- initialize_walker_state only reads these four attributes.""" atom_coords = jnp.array([[0.0, 0.0, 0.0]]) atom_charges = jnp.array([1.0]) n_electrons = 3 n_alpha = 2
[docs] class TestMCMCUtils(unittest.TestCase):
[docs] def test_save_load_optimization_history(self): """Test round-trip save and load of VMC optimization history data to HDF5. Ensures that: 1. Tuples, lists, and dict formats in the PyTree are correctly maintained after round-tripping. 2. Leaves of the PyTree (across different steps) are correctly stacked sequentially (axis=0). 3. Real and complex float arrays are seamlessly restored. """ params_history = [ ( {"weight": np.array([1.0, 2.0]), "bias": np.array([0.5])}, [np.array([1+1j, 2-1j])] ), ( {"weight": np.array([3.0, 4.0]), "bias": np.array([1.5])}, [np.array([3+2j, 4-2j])] ), ( {"weight": np.array([5.0, 6.0]), "bias": np.array([2.5])}, [np.array([5+3j, 6-3j])] ) ] data = { 'cost': np.array([10.5, 8.2, 7.1]), 'energies': np.array([-1.1, -1.5, -1.8]), 'stds': np.array([0.1, 0.05, 0.02]), 'acceptance': np.array([0.6, 0.7, 0.65]), 'params': params_history } with tempfile.NamedTemporaryFile(suffix='.h5') as tmp: filepath = tmp.name saved_path = save_optimization_history(data, filepath) self.assertEqual(saved_path, filepath) loaded_data = load_optimization_history(filepath) self.assertTrue('cost' in loaded_data) np.testing.assert_array_almost_equal(loaded_data['cost'], data['cost']) np.testing.assert_array_almost_equal(loaded_data['energies'], data['energies']) stacked_params = loaded_data['params'] # Should have restored top level as a tuple self.assertTrue(isinstance(stacked_params, tuple), "Top level param should be a tuple") self.assertEqual(len(stacked_params), 2) # Element 0 of tuple is a dict self.assertTrue(isinstance(stacked_params[0], dict)) self.assertIn("weight", stacked_params[0]) self.assertIn("bias", stacked_params[0]) # weight shape should be (3, 2) since there are 3 steps and the array is len 2 self.assertEqual(stacked_params[0]["weight"].shape, (3, 2)) np.testing.assert_array_almost_equal(stacked_params[0]["weight"][0], np.array([1.0, 2.0])) np.testing.assert_array_almost_equal(stacked_params[0]["weight"][2], np.array([5.0, 6.0])) # Element 1 of tuple is a list self.assertTrue(isinstance(stacked_params[1], list)) self.assertEqual(len(stacked_params[1]), 1) # 3 steps, 2 elements per array -> shape (3, 2) complex_arr = stacked_params[1][0] self.assertEqual(complex_arr.shape, (3, 2)) self.assertTrue(np.iscomplexobj(complex_arr)) np.testing.assert_array_almost_equal(complex_arr[0], np.array([1+1j, 2-1j])) np.testing.assert_array_almost_equal(complex_arr[2], np.array([5+3j, 6-3j])) last_params = jax.tree_util.tree_map(lambda x: x[-1], stacked_params) self.assertTrue(isinstance(last_params, tuple)) np.testing.assert_array_almost_equal(last_params[0]["weight"], np.array([5.0, 6.0])) np.testing.assert_array_almost_equal(last_params[1][0], np.array([5+3j, 6-3j]))
[docs] def test_save_load_walkers(self): """Round-trip a Walker's full state (positions plus cached psi/det/grad/lap fields, and the det_up/det_down (sign, log|det|) tuples) through save_walkers/load_walkers -- the continued-walkers checkpoint mechanism. """ n_walkers, n_alpha, n_beta = 4, 2, 1 n_electrons = n_alpha + n_beta walkers = Walker( positions=jnp.arange(n_walkers * n_electrons * 3, dtype=jnp.float64).reshape( n_walkers, n_electrons, 3), det_up=(jnp.ones(n_walkers), jnp.full(n_walkers, -0.5)), det_down=(jnp.ones(n_walkers), jnp.full(n_walkers, -0.25)), slater_up=jnp.zeros((n_walkers, n_alpha, n_alpha)), slater_down=jnp.zeros((n_walkers, n_beta, n_beta)), inv_up=jnp.zeros((n_walkers, n_alpha, n_alpha)), inv_down=jnp.zeros((n_walkers, n_beta, n_beta)), grad_up=jnp.zeros((n_walkers, n_alpha, n_alpha, 3)), grad_down=jnp.zeros((n_walkers, n_beta, n_beta, 3)), lap_up=jnp.zeros((n_walkers, n_alpha, n_alpha)), lap_down=jnp.zeros((n_walkers, n_beta, n_beta)), move_mask=jnp.ones((n_walkers, n_electrons), dtype=bool), log_psi=jnp.linspace(-3.0, -1.0, n_walkers), psi_sign=jnp.ones(n_walkers), log_jastrow=jnp.linspace(0.1, 0.4, n_walkers), ) with tempfile.NamedTemporaryFile(suffix='.h5') as tmp: filepath = tmp.name saved_path = save_walkers(walkers, filepath) self.assertEqual(saved_path, filepath) loaded = load_walkers(filepath) self.assertIsInstance(loaded, Walker) np.testing.assert_array_almost_equal(loaded.positions, walkers.positions) np.testing.assert_array_almost_equal(loaded.log_psi, walkers.log_psi) np.testing.assert_array_almost_equal(loaded.log_jastrow, walkers.log_jastrow) np.testing.assert_array_equal(loaded.move_mask, walkers.move_mask) # det_up/det_down are (sign, log|det|) tuples -- confirm the tuple # structure survives, not just the leaf values. self.assertIsInstance(loaded.det_up, tuple) self.assertEqual(len(loaded.det_up), 2) np.testing.assert_array_almost_equal(loaded.det_up[1], walkers.det_up[1]) np.testing.assert_array_almost_equal(loaded.det_down[1], walkers.det_down[1])
[docs] def test_resample_walkers_upsamples_and_jitters(self): """resample_walkers should (a) produce exactly target_n_walkers walkers, (b) reset the cached psi/det/grad/lap fields (same convention as a cold init, since jittered positions invalidate them), and (c) place each resampled position near (not exactly on top of) one of the source positions. """ ansatz = _FakeDetAnsatz() source_n = 5 source_positions = jnp.arange( source_n * ansatz.n_electrons * 3, dtype=jnp.float32 ).reshape(source_n, ansatz.n_electrons, 3) source = initialize_walker_state(ansatz, source_positions) target_n = 23 # deliberately not a multiple of source_n step_size = 0.05 resampled = resample_walkers( ansatz, source, target_n_walkers=target_n, step_size=step_size, key=jax.random.PRNGKey(7)) self.assertEqual(resampled.positions.shape, (target_n, ansatz.n_electrons, 3)) # Cached fields must be invalidated -- jittered positions make the # copied cache values wrong, and move_mask=True signals "recompute". self.assertTrue(bool(jnp.all(resampled.move_mask))) self.assertTrue(bool(jnp.all(resampled.log_psi == 0))) # Each resampled walker's position should be within a small # multiple of step_size of SOME source position (bootstrap + # jitter, not an arbitrary new location). deltas = resampled.positions[:, None, :, :] - source_positions[None, :, :, :] nearest_dist = jnp.min(jnp.sqrt(jnp.sum(deltas ** 2, axis=(-1, -2))), axis=1) self.assertTrue(bool(jnp.all(nearest_dist < 10 * step_size))) # And it shouldn't be an EXACT duplicate (jitter actually applied). self.assertTrue(bool(jnp.all(nearest_dist > 0)))
if __name__ == '__main__': unittest.main()