From 9cdf0e3caba65f91337e168dd31f6e486b5341fa Mon Sep 17 00:00:00 2001 From: Lester Hedges Date: Tue, 29 Sep 2026 21:45:27 +0100 Subject: [PATCH] Backport fix from PR #232. [ci skip] --- src/somd2/runner/_repex.py | 30 ++++++++++++------------- tests/runner/test_repex.py | 46 +++++++++++++++++++++++++++++++++----- 2 files changed, 55 insertions(+), 21 deletions(-) diff --git a/src/somd2/runner/_repex.py b/src/somd2/runner/_repex.py index 8d95e97..828739b 100644 --- a/src/somd2/runner/_repex.py +++ b/src/somd2/runner/_repex.py @@ -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. @@ -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. @@ -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 @@ -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): """ @@ -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, @@ -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 @@ -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 @@ -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): """ diff --git a/tests/runner/test_repex.py b/tests/runner/test_repex.py index 1305be8..c9232d6 100644 --- a/tests/runner/test_repex.py +++ b/tests/runner/test_repex.py @@ -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", [ @@ -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 @@ -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