Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
44 changes: 44 additions & 0 deletions src/somd2/runner/_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -646,6 +646,9 @@ def __init__(self, system, config):
self._ec_rows = {}
self._max_ec_rows = 10000

# Per-window GCMC ghost residue lines collected since the last checkpoint.
self._ghost_rows = {}

# Per-window cache of the integrator's integration force groups bitmask.
self._integration_groups = {}

Expand Down Expand Up @@ -2540,6 +2543,8 @@ def _checkpoint(
if not is_post_equilibration:
self._flush_energy_components(index)

self._flush_ghost_residues(index)

except Exception as e:
return index, e

Expand Down Expand Up @@ -2788,6 +2793,45 @@ def _flush_energy_components(self, index):
)
_pq_local.write_table(table, filepath)

def _save_ghost_residues(self, index, gcmc_sampler):
"""
Record the current GCMC ghost residue indices for a window. This must
be called at the point the matching trajectory frame is taken. The
lines are written by _flush_ghost_residues() at checkpoint time, along
with the frames.

Parameters
----------

index : int
The index of the window or replica.

gcmc_sampler : loch.GCMCSampler
The GCMC sampler for the window.
"""
ghost_residues = gcmc_sampler.ghost_residues()
self._ghost_rows.setdefault(index, []).append(
f"{', '.join([str(x) for x in ghost_residues])}\n"
)

def _flush_ghost_residues(self, index):
"""
Append the GCMC ghost residue lines buffered by _save_ghost_residues()
to the ghost residue file for a window.

Parameters
----------

index : int
The index of the window or replica.
"""
rows = self._ghost_rows.pop(index, [])
if not rows:
return

with open(self._filenames[index]["gcmc_ghosts"], "a") as f:
f.writelines(rows)

def _restore_backup_files(self):
"""
Restore backup files in the working directory.
Expand Down
80 changes: 37 additions & 43 deletions src/somd2/runner/_repex.py
Original file line number Diff line number Diff line change
Expand Up @@ -2064,6 +2064,27 @@ def run(self):
_logger.error("Commit cancelled. Exiting.")
_sys.exit(1)

# Assemble an energy matrix from the results.
_logger.info("Assembling energy matrix")
energy_matrix = self._assemble_results(results)

# Mix the replicas.
_logger.info("Mixing replicas")
old_states = self._dynamics_cache.get_states()
self._dynamics_cache.set_states(
self._mix_replicas(
self._config.num_lambda,
energy_matrix,
self._dynamics_cache.get_proposed(),
self._dynamics_cache.get_accepted(),
)
)

# This only permutes the stored states. They are pushed into the
# contexts by load_replica() at the start of the next block, which
# is also where the pre-run state for crash recovery is captured.
self._dynamics_cache.mix_states(old_states)

# Checkpoint. This happens once the whole cycle is complete, with
# every checkpoint file written under a single lock, so that an
# external process reading the output directory always sees a
Expand Down Expand Up @@ -2129,26 +2150,7 @@ def run(self):
_logger.error("Checkpoint cancelled. Exiting.")
_sys.exit(1)

# Assemble an energy matrix from the results.
_logger.info("Assembling energy matrix")
energy_matrix = self._assemble_results(results)

# Mix the replicas.
_logger.info("Mixing replicas")
old_states = self._dynamics_cache.get_states()
self._dynamics_cache.set_states(
self._mix_replicas(
self._config.num_lambda,
energy_matrix,
self._dynamics_cache.get_proposed(),
self._dynamics_cache.get_accepted(),
)
)

# This only permutes the stored states. They are pushed into the
# contexts by load_replica() at the start of the next block, which
# is also where the pre-run state for crash recovery is captured.
self._dynamics_cache.mix_states(old_states)
self._save_repex_state(final=not is_checkpoint)

# This is a checkpoint cycle.
if is_checkpoint:
Expand All @@ -2158,19 +2160,12 @@ def run(self):
# Advance the checkpoint threshold.
next_checkpoint += cycles_per_checkpoint

self._save_repex_state()

dynamics_executor.shutdown(wait=True)
checkpoint_executor.shutdown(wait=True)

# Record the end time for the production block.
prod_end = time()

# Save the final state, unless the last cycle was a checkpoint cycle
# and has just done so.
if not is_checkpoint:
self._save_repex_state(final=True)

# Record the end time.
end = time()

Expand Down Expand Up @@ -2285,10 +2280,10 @@ def _run_block(
finally:
gcmc_sampler.pop()

# Write ghost residues immediately after the GCMC move so the
# Record ghost residues immediately after the GCMC move so the
# ghost state and frame (saved during dynamics) are consistent.
if write_gcmc_ghosts:
gcmc_sampler.write_ghost_residues()
self._save_ghost_residues(index, gcmc_sampler)

# Perform a terminal flip move before dynamics if requested.
if self._terminal_flip_samplers is not None and is_terminal_flip:
Expand Down Expand Up @@ -3030,7 +3025,8 @@ def _merge_gcmc_stats(self):
def _save_repex_state(self, final=False):
"""
Save the transition matrix and pickle the dynamics cache, backing up
the previous pickle, under the file lock.
the previous pickle. Must be called with the file lock held, alongside
the checkpoint files.

