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
30 changes: 14 additions & 16 deletions src/somd2/runner/_repex.py
Original file line number Diff line number Diff line change
Expand Up @@ -237,10 +237,12 @@ def __setstate__(self, state):

# Convert a legacy checkpoint to the current convention, in which the
# stored state of a replica is its own, with the last mix already
# applied.
# applied. Legacy states hold the destination of each replica's
# configuration, so are inverted.
if is_legacy:
self._openmm_states = [self._openmm_states[s] for s in self._states]
self._gcmc_states = [self._gcmc_states[s] for s in self._states]
sources = _np.argsort(self._states)
self._openmm_states = [self._openmm_states[s] for s in sources]
self._gcmc_states = [self._gcmc_states[s] for s in sources]

# Every replica is seeded from its stored state on a restart, since the
# contexts are created from the input system rather than the checkpoint.
Expand Down Expand Up @@ -951,7 +953,7 @@ def store_replica(self, replica):
finally:
gcmc_sampler.pop()

def mix_states(self, old_states):
def mix_states(self):
"""
Apply the result of a replica mix.

Expand All @@ -966,12 +968,6 @@ def mix_states(self, old_states):
may be loaded after another replica has already stored its post-run
state; reading through the indirection at that point would pick up the
new state rather than the pre-mix one.

Parameters
----------

old_states : numpy.ndarray
The state indices from before the last replica mix.
"""
# Permute the travelling state. This is a reference shuffle, so it is
# cheap even for large systems. Statistics and output files stay with
Expand All @@ -989,9 +985,9 @@ def mix_states(self, old_states):
for i, (state, moved) in enumerate(zip(self._states, self._state_moved))
]

# Update the swap matrix.
# Update the swap matrix with each configuration's move.
for i, state in enumerate(self._states):
self._num_swaps[old_states[i], state] += 1
self._num_swaps[state, i] += 1

def get_proposed(self):
"""
Expand Down Expand Up @@ -2070,7 +2066,6 @@ def run(self):

# 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,
Expand All @@ -2083,7 +2078,7 @@ def run(self):
# 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._dynamics_cache.mix_states()

# Checkpoint. This happens once the whole cycle is complete, with
# every checkpoint file written under a single lock, so that an
Expand Down Expand Up @@ -2956,7 +2951,8 @@ def _mix_replicas(num_replicas, energy_matrix, proposed, accepted):
-------

states: np.ndarray
The new states.
The new states, where states[i] is the replica whose configuration
seeds replica i.
"""

# Adapted from OpenMMTools: https://github.com/choderalab/openmmtools
Expand Down Expand Up @@ -2996,7 +2992,9 @@ def _mix_replicas(num_replicas, energy_matrix, proposed, accepted):
accepted[state_i, state_j] += 1
accepted[state_j, state_i] += 1

return states
# Here states[i] is the state that replica i's configuration moves to,
# whereas the configurations are moved into fixed states, so invert.
return _np.argsort(states)

def _merge_gcmc_stats(self):
"""
Expand Down
46 changes: 41 additions & 5 deletions tests/runner/test_repex.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,43 @@ def test_repex_mixing():
assert (off_diagonal == 0).all()


def test_repex_mixing_moves_configurations():
"""
Validate that each configuration is moved to the state that the mixing
accepted, for a mix that can only be a 3-cycle.
"""
from somd2.runner._repex import DynamicsCache

num_replicas = 3

# The configuration from replica i is only favourable in state i + 1.
energy_matrix = 10000 * np.ones((num_replicas, num_replicas))
for i in range(num_replicas):
energy_matrix[i, (i + 1) % num_replicas] = -10000

proposed = np.zeros((num_replicas, num_replicas), dtype=np.int32)
accepted = np.zeros((num_replicas, num_replicas), dtype=np.int32)

np.random.seed(42)
states = RepexRunner._mix_replicas(num_replicas, energy_matrix, proposed, accepted)

cache = object.__new__(DynamicsCache)
cache._openmm_states = list(range(num_replicas))
cache._gcmc_states = list(range(num_replicas))
cache._state_moved = [False] * num_replicas
cache._num_swaps = np.zeros((num_replicas, num_replicas))
cache._states = states
cache.mix_states()

# Each state now holds the configuration that is favourable there.
for state, config in enumerate(cache._openmm_states):
assert energy_matrix[config, state] == -10000

# The swap matrix records each configuration's move.
for config in range(num_replicas):
assert cache._num_swaps[config, (config + 1) % num_replicas] == 1


@pytest.mark.parametrize(
"rest2_scale, is_valid",
[
Expand Down Expand Up @@ -735,11 +772,10 @@ def test_gcmc_state_follows_replica():

# Mix twice, since a slot is re-used within a cycle.
for states in ([2, 0, 3, 1], [1, 3, 0, 2]):
old_states = list(range(num_replicas))
expected = [cache._gcmc_states[state] for state in states]

cache._states = states
cache.mix_states(old_states)
cache.mix_states()

# The water occupancy follows the same permutation as the positions.
assert cache._gcmc_states == expected
Expand Down Expand Up @@ -831,9 +867,9 @@ def test_legacy_checkpoint_restore():
cache.__setstate__(dict(legacy))

# Converted to the current convention: each replica's own state, with the
# last mix applied.
assert cache._openmm_states == ["state2", "state0", "state1", "state3"]
assert cache._gcmc_states == ["water2", "water0", "water1", "water3"]
# last mix applied. The legacy states hold each configuration's destination.
assert cache._openmm_states == ["state1", "state2", "state0", "state3"]
assert cache._gcmc_states == ["water1", "water2", "water0", "water3"]

# Every replica is seeded from its stored state on a restart.
assert cache._state_moved == [True] * n
Expand Down