From 7241025babe4ee09af098eccffa5e7f0311ff5e7 Mon Sep 17 00:00:00 2001 From: clonker <1685266+clonker@users.noreply.github.com> Date: Sun, 29 Mar 2026 17:38:12 +0200 Subject: [PATCH] Fix restrict_to_submodel iterating over therm states instead of trajectories When ttrajs are provided (replica exchange), n_therm_states can differ from len(dtrajs). The loop in restrict_to_submodel was using n_therm_states, which would skip trajectories or cause an IndexError. --- src/deeptime/markov/msm/tram/_tram_dataset.py | 2 +- tests/markov/msm/test_tram_datatset.py | 25 +++++++++++++++++++ 2 files changed, 26 insertions(+), 1 deletion(-) diff --git a/src/deeptime/markov/msm/tram/_tram_dataset.py b/src/deeptime/markov/msm/tram/_tram_dataset.py index bad6e97f..76c881db 100644 --- a/src/deeptime/markov/msm/tram/_tram_dataset.py +++ b/src/deeptime/markov/msm/tram/_tram_dataset.py @@ -340,7 +340,7 @@ def restrict_to_submodel(self, submodel): if isinstance(submodel, list) or isinstance(submodel, np.ndarray): submodel = self._submodel_from_states(submodel) - for k in range(self.n_therm_states): + for k in range(len(self.dtrajs)): # Get largest connected set # Assign -1 to all indices not in the submodel. restricted_dtraj = submodel.transform_discrete_trajectories_to_submodel(self.dtrajs[k]) diff --git a/tests/markov/msm/test_tram_datatset.py b/tests/markov/msm/test_tram_datatset.py index eb5744d2..8cdb82ac 100644 --- a/tests/markov/msm/test_tram_datatset.py +++ b/tests/markov/msm/test_tram_datatset.py @@ -155,6 +155,31 @@ def test_restrict_to_submodel_with_indices_input(test_input, submodel, expected) np.testing.assert_equal(tram_data.dtrajs, expected) +def test_restrict_to_submodel_with_ttrajs(): + # 3 trajectories but only 2 thermodynamic states (replica exchange scenario). + # Previously this would fail because restrict_to_submodel iterated over + # n_therm_states instead of len(dtrajs), skipping the third trajectory. + dtrajs = [np.asarray([0, 1, 2, 3, 1]), + np.asarray([2, 3, 2, 1, 0]), + np.asarray([1, 2, 3, 0, 1])] + ttrajs = [np.asarray([0, 0, 1, 1, 0]), + np.asarray([1, 1, 0, 0, 1]), + np.asarray([0, 1, 1, 0, 0])] + bias_matrices = make_matching_bias_matrix(dtrajs, n_therm_states=2) + tram_data = TRAMDataset(dtrajs=dtrajs, ttrajs=ttrajs, bias_matrices=bias_matrices) + + assert tram_data.n_therm_states == 2 + assert len(tram_data.dtrajs) == 3 + + tram_data.restrict_to_submodel([1, 2, 3]) + + # All 3 trajectories should be restricted: state 0 becomes -1 + assert len(tram_data.dtrajs) == 3 + for dtraj in tram_data.dtrajs: + assert 0 not in dtraj + assert -1 in dtraj + + @pytest.mark.parametrize( "lagtime", [1, 3] )