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
13 changes: 6 additions & 7 deletions src/loch/_sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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."
)
Expand Down
66 changes: 66 additions & 0 deletions tests/test_ghost_file.py
Original file line number Diff line number Diff line change
@@ -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"