diff --git a/package/CHANGELOG b/package/CHANGELOG index 45647f3904..524f3918cb 100644 --- a/package/CHANGELOG +++ b/package/CHANGELOG @@ -23,6 +23,8 @@ The rules for this file: * 2.11.0 Fixes + * `AnalysisBase.run()` now interprets NumPy boolean arrays passed to `frames` + as frame masks instead of integer frame indices (Issue #5472) * Mass (u) has been added to MDAnalysis base units and clarified that physical constants use CODATA 2010 values (Issue #3944, PR #5439) * Added `.gitattributes` to enforce LF (\n) as line endings and renormalized diff --git a/package/MDAnalysis/analysis/base.py b/package/MDAnalysis/analysis/base.py index f18b866951..8ab19be6db 100644 --- a/package/MDAnalysis/analysis/base.py +++ b/package/MDAnalysis/analysis/base.py @@ -616,7 +616,7 @@ def _setup_computation_groups( else: used_frames = frames - if all(isinstance(obj, bool) for obj in used_frames): + if all(isinstance(obj, (bool, np.bool_)) for obj in used_frames): arange = np.arange(len(used_frames)) used_frames = arange[used_frames] @@ -781,8 +781,10 @@ def run( step : int, optional number of frames to skip between each analysed frame frames : array_like, optional - array of integers or booleans to slice trajectory; ``frames`` can - only be used *instead* of ``start``, ``stop``, and ``step``. Setting + array of integers or booleans to slice trajectory. Boolean arrays, + including NumPy arrays, can be used as masks for fancy indexing. + ``frames`` can only be used *instead* of ``start``, ``stop``, and + ``step``. Setting *both* ``frames`` and at least one of ``start``, ``stop``, ``step`` to a non-default value will raise a :exc:`ValueError`. @@ -831,6 +833,10 @@ def run( Introduced ``backend``, ``n_workers``, ``n_parts`` and ``unsupported_backend`` keywords, and refactored the method logic to support parallelizable execution. + + .. versionchanged:: 2.11.0 + NumPy boolean arrays passed to ``frames`` are now interpreted as + boolean masks. """ # default to serial execution backend = "serial" if backend is None else backend diff --git a/testsuite/MDAnalysisTests/analysis/test_base.py b/testsuite/MDAnalysisTests/analysis/test_base.py index 0dd872bde5..7fea26f1d5 100644 --- a/testsuite/MDAnalysisTests/analysis/test_base.py +++ b/testsuite/MDAnalysisTests/analysis/test_base.py @@ -337,6 +337,27 @@ def test_start_stop_step(u, run_kwargs, frames): }, (0, 1, 3, 5, 6, 8), ), + pytest.param( + { + "frames": np.array( + [ + True, + True, + False, + True, + False, + True, + True, + False, + True, + False, + ], + dtype=bool, + ) + }, + (0, 1, 3, 5, 6, 8), + id="numpy-bool-mask", + ), ], ) def test_frame_slice(u_xtc, run_kwargs, frames): @@ -370,6 +391,27 @@ def test_frame_slice(u_xtc, run_kwargs, frames): }, (0, 1, 3, 5, 6, 8), ), + pytest.param( + { + "frames": np.array( + [ + True, + True, + False, + True, + False, + True, + True, + False, + True, + False, + ], + dtype=bool, + ) + }, + (0, 1, 3, 5, 6, 8), + id="numpy-bool-mask", + ), ], ) def test_frame_slice_parallel(run_kwargs, frames, client_FrameAnalysis):