diff --git a/src/loch/_sampler.py b/src/loch/_sampler.py index 3c11e24..f7a43a1 100644 --- a/src/loch/_sampler.py +++ b/src/loch/_sampler.py @@ -259,8 +259,8 @@ def __init__( Whether to swap the end states of the alchemical systems. restart: bool - Whether this is a restart simulation. If True, then data will - be appended to existing log files. + Whether this is a restart simulation. If True, then ghost + residues will be appended to an existing 'ghost_file'. overwrite: bool Overwrite existing log files. @@ -413,12 +413,11 @@ def __init__( if not isinstance(ghost_file, str): raise ValueError("'ghost_file' must be of type 'str'") self._ghost_file = ghost_file - if not isinstance(ghost_file, str): - raise ValueError("'ghost_file' must be of type 'str'") - self._ghost_file = ghost_file - if _os.path.exists(self._ghost_file): - if not self._restart and not self._overwrite: + # On restart, keep the existing ghost residues so that they stay + # aligned with the trajectory frames written before the restart. + if _os.path.exists(self._ghost_file) and not self._restart: + if not self._overwrite: raise ValueError( "'ghost_file' already exists. Use 'overwrite=True' to overwrite it." ) diff --git a/tests/test_ghost_file.py b/tests/test_ghost_file.py new file mode 100644 index 0000000..bb8d108 --- /dev/null +++ b/tests/test_ghost_file.py @@ -0,0 +1,66 @@ +import os + +import pytest + +from loch import GCMCSampler + + +def make_sampler(mols, ghost_file, restart=False, overwrite=False): + return GCMCSampler( + mols, + cutoff_type="rf", + cutoff="10 A", + ghost_file=str(ghost_file), + log_file=None, + restart=restart, + overwrite=overwrite, + test=True, + platform="cuda", + seed=42, + ) + + +@pytest.mark.skipif( + "CUDA_VISIBLE_DEVICES" not in os.environ, + reason="Requires CUDA enabled GPU.", +) +@pytest.mark.parametrize( + "restart, overwrite, expected", + [ + (True, False, "1, 2, 3\n"), + (True, True, "1, 2, 3\n"), + (False, True, ""), + ], +) +def test_existing_ghost_file(water_box, tmp_path, restart, overwrite, expected): + """ + An existing ghost file is kept on restart, so that it stays aligned with + the trajectory, and is only cleared when overwriting a fresh run. + """ + mols, _ = water_box + + ghost_file = tmp_path / "ghosts.txt" + ghost_file.write_text("1, 2, 3\n") + + make_sampler(mols, ghost_file, restart=restart, overwrite=overwrite) + + assert ghost_file.read_text() == expected + + +@pytest.mark.skipif( + "CUDA_VISIBLE_DEVICES" not in os.environ, + reason="Requires CUDA enabled GPU.", +) +def test_existing_ghost_file_raises(water_box, tmp_path): + """ + An existing ghost file can't be silently overwritten by a fresh run. + """ + mols, _ = water_box + + ghost_file = tmp_path / "ghosts.txt" + ghost_file.write_text("1, 2, 3\n") + + with pytest.raises(ValueError, match="ghost_file"): + make_sampler(mols, ghost_file) + + assert ghost_file.read_text() == "1, 2, 3\n"