Parameters
----------
Expand All @@ -3040,21 +3036,19 @@ def _save_repex_state(self, final=False):
"""
label = "final replica exchange" if final else "replica exchange"

lock = _FileLock(self._lock_file)
with lock.acquire(timeout=self._config.timeout.to("seconds")):
_logger.info(f"Saving {label} transition matrix")
self._save_transition_matrix()
_logger.info(f"Saving {label} transition matrix")
self._save_transition_matrix()

if self._repex_state.exists():
_copyfile(
self._repex_state,
self._repex_state.with_suffix(".pkl.bak"),
)
if self._repex_state.exists():
_copyfile(
self._repex_state,
self._repex_state.with_suffix(".pkl.bak"),
)

_logger.info(f"Saving {label} state")
self._save_sampler_stats()
with open(self._repex_state, "wb") as f:
_pickle.dump(self._dynamics_cache, f)
_logger.info(f"Saving {label} state")
self._save_sampler_stats()
with open(self._repex_state, "wb") as f:
_pickle.dump(self._dynamics_cache, f)

def _save_sampler_stats(self):
"""
Expand Down
25 changes: 20 additions & 5 deletions src/somd2/runner/_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -1033,10 +1033,10 @@ def generate_lam_vals(lambda_base, increment=0.001):
)
)

# Write ghost residues immediately before the dynamics
# Record ghost residues immediately before the dynamics
# block if a frame will be saved within it.
if save_frames and runtime + block_size >= next_frame:
gcmc_sampler.write_ghost_residues()
self._save_ghost_residues(index, gcmc_sampler)
next_frame += self._config.frame_frequency

# Run the dynamics block.
Expand Down Expand Up @@ -1250,7 +1250,14 @@ def generate_lam_vals(lambda_base, increment=0.001):
# Acquire the file lock to ensure that the checkpoint files are
# in a consistent state if read by another process.
with lock.acquire(timeout=self._config.timeout.to("seconds")):
self._checkpoint(
# Backup any existing checkpoint files.
index, error = self._backup_checkpoint(index)

if error is not None:
raise error

# Write the checkpoint files.
index, error = self._checkpoint(
system,
index,
block,
Expand All @@ -1262,6 +1269,14 @@ def generate_lam_vals(lambda_base, increment=0.001):
gcmc_sampler=gcmc_sampler,
)

if error is not None:
raise error

# Save sampler statistics alongside the checkpoint.
self._save_sampler_stats(
index, gcmc_sampler, terminal_flip_sampler
)

# Delete all trajectory frames from the Sire system within the
# dynamics object.
dynamics._d._sire_mols.delete_all_frames()
Expand Down Expand Up @@ -1370,10 +1385,10 @@ def generate_lam_vals(lambda_base, increment=0.001):
getPositions=True, getVelocities=True
)

# Write ghost residues immediately before the dynamics
# Record ghost residues immediately before the dynamics
# block if a frame will be saved within it.
if save_frames and runtime + block_size >= next_frame:
gcmc_sampler.write_ghost_residues()
self._save_ghost_residues(index, gcmc_sampler)
next_frame += self._config.frame_frequency

# Run the dynamics block.
Expand Down
69 changes: 69 additions & 0 deletions tests/runner/test_gcmc.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,3 +61,72 @@ def test_runner_gcmc_without_a_selection(ethane_methanol):
]
assert counts, "no water count was logged"
assert all(count > 0 for count in counts), f"zero water count logged: {counts}"


@pytest.mark.skipif(not has_cuda, reason="CUDA not available.")
def test_runner_gcmc_ghosts_written_at_checkpoint(ethane_methanol):
"""
Validate that ghost residues are only written to file at a checkpoint,
under the file lock, with one line per trajectory frame.
"""
pytest.importorskip("loch")

