From 6d7b7aa6f7481df1316b0fee76d8f8db6c78a022 Mon Sep 17 00:00:00 2001 From: Lester Hedges Date: Mon, 28 Sep 2026 11:52:44 +0100 Subject: [PATCH] Backport fix from PR #225, plus the missing tests from PR #218. [ci skip] --- src/somd2/runner/_base.py | 44 ++++++++++++++++++++ src/somd2/runner/_repex.py | 80 +++++++++++++++++-------------------- src/somd2/runner/_runner.py | 25 +++++++++--- tests/runner/test_gcmc.py | 69 ++++++++++++++++++++++++++++++++ tests/runner/test_repex.py | 59 +++++++++++++++++++++++++-- 5 files changed, 225 insertions(+), 52 deletions(-) diff --git a/src/somd2/runner/_base.py b/src/somd2/runner/_base.py index e79fc72..91383d3 100644 --- a/src/somd2/runner/_base.py +++ b/src/somd2/runner/_base.py @@ -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 = {} @@ -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 @@ -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. diff --git a/src/somd2/runner/_repex.py b/src/somd2/runner/_repex.py index f87a80c..8d95e97 100644 --- a/src/somd2/runner/_repex.py +++ b/src/somd2/runner/_repex.py @@ -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 @@ -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: @@ -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() @@ -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: @@ -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 ---------- @@ -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): """ diff --git a/src/somd2/runner/_runner.py b/src/somd2/runner/_runner.py index 5b6283e..681cd0d 100644 --- a/src/somd2/runner/_runner.py +++ b/src/somd2/runner/_runner.py @@ -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. @@ -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, @@ -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() @@ -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. diff --git a/tests/runner/test_gcmc.py b/tests/runner/test_gcmc.py index c952966..b8228dd 100644 --- a/tests/runner/test_gcmc.py +++ b/tests/runner/test_gcmc.py @@ -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() diff --git a/tests/runner/test_repex.py b/tests/runner/test_repex.py index d0a7bea..1305be8 100644 --- a/tests/runner/test_repex.py +++ b/tests/runner/test_repex.py @@ -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.") @@ -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):