"""Sampling procedures for VMC simulation.
This module contains functions for burn-in procedures and main sampling loops,
including both standard MCMC and importance sampling variants.
"""
import gc
import logging
import time
import numpy as np
import jax
import jax.numpy as jnp
from jax import random
from jax.sharding import PartitionSpec as P
import folx
from typing import Dict, Any
from .metropolis import (
metropolis_hastings, metropolis_hastings_importance_sampling,
make_mcmc_step, make_mcmc_step_importance
)
from .walker import initialize_walkers
from .mcmc_utils import prepare_sampling_results, report_progress
from functools import partial
from .sharding import (
get_vmap_fn, shard_map_wrap,
create_mesh, pad_n_walkers, n_devices, is_multi_gpu,
initialize_walkers_sharded, replicate
)
logger = logging.getLogger(__name__)
[docs]
def burn_in(ansatz,
walkers,
n_steps=2000,
step_size=0.01,
key=None,
params=None,
report_interval=100,
move_type="one",
max_vmap_batch_size=0,
mesh=None,
adapt_step_size=True):
"""Perform burn-in steps for MCMC sampling.
Args:
ansatz: Wavefunction object
walkers: Initial walker configurations
n_steps: Number of burn-in steps
step_size: Step size for MCMC proposals, std dev of Gaussian
key: PRNG key
params: Parameters for the ansatz, including jastrow and linear coefficients
report_interval: How often to print progress (and, when
adapt_step_size is True, how often to adapt step_size)
adapt_step_size: Adapt step_size after each COMPLETED
report_interval window, using the mean acceptance over that
window (never before the first window completes). Callers that
adapt step_size themselves between calls (e.g. adaptive_burn_in,
which adapts once per chunk on the chunk-mean acceptance) pass
False so the two controllers don't fight.
Returns:
Tuple of (equilibrated_walkers, acceptance_history, new_key,
step_size) — step_size reflects any adaptation during the burn-in.
"""
acceptance_history = []
if n_steps <= 0:
return walkers, acceptance_history, key, step_size
logger.info(f"Starting burn-in with {n_steps} steps...")
# Warm up walker cache (populate log_psi, psi_sign) so that
# _one_electron_move can reuse cached values instead of recomputing.
vmap_fn = get_vmap_fn(max_vmap_batch_size=max_vmap_batch_size)
_warmup_ansatz = vmap_fn(lambda w, p: ansatz(w, p), in_axes=(0, None))
_, walkers = _warmup_ansatz(walkers, params)
if mesh is not None:
axis_name = "walkers"
local_batch_ansatz = (
folx.batched_vmap(
lambda w, p: ansatz(w, p),
in_axes=(0, None),
max_batch_size=max_vmap_batch_size,
) if max_vmap_batch_size > 0 else
jax.vmap(lambda w, p: ansatz(w, p), in_axes=(0, None))
)
def _mcmc_step_sharded(walkers, key, params, step_size):
key = random.fold_in(key, jax.lax.axis_index(axis_name))
walkers_out, acceptance = metropolis_hastings(
ansatz, walkers, step_size, key, params,
move_type=move_type, batch_ansatz=local_batch_ansatz
)
acceptance = (
jax.lax.psum(acceptance, axis_name)
/ jax.lax.psum(jnp.array(1.0, dtype=acceptance.dtype), axis_name)
)
return walkers_out, acceptance
sharded_step = shard_map_wrap(
_mcmc_step_sharded,
mesh=mesh,
in_specs=(P(axis_name), P(), P(), P()),
out_specs=(P(axis_name), P()),
)
def mcmc_step(_ansatz, walkers, step_size, key, params):
return sharded_step(walkers, key, params, step_size)
else:
# JIT-compile the MCMC step function to speed up the loop.
# We partial out move_type since it's a static string argument.
# step_size is passed as argument so it can vary without recompilation.
mcmc_step = partial(
metropolis_hastings, move_type=move_type, batch_ansatz=vmap_fn(
lambda w, p: ansatz(w, p),
in_axes=(0, None)
)
)
mcmc_step = jax.jit(mcmc_step)
start_time = time.time()
for step in range(n_steps):
key, subkey = random.split(key)
walkers, acceptance = mcmc_step(
ansatz, walkers, step_size, subkey, params)
acceptance_float = float(acceptance)
acceptance_history.append(acceptance_float)
if step % report_interval == 0:
logger.info(f"Burn-in step {step}/{n_steps}, acceptance: {acceptance_float:.3f}, time: {time.time() - start_time:.2f}s")
start_time = time.time()
gc.collect()
if adapt_step_size and (step + 1) % report_interval == 0:
recent = acceptance_history[-report_interval:]
mean_acceptance = float(np.mean(recent))
step_size *= float(np.clip(mean_acceptance / 0.5, 0.5, 2.0))
logger.info("Burn-in complete.")
return walkers, acceptance_history, key, step_size
[docs]
def adaptive_burn_in(
ref_det,
full_ansatz,
walkers,
params,
step_size=0.01,
key=None,
move_type="one",
max_vmap_batch_size=0,
mesh=None,
chunk_size=500,
max_steps=50000,
acceptance_target=0.5,
acceptance_tol=0.02,
stability_window=3,
energy_stability_atol=0.05,
variance_stability_rtol=0.02,
):
"""Burn in until the ensemble's E_L/Var estimates stabilize, instead of
a fixed step count (a count tuned for one system size under-provisions
a larger one, since equilibration time grows with system size).
Runs in chunks of `chunk_size` sweeps; after each chunk:
1. PRE-GATE on acceptance within `acceptance_tol` of
`acceptance_target`. Acceptance reflects only step-size
adaptation, not global |Psi|^2 mixing -- a cheap pre-check, not
the stopping decision.
2. Once the pre-gate passes, compute batch-mean E_L and Var via
`full_ansatz` (the physical trial wavefunction -- NOT `ref_det`,
which only defines the sampling distribution) over a sliding
window of the last `stability_window` chunks. Terminate once E is
within `energy_stability_atol` (absolute: E crosses zero during
equilibration) and Var within `variance_stability_rtol`
(relative: Var is strictly positive) across the window.
3. Hard-capped at `max_steps` total sweeps.
If `walkers` came from `mcmc_utils.resample_walkers`, this function's
sweep counter is "sweeps since resample" by construction, so the
stability window can't read "stable" off still-correlated bootstrap
duplicates.
Args:
ref_det: The determinant (or other) ansatz that defines the MCMC
proposal/target distribution -- same role as `ansatz` in
`burn_in`.
full_ansatz: The physical trial wavefunction (e.g. SlaterJastrow)
whose local_energy is the actual quantity of interest for
the stability check.
walkers: Initial walker configurations.
params: Full [jastrow_params, linear_coeffs] for `full_ansatz`.
step_size: Initial MCMC proposal step size.
key: PRNG key.
move_type, max_vmap_batch_size, mesh: forwarded to burn_in.
chunk_size: Sweeps per chunk (one stability check per chunk, and
one step-size adaptation per chunk using the chunk-MEAN
acceptance -- burn_in's internal per-interval adaptation
is disabled here so the two controllers don't fight).
max_steps: Hard cap on total sweeps; the stability criterion, not
the cap, should normally terminate.
acceptance_target: Pre-gate center, matching burn_in's step-size
adaptation target.
acceptance_tol: Pre-gate band around the target.
stability_window: Number of consecutive chunks required stable.
energy_stability_atol: Absolute energy tolerance (Ha) for the
window range; scale with system size (equilibrium
fluctuations grow with it).
variance_stability_rtol: Relative tolerance for Var's window
range, above the plateau noise floor and below the
pre-plateau transition.
Returns:
Tuple of (equilibrated_walkers, chunk_history, new_key, step_size,
total_steps_run). chunk_history is a list of per-chunk dicts with
keys: steps_so_far, acceptance, mean_energy, variance (the latter
two are None for chunks skipped by the acceptance pre-gate).
"""
if key is None:
key = random.PRNGKey(int(time.time()))
if not (0.0 < acceptance_target <= 1.0):
raise ValueError(
f"acceptance_target must be a probability in (0, 1]; got "
f"{acceptance_target!r} (it is the divisor of the step-size "
f"controller and the centre of the acceptance pre-gate).")
if chunk_size <= 0:
raise ValueError(
f"chunk_size must be positive; got {chunk_size!r} (a non-positive "
f"chunk never advances total_steps and would loop forever).")
if max_steps <= 0:
raise ValueError(
f"max_steps must be positive; got {max_steps!r}.")
vmap_fn = get_vmap_fn(max_vmap_batch_size=max_vmap_batch_size, mesh=mesh)
batch_local_energy = jax.jit(vmap_fn(
lambda w, p: full_ansatz.local_energy(w, p)[0],
in_axes=(0, None),
out_axes=0,
))
chunk_history = []
e_window = []
var_window = []
total_steps = 0
while total_steps < max_steps:
this_chunk = min(chunk_size, max_steps - total_steps)
walkers, acc_hist, key, step_size = burn_in(
ref_det, walkers, this_chunk, step_size, key, params=params,
report_interval=chunk_size, move_type=move_type,
max_vmap_batch_size=max_vmap_batch_size, mesh=mesh,
adapt_step_size=False,
)
total_steps += this_chunk
chunk_acceptance = float(np.mean(acc_hist)) if acc_hist else None
if chunk_acceptance is not None:
step_size *= float(np.clip(chunk_acceptance / acceptance_target, 0.5, 2.0))
record = {"steps_so_far": total_steps, "acceptance": chunk_acceptance,
"mean_energy": None, "variance": None}
if chunk_acceptance is not None and abs(chunk_acceptance - acceptance_target) <= acceptance_tol:
energies = np.asarray(jax.device_get(batch_local_energy(walkers, params))).reshape(-1)
e_mean = float(np.mean(energies))
var = float(np.mean((energies - e_mean) ** 2))
record["mean_energy"] = e_mean
record["variance"] = var
e_window.append(e_mean)
var_window.append(var)
e_window = e_window[-stability_window:]
var_window = var_window[-stability_window:]
if len(e_window) == stability_window:
e_range = max(e_window) - min(e_window)
var_range = max(var_window) - min(var_window)
var_scale = max(abs(np.mean(var_window)), 1e-12)
if e_range <= energy_stability_atol and var_range / var_scale <= variance_stability_rtol:
chunk_history.append(record)
logger.info(
f"adaptive_burn_in converged after {total_steps} sweeps "
f"(E window {e_window}, Var window {var_window})."
)
return walkers, chunk_history, key, step_size, total_steps
else:
# Pre-gate not yet passed -- acceptance still settling.
# Reset the stability window: a chunk that skipped the E/Var
# check contributes no evidence either way, and letting a
# stale window from before a pre-gate dip carry over risks
# false "stable" on a window that isn't contiguous.
e_window = []
var_window = []
chunk_history.append(record)
logger.info(
f"adaptive_burn_in hit max_steps={max_steps} without meeting the "
f"stability criterion -- returning current state; consider "
f"raising max_steps or loosening the stability tolerances."
)
return walkers, chunk_history, key, step_size, total_steps
[docs]
def burn_in_with_importance(ansatz, walkers, n_steps, time_step, key, params, report_interval=100, mesh=None):
"""Perform burn-in steps for MCMC sampling with importance sampling.
Args:
ansatz: Wavefunction object
walkers: Initial walker configurations
n_steps: Number of burn-in steps
time_step: Time step for the drift-diffusion process
key: PRNG key
params: Parameters for the ansatz, including jastrow and linear coefficients
report_interval: How often to print progress
Returns:
Tuple of (equilibrated_walkers, acceptance_history, new_key,
time_step)
"""
acceptance_history = []
if n_steps <= 0:
return walkers, acceptance_history, key
logger.info(f"Starting burn-in with {n_steps} steps using importance sampling...")
if mesh is not None:
axis_name = "walkers"
def _mcmc_step_sharded(walkers, key, params, time_step):
key = random.fold_in(key, jax.lax.axis_index(axis_name))
walkers_out, acceptance = metropolis_hastings_importance_sampling(
ansatz, walkers, time_step, key, params
)
acceptance = (
jax.lax.psum(acceptance, axis_name)
/ jax.lax.psum(jnp.array(1.0, dtype=acceptance.dtype), axis_name)
)
return walkers_out, acceptance
sharded_step = shard_map_wrap(
_mcmc_step_sharded,
mesh=mesh,
in_specs=(P(axis_name), P(), P(), P()),
out_specs=(P(axis_name), P()),
)
def mcmc_step(_ansatz, walkers, time_step, key, params):
return sharded_step(walkers, key, params, time_step)
else:
mcmc_step = metropolis_hastings_importance_sampling
mcmc_step = jax.jit(mcmc_step)
time_start = time.time()
for step in range(n_steps):
key, subkey = random.split(key)
walkers, acceptance = mcmc_step(
ansatz, walkers, time_step, subkey, params)
acceptance_history.append(acceptance)
if step % report_interval == 0:
logger.info(f"Burn-in step {step}/{n_steps}, acceptance: {acceptance_history[-1]}, time: {time.time() - time_start:.2f}s")
time_step *= acceptance_history[-1]/0.5
time_start = time.time()
logger.info("Burn-in complete.")
return walkers, acceptance_history, key, time_step
[docs]
def sample(
ansatz,
n_walkers: int = 100,
n_steps: int = 1000,
step_size: float = 1.0,
thinning: int = 10,
burn_in_steps: int = 1000,
initial_walkers=None,
use_importance_sampling: bool = False,
params=None,
key=None,
move_type: str = "one",
report_interval: int = 100,
max_vmap_batch_size: int = 0
) -> Dict[str, Any]:
"""Perform MCMC sampling for quantum wavefunction.
Args:
ansatz: Wavefunction object with __call__ method that returns ψ(R)
n_walkers: Number of parallel walkers
n_steps: Number of MCMC steps for each walker
step_size: Standard deviation of Gaussian proposal for regular MCMC
or time step for importance sampling (typically 0.01-0.05)
thinning: Keep only every `thinning` steps to reduce autocorrelation
burn_in_steps: Number of initial MCMC steps to discard (equilibration)
initial_walkers: Optional initial positions, otherwise initialized near nuclei
use_importance_sampling: Whether to use importance sampling with drift
params: Parameters for the ansatz, including jastrow and linear coefficients
key: PRNG key
Returns:
Dictionary with sampling results and statistics
"""
if key is None:
key = random.PRNGKey(int(time.time()))
mesh = None
if is_multi_gpu():
num_devices = n_devices()
mesh = create_mesh()
padded_n = pad_n_walkers(n_walkers, num_devices)
if padded_n != n_walkers:
print(f"Padding n_walkers from {n_walkers} to {padded_n} "
f"(divisible by {num_devices} devices)")
n_walkers = padded_n
print(f"Multi-GPU auto-detected: {num_devices} devices, "
f"{n_walkers // num_devices} walkers/device")
if mesh is not None:
walkers = initialize_walkers_sharded(
ansatz, n_walkers, mesh, initial_walkers=initial_walkers, key=key
)
if params is not None:
params = replicate(params, mesh)
key = replicate(key, mesh)
else:
walkers = initialize_walkers(ansatz, n_walkers, initial_walkers, key)
logger.info("Starting production sampling...")
logger.info(f"Burn-in steps = {burn_in_steps}")
logger.info(f"Number of walkers = {n_walkers}")
logger.info(f"Number of steps = {n_steps}")
logger.info(f"Thinning factor = {thinning}")
logger.info(f"Step size = {step_size:.4f}")
logger.info(f"Using importance sampling: {use_importance_sampling}")
logger.info(f"Move type: {move_type}")
if use_importance_sampling:
walkers, acceptance_history, key, step_size = burn_in_with_importance(
ansatz, walkers, burn_in_steps, step_size, key, params, mesh=mesh)
else:
walkers, acceptance_history, key, step_size = burn_in(
ansatz, walkers, burn_in_steps, step_size, key=key, params=params,
move_type=move_type, max_vmap_batch_size=max_vmap_batch_size, mesh=mesh)
collected_samples = []
collected_energies = []
step_times = []
vmap_fn = get_vmap_fn(max_vmap_batch_size=max_vmap_batch_size, mesh=mesh)
if use_importance_sampling:
mcmc_step = make_mcmc_step_importance(ansatz, step_size, mesh=mesh)
else:
mcmc_step = make_mcmc_step(
ansatz, step_size, move_type,
max_vmap_batch_size=max_vmap_batch_size, mesh=mesh
)
batch_local_energy = jax.jit(vmap_fn(
lambda w, p: ansatz.local_energy(w, p)[0],
in_axes=(0, None)
))
start_time = time.time()
for step in range(n_steps):
key, subkey = random.split(key)
walkers, acceptance = mcmc_step(
ansatz, walkers, subkey, params)
acceptance_history.append(acceptance)
if step % thinning == 0:
energies = batch_local_energy(walkers, params)
# Convert to numpy to avoid holding JAX device references
collected_samples.append(np.array(walkers.positions))
collected_energies.append(np.array(energies))
if step % report_interval == 0 or step == n_steps - 1:
step_time = time.time() - start_time
step_times.append(step_time)
if not collected_energies and step == 0:
energies = batch_local_energy(walkers, params)
logger.info(f"Batch mean energy: {jnp.mean(energies):.6f}")
report_progress(step, n_steps, acceptance_history, step_times,
collected_energies if collected_energies else None)
start_time = time.time()
gc.collect()
return prepare_sampling_results(
collected_samples, collected_energies, acceptance_history, walkers, step_times)