import sire as sr
from somd2.runner import _runner as runner_module

with tempfile.TemporaryDirectory() as tmpdir:
config = Config(
runtime="16fs",
output_directory=tmpdir,
energy_frequency="4fs",
checkpoint_frequency="8fs",
frame_frequency="4fs",
platform="cuda",
max_threads=1,
num_lambda=2,
gcmc=True,
gcmc_selection="resname LIG",
gcmc_frequency="4fs",
)

runner = Runner(ethane_methanol, config)

# Windows normally run in spawned processes, so run one in this process
# for the patched lock to take effect.
index = 0
lam = runner._lambda_values[index]
ghost_file = Path(tmpdir) / f"gcmc_ghosts_{lam:.5f}.txt"

def num_lines():
return (
len(ghost_file.read_text().splitlines()) if ghost_file.exists() else 0
)

# Record the line count whenever the lock is released, and check that
# nothing is written while it isn't held.
released = [num_lines()]
real_filelock = runner_module._FileLock

class CheckingFileLock(real_filelock):
def acquire(self, *args, **kwargs):
assert num_lines() == released[-1], "ghost file written outside lock"
return super().acquire(*args, **kwargs)

def release(self, *args, **kwargs):
released.append(num_lines())
return super().release(*args, **kwargs)

runner_module._FileLock = CheckingFileLock
try:
runner._run(runner._system.clone(), index, device=0)
finally:
runner_module._FileLock = real_filelock

assert num_lines() == released[-1], "ghost file written outside lock"

traj = sr.load(
str(Path(tmpdir) / "system0.prm7"),
str(Path(tmpdir) / f"traj_{lam:.5f}.dcd"),
)
assert num_lines() > 0
assert traj.num_frames() == num_lines()
59 changes: 55 additions & 4 deletions tests/runner/test_repex.py
Original file line number Diff line number Diff line change
Expand Up @@ -457,10 +457,51 @@ def acquire(self, *args, **kwargs):
finally:
repex_module._FileLock = real_filelock

# Two cycles, each taking the lock once for the checkpoint files and once
# for the repex state, plus a final acquisition. This must not scale with
# the number of passes.
assert len(acquisitions) == 5
# Two checkpoint cycles, each taking the lock once for the checkpoint files
# and the repex state together. This must not scale with the number of passes.
assert len(acquisitions) == 2


@pytest.mark.skipif(not has_cuda, reason="CUDA not available.")
@pytest.mark.parametrize(
"runtime, checkpoint_frequency, expected",
[("8fs", "4fs", [False, False]), ("12fs", "8fs", [False, True])],
)
def test_repex_state_saved_once(
ethane_methanol, runtime, checkpoint_frequency, expected
):
"""
Validate that the replica exchange state is saved once per checkpoint
cycle, with a separate final save only when the last cycle is not a
checkpoint cycle.
"""
with tempfile.TemporaryDirectory() as tmpdir:
config = {
"runtime": runtime,
"restart": False,
"output_directory": tmpdir,
"energy_frequency": "4fs",
"checkpoint_frequency": checkpoint_frequency,
"frame_frequency": "4fs",
"platform": "cuda",
"max_threads": 1,
"num_lambda": 2,
"replica_exchange": True,
}
runner = RepexRunner(ethane_methanol, Config(**config))

saves = []
save = runner._save_repex_state

def counting_save(final=False):
saves.append(final)
return save(final=final)

runner._save_repex_state = counting_save
runner.run()

assert saves == expected
assert (Path(tmpdir) / "repex_state.pkl").exists()


@pytest.mark.skipif(not has_cuda, reason="CUDA not available.")
Expand Down Expand Up @@ -509,6 +550,16 @@ def test_repex_gcmc_bounded_contexts(ethane_methanol, max_contexts):
assert len(set(counts)) == 1, f"unbalanced ghost files: {counts}"
assert counts[0] > 0

# Each ghost line pairs with a trajectory frame.
import sire as sr

for lam, count in zip(runner._lambda_values, counts):
traj = sr.load(
str(Path(tmpdir) / "system0.prm7"),
str(Path(tmpdir) / f"traj_{lam:.5f}.dcd"),
)
assert traj.num_frames() == count


@pytest.mark.skipif(not has_cuda, reason="CUDA not available.")
def test_repex_gcmc_without_a_selection(ethane_methanol):
Expand Down