From 5af5755f32a8c5db6231119f7824e9431c5502ff Mon Sep 17 00:00:00 2001 From: Pardhav Maradani Date: Fri, 10 Jul 2026 19:32:54 +0530 Subject: [PATCH 1/7] Add widget for MSD analysis --- CHANGELOG.md | 1 + mdadash/backend/analyses/__init__.py | 3 +- mdadash/backend/analyses/msd.py | 168 +++++++++++++++++++++++++++ mdadash/backend/tests/test_server.py | 29 +++++ mdadash/backend/tests/utils.py | 6 +- 5 files changed, 203 insertions(+), 4 deletions(-) create mode 100644 mdadash/backend/analyses/msd.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 8e64e19..991110e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -31,6 +31,7 @@ The rules for this file: - Added pause support from within widgets (PR #27) - Added alerts support (PR #29) +- Added widget for MSD Analysis (PR #35) ### Fixed diff --git a/mdadash/backend/analyses/__init__.py b/mdadash/backend/analyses/__init__.py index 97daa12..6921d8f 100644 --- a/mdadash/backend/analyses/__init__.py +++ b/mdadash/backend/analyses/__init__.py @@ -2,7 +2,7 @@ Module that has all the analyses widgets """ -from . import com_distance, dssp, energies, janin, ramachandran, rog +from . import com_distance, dssp, energies, janin, msd, ramachandran, rog __all__ = [ "energies", @@ -11,4 +11,5 @@ "dssp", "ramachandran", "janin", + "msd", ] diff --git a/mdadash/backend/analyses/msd.py b/mdadash/backend/analyses/msd.py new file mode 100644 index 0000000..c00673b --- /dev/null +++ b/mdadash/backend/analyses/msd.py @@ -0,0 +1,168 @@ +import logging +from functools import partial + +import matplotlib.pyplot as plt +from IPython.display import display +from joblib import delayed +from MDAnalysis.analysis import msd +from tqdm import tqdm + +from mdadash.backend.widgets.base import WidgetBase + +logger = logging.getLogger(__name__) + +# `MDAnalysis.analysis.msd.EinsteinMSD` shows a progress bar for the simple +# and FFT cases. Disable this to prevent console output of the progress bar +tqdm.__init__ = partial(tqdm.__init__, disable=True) + + +class MSDAnalysis(WidgetBase): + name = "MSD Analysis" + description = "Mean squared displacement analysis" + + _inputs = [ + { + "attribute": "_run_mode", + "name": "Run mode", + "description": "The mode in which the widget is run", + "type": "select", + "items": [ + "serial", + "parallel", + ], + }, + { + "attribute": "selection", + "name": "Selection", + "description": "MDAnalysis selection phrase", + "type": "str", + }, + { + "attribute": "msd_type", + "name": "MSD type", + "description": "Desired dimensions to be included in the MSD", + "type": "select", + "items": [ + "xyz", + "xy", + "yz", + "xz", + "x", + "y", + "z", + ], + }, + { + "attribute": "fft", + "name": "FFT", + "description": "Use a fast FFT based computation", + "type": "bool", + }, + { + "attribute": "non_linear", + "name": "Non-linear", + "description": "Frames are non-linear", + "type": "bool", + }, + { + "attribute": "log_scale", + "name": "Log scale", + "description": "Use a log scale for the axes", + "type": "bool", + }, + { + "attribute": "custom_title", + "name": "Custom title", + "description": "Custom title for the plot", + "type": "str", + }, + ] + + def __init__(self): + super().__init__() + self.msd = None + self.selection = "all" + self.msd_type = "xyz" + self.fft = False + self.non_linear = False + self.log_scale = False + self.title = "MSD" + self.custom_title = None + self._setup_plot() + + def _setup_plot(self): + """Setup matplotlib plot""" + self.fig, self.ax = plt.subplots() + (self.plot,) = self.ax.plot([], []) + self.ax.set_xlabel("Lag time") + self.ax.set_ylabel("MSD") + self.ax.grid(True) + self._set_title() + self._set_axes_scale() + + def _set_title(self): + """Set plot title""" + self.ax.set_title(self.custom_title if self.custom_title else self.title) + + def _set_axes_scale(self): + """Set axes scale""" + self.ax.set_xscale("log" if self.log_scale else "linear") + self.ax.set_yscale("log" if self.log_scale else "linear") + + def _create_msd(self): + """Create msd instance""" + self.msd = msd.EinsteinMSD( + self.u, + select=self.selection, + msd_type=self.msd_type, + fft=self.fft, + non_linear=self.non_linear, + ) + self.title = f"MSD of '{self.selection}'" + self._set_title() + + def on_post_create(self): + """on_post_create handler""" + self._set_title() + self._set_axes_scale() + + def on_post_connect(self): + """on_post_connect handler""" + self._create_msd() + + def on_input_change(self, attribute, _old_value, new_value): + """on_input_change handler""" + if attribute == "custom_title": + self._set_title() + elif attribute == "log_scale": + self._set_axes_scale() + else: + self._create_msd() + + def _compute(self): + """Run MSD for the current timesteps window""" + self.msd.run() + return ( + self.msd.results.delta_t_values, + self.msd.results.timeseries, + ) + + def _update_plot(self, values): + """Update plot with computed values""" + self.plot.set_data(*values) + self.ax.relim() + self.ax.autoscale_view() + self.fig.canvas.draw() + display(self.fig) + + def run_every_frame(self): + """every-frame run handler""" + self._update_plot(self._compute()) + + def get_parallel_job(self, batch_size): + """get parallel job handler""" + return delayed(self._compute)() + + def apply_parallel_results(self, values): + """apply parallel results handler""" + self._update_plot(values) diff --git a/mdadash/backend/tests/test_server.py b/mdadash/backend/tests/test_server.py index fe7aa46..8bb7bb2 100644 --- a/mdadash/backend/tests/test_server.py +++ b/mdadash/backend/tests/test_server.py @@ -604,6 +604,35 @@ async def test_widget_run_janin(_client, imd_server): await disconnect_from_simulation() +async def test_widget_run_msd_serial(_client, imd_server): + uuid = await add_widget("MSD Analysis") + await connect_to_simulation(imd_server, step=1, batch_size=2) + inputs = [ + ("selection", "resid 1"), + ("custom_title", ""), + ("log_scale", False), + ] + await check_input_changes(uuid, inputs) + await resume_simulation(imd_server) + assert await sio_event_emitted(sio, "widgets:output", n=1) + await remove_widget(uuid) + await disconnect_from_simulation() + + +async def test_widget_run_msd_parallel(_client, imd_server): + uuid = await add_widget("MSD Analysis") + await connect_to_simulation(imd_server, step=1, batch_size=2) + inputs = [ + ("selection", "resid 1"), + ("_run_mode", "parallel"), + ] + await check_input_changes(uuid, inputs) + await resume_simulation(imd_server) + assert await sio_event_emitted(sio, "widgets:output", n=1) + await remove_widget(uuid) + await disconnect_from_simulation() + + def test_state_load(tmp_path): # test with no state file sm = StateManager("") diff --git a/mdadash/backend/tests/utils.py b/mdadash/backend/tests/utils.py index 72e73ac..e252adf 100644 --- a/mdadash/backend/tests/utils.py +++ b/mdadash/backend/tests/utils.py @@ -52,13 +52,13 @@ async def check_input_changes(uuid, inputs, status="ok"): assert response["status"] == status -async def connect_to_simulation(imd_server): +async def connect_to_simulation(imd_server, step=2, batch_size=1): main.mdadash.sm.universe_configs[0].update( { "topology": str(TPR), "trajectory": f"imd://localhost:{imd_server.port}", - "step": 2, - "batch_size": 1, + "step": step, + "batch_size": batch_size, } ) handler = sio.handlers["/"]["connect_to_simulations"] From ac2f2d0d8b57ced9d0e5bfad51507565c5d76cf9 Mon Sep 17 00:00:00 2001 From: Pardhav Maradani Date: Fri, 10 Jul 2026 19:49:37 +0530 Subject: [PATCH 2/7] Comment out tqdm disabling --- mdadash/backend/analyses/msd.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/mdadash/backend/analyses/msd.py b/mdadash/backend/analyses/msd.py index c00673b..f92ee9b 100644 --- a/mdadash/backend/analyses/msd.py +++ b/mdadash/backend/analyses/msd.py @@ -1,19 +1,19 @@ import logging -from functools import partial +# from functools import partial import matplotlib.pyplot as plt from IPython.display import display from joblib import delayed from MDAnalysis.analysis import msd -from tqdm import tqdm +# from tqdm import tqdm from mdadash.backend.widgets.base import WidgetBase logger = logging.getLogger(__name__) # `MDAnalysis.analysis.msd.EinsteinMSD` shows a progress bar for the simple # and FFT cases. Disable this to prevent console output of the progress bar -tqdm.__init__ = partial(tqdm.__init__, disable=True) +# tqdm.__init__ = partial(tqdm.__init__, disable=True) class MSDAnalysis(WidgetBase): From 274c10902b0573853090fcf5840fd2618df27fd3 Mon Sep 17 00:00:00 2001 From: Pardhav Maradani Date: Fri, 10 Jul 2026 20:27:19 +0530 Subject: [PATCH 3/7] Disable tqdm via env var due to MDA Bug #5144 --- mdadash/backend/analyses/msd.py | 6 ------ mdadash/backend/main.py | 2 ++ 2 files changed, 2 insertions(+), 6 deletions(-) diff --git a/mdadash/backend/analyses/msd.py b/mdadash/backend/analyses/msd.py index f92ee9b..29dca51 100644 --- a/mdadash/backend/analyses/msd.py +++ b/mdadash/backend/analyses/msd.py @@ -1,20 +1,14 @@ import logging -# from functools import partial import matplotlib.pyplot as plt from IPython.display import display from joblib import delayed from MDAnalysis.analysis import msd -# from tqdm import tqdm from mdadash.backend.widgets.base import WidgetBase logger = logging.getLogger(__name__) -# `MDAnalysis.analysis.msd.EinsteinMSD` shows a progress bar for the simple -# and FFT cases. Disable this to prevent console output of the progress bar -# tqdm.__init__ = partial(tqdm.__init__, disable=True) - class MSDAnalysis(WidgetBase): name = "MSD Analysis" diff --git a/mdadash/backend/main.py b/mdadash/backend/main.py index afa0a12..d86338d 100644 --- a/mdadash/backend/main.py +++ b/mdadash/backend/main.py @@ -19,6 +19,8 @@ logger = logging.getLogger(__name__) +os.environ["TQDM_DISABLE"] = "True" + @asynccontextmanager async def lifespan(_app: FastAPI): From 023257e3a348098ed1c9e98e02e6861a27a3ba7f Mon Sep 17 00:00:00 2001 From: Pardhav Maradani Date: Sat, 11 Jul 2026 21:15:08 +0530 Subject: [PATCH 4/7] - Remove use of MDAnalysis.analysis.msd.EinsteinMSD - Add a new SlidingWindowMSD class that is O(n) for each window --- mdadash/backend/analyses/msd.py | 103 ++++++++++++++++++++++---------- mdadash/backend/kernel/core.py | 10 +++- mdadash/backend/main.py | 2 - 3 files changed, 78 insertions(+), 37 deletions(-) diff --git a/mdadash/backend/analyses/msd.py b/mdadash/backend/analyses/msd.py index 29dca51..0ba9147 100644 --- a/mdadash/backend/analyses/msd.py +++ b/mdadash/backend/analyses/msd.py @@ -1,9 +1,11 @@ import logging +from collections import defaultdict, deque import matplotlib.pyplot as plt +import MDAnalysis as mda +import numpy as np from IPython.display import display from joblib import delayed -from MDAnalysis.analysis import msd from mdadash.backend.widgets.base import WidgetBase @@ -46,18 +48,6 @@ class MSDAnalysis(WidgetBase): "z", ], }, - { - "attribute": "fft", - "name": "FFT", - "description": "Use a fast FFT based computation", - "type": "bool", - }, - { - "attribute": "non_linear", - "name": "Non-linear", - "description": "Frames are non-linear", - "type": "bool", - }, { "attribute": "log_scale", "name": "Log scale", @@ -77,8 +67,6 @@ def __init__(self): self.msd = None self.selection = "all" self.msd_type = "xyz" - self.fft = False - self.non_linear = False self.log_scale = False self.title = "MSD" self.custom_title = None @@ -88,9 +76,9 @@ def _setup_plot(self): """Setup matplotlib plot""" self.fig, self.ax = plt.subplots() (self.plot,) = self.ax.plot([], []) - self.ax.set_xlabel("Lag time") - self.ax.set_ylabel("MSD") - self.ax.grid(True) + self.ax.set_xlabel(r"Lag time $\Delta$t (ps)") + self.ax.set_ylabel(r"MSD ($\AA^2$)") + self.ax.grid(True, linestyle="--", alpha=0.6) self._set_title() self._set_axes_scale() @@ -105,12 +93,10 @@ def _set_axes_scale(self): def _create_msd(self): """Create msd instance""" - self.msd = msd.EinsteinMSD( + self.msd = SlidingWindowMSD( self.u, select=self.selection, msd_type=self.msd_type, - fft=self.fft, - non_linear=self.non_linear, ) self.title = f"MSD of '{self.selection}'" self._set_title() @@ -130,20 +116,18 @@ def on_input_change(self, attribute, _old_value, new_value): self._set_title() elif attribute == "log_scale": self._set_axes_scale() + elif attribute == "_run_mode": + pass else: self._create_msd() - def _compute(self): + def _compute(self, parallel: bool = False): """Run MSD for the current timesteps window""" - self.msd.run() - return ( - self.msd.results.delta_t_values, - self.msd.results.timeseries, - ) + return self.msd.run(parallel=parallel) - def _update_plot(self, values): + def _update_plot(self, x_values, y_values): """Update plot with computed values""" - self.plot.set_data(*values) + self.plot.set_data(x_values, y_values) self.ax.relim() self.ax.autoscale_view() self.fig.canvas.draw() @@ -151,12 +135,67 @@ def _update_plot(self, values): def run_every_frame(self): """every-frame run handler""" - self._update_plot(self._compute()) + x_values, y_values, _ = self._compute() + self._update_plot(x_values, y_values) def get_parallel_job(self, batch_size): """get parallel job handler""" - return delayed(self._compute)() + return delayed(self._compute)(parallel=True) def apply_parallel_results(self, values): """apply parallel results handler""" - self._update_plot(values) + x_values, y_values, self.msd.msd_dict = values + self._update_plot(x_values, y_values) + + +class SlidingWindowMSD: + """Sliding Window MSD + + Calculate MSD for a sliding window of frames + + """ + + def __init__(self, u: mda.Universe, select: str = "all", msd_type: str = "xyz"): + self.u = u + self.select = select + self.msd_type = msd_type + self._parse_msd_type() + self.ag = u.select_atoms(self.select) + self.msd_dict = defaultdict(lambda: deque(maxlen=self.u.trajectory.buffer_size)) + self.msd_dict[0] = deque([0]) + + def _parse_msd_type(self): + """Sets up the desired dimensionality of the MSD.""" + keys = { + "x": [0], + "y": [1], + "z": [2], + "xy": [0, 1], + "xz": [0, 2], + "yz": [1, 2], + "xyz": [0, 1, 2], + } + self._dim = keys[self.msd_type.lower()] + + def run(self, parallel: bool = False) -> tuple: + """Run MSD for the current window""" + + time_current = self.u.trajectory.ts.time + positions_current = self.ag.positions[:, self._dim] + + for i in range(0, len(self.u.trajectory) - 1): + ts = self.u.trajectory[i] # set the buffered trajectory frame + delta_t = round(time_current - ts.time, 2) + disp = positions_current - self.ag.positions[:, self._dim] + squared_disp = np.sum(disp**2, axis=1) + msd = np.mean(squared_disp) + self.msd_dict[delta_t].append(msd) + + delta_t_values = sorted(self.msd_dict.keys()) + avg_msds = np.array([np.mean(self.msd_dict[dt]) for dt in delta_t_values]) + + return ( + delta_t_values, + avg_msds, + self.msd_dict if parallel else None, + ) diff --git a/mdadash/backend/kernel/core.py b/mdadash/backend/kernel/core.py index 0cb4021..e6ff831 100644 --- a/mdadash/backend/kernel/core.py +++ b/mdadash/backend/kernel/core.py @@ -65,10 +65,10 @@ class BufferedTrajectory: """ - def __init__(self, trajectory: mda.Universe.trajectory, batch_size: int): + def __init__(self, trajectory: mda.Universe.trajectory, buffer_size: int): self._trajectory = trajectory - self._batch_size = batch_size - self._buffer = deque(maxlen=batch_size) + self._buffer_size = buffer_size + self._buffer = deque(maxlen=buffer_size) self._buffer.append(trajectory.ts.copy()) BufferedTrajectory.next.__doc__ = type(trajectory).next.__doc__ @@ -76,6 +76,10 @@ def __init__(self, trajectory: mda.Universe.trajectory, batch_size: int): def n_frames(self): return len(self._buffer) + @property + def buffer_size(self): + return self._buffer_size + def __len__(self): return len(self._buffer) diff --git a/mdadash/backend/main.py b/mdadash/backend/main.py index d86338d..afa0a12 100644 --- a/mdadash/backend/main.py +++ b/mdadash/backend/main.py @@ -19,8 +19,6 @@ logger = logging.getLogger(__name__) -os.environ["TQDM_DISABLE"] = "True" - @asynccontextmanager async def lifespan(_app: FastAPI): From 0fd9de212d5a4b592cf28513228467b8c781bc34 Mon Sep 17 00:00:00 2001 From: Pardhav Maradani Date: Sun, 12 Jul 2026 08:04:12 +0530 Subject: [PATCH 5/7] - Update after merge from main --- mdadash/backend/analyses/msd.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mdadash/backend/analyses/msd.py b/mdadash/backend/analyses/msd.py index 0ba9147..2d35553 100644 --- a/mdadash/backend/analyses/msd.py +++ b/mdadash/backend/analyses/msd.py @@ -138,7 +138,7 @@ def run_every_frame(self): x_values, y_values, _ = self._compute() self._update_plot(x_values, y_values) - def get_parallel_job(self, batch_size): + def get_parallel_job(self): """get parallel job handler""" return delayed(self._compute)(parallel=True) From aa28e3529302cb58d4abdad60f1f042df05b0768 Mon Sep 17 00:00:00 2001 From: Pardhav Maradani Date: Tue, 14 Jul 2026 21:43:07 +0530 Subject: [PATCH 6/7] - Add support to compute and display particle MSDs - Remove windowed averages - Use pure numpy arrays instead of dicts and loops --- mdadash/backend/analyses/msd.py | 95 ++++++++++++++++++++++------ mdadash/backend/tests/test_server.py | 2 + mdadash/backend/widgets/base.py | 6 +- 3 files changed, 82 insertions(+), 21 deletions(-) diff --git a/mdadash/backend/analyses/msd.py b/mdadash/backend/analyses/msd.py index 2d35553..c7afb43 100644 --- a/mdadash/backend/analyses/msd.py +++ b/mdadash/backend/analyses/msd.py @@ -1,11 +1,11 @@ import logging -from collections import defaultdict, deque import matplotlib.pyplot as plt import MDAnalysis as mda import numpy as np from IPython.display import display from joblib import delayed +from matplotlib.collections import LineCollection from mdadash.backend.widgets.base import WidgetBase @@ -48,6 +48,12 @@ class MSDAnalysis(WidgetBase): "z", ], }, + { + "attribute": "show_particle_msds", + "name": "Show particle MSDs", + "description": "Show MSDs for individual particles of the selection", + "type": "bool", + }, { "attribute": "log_scale", "name": "Log scale", @@ -68,6 +74,7 @@ def __init__(self): self.selection = "all" self.msd_type = "xyz" self.log_scale = False + self.show_particle_msds = False self.title = "MSD" self.custom_title = None self._setup_plot() @@ -75,7 +82,11 @@ def __init__(self): def _setup_plot(self): """Setup matplotlib plot""" self.fig, self.ax = plt.subplots() - (self.plot,) = self.ax.plot([], []) + # use non-empty values to prevent initial exception + # if widget is configured to use log scale + (self.plot,) = self.ax.plot([1], [1], color="red", zorder=2) + self.lc = LineCollection([], colors="gray", alpha=0.2, lw=0.5, zorder=1) + self.ax.add_collection(self.lc) self.ax.set_xlabel(r"Lag time $\Delta$t (ps)") self.ax.set_ylabel(r"MSD ($\AA^2$)") self.ax.grid(True, linestyle="--", alpha=0.6) @@ -97,6 +108,7 @@ def _create_msd(self): self.u, select=self.selection, msd_type=self.msd_type, + show_particle_msds=self.show_particle_msds, ) self.title = f"MSD of '{self.selection}'" self._set_title() @@ -125,9 +137,10 @@ def _compute(self, parallel: bool = False): """Run MSD for the current timesteps window""" return self.msd.run(parallel=parallel) - def _update_plot(self, x_values, y_values): + def _update_plot(self, x, y1, y2): """Update plot with computed values""" - self.plot.set_data(x_values, y_values) + self.plot.set_data(x, y1) + self.lc.set_segments(y2 if self.show_particle_msds else []) self.ax.relim() self.ax.autoscale_view() self.fig.canvas.draw() @@ -135,8 +148,8 @@ def _update_plot(self, x_values, y_values): def run_every_frame(self): """every-frame run handler""" - x_values, y_values, _ = self._compute() - self._update_plot(x_values, y_values) + x, y1, y2, _ = self._compute() + self._update_plot(x, y1, y2) def get_parallel_job(self): """get parallel job handler""" @@ -144,8 +157,13 @@ def get_parallel_job(self): def apply_parallel_results(self, values): """apply parallel results handler""" - x_values, y_values, self.msd.msd_dict = values - self._update_plot(x_values, y_values) + x, y1, y2, (v1, v2, v3, v4) = values + self._update_plot(x, y1, y2) + # update msd state + self.msd.msd_sums = v1 + self.msd.msd_counts = v2 + self.msd.particle_msd_sums = v3 + self.msd.particle_msd_counts = v4 class SlidingWindowMSD: @@ -155,14 +173,28 @@ class SlidingWindowMSD: """ - def __init__(self, u: mda.Universe, select: str = "all", msd_type: str = "xyz"): + def __init__( + self, + u: mda.Universe, + select: str = "all", + msd_type: str = "xyz", + show_particle_msds: bool = False, + ): self.u = u self.select = select self.msd_type = msd_type + self.show_particle_msds = show_particle_msds self._parse_msd_type() self.ag = u.select_atoms(self.select) - self.msd_dict = defaultdict(lambda: deque(maxlen=self.u.trajectory.buffer_size)) - self.msd_dict[0] = deque([0]) + self.n_atoms = self.ag.atoms.n_atoms + self.n_lags = u.trajectory.buffer_size + self.msd_sums = np.zeros(self.n_lags) + self.msd_counts = np.zeros(self.n_lags, dtype=int) + self.msd_counts[0] = 1 + if self.show_particle_msds: + self.particle_msd_sums = np.zeros((self.n_lags, self.n_atoms)) + self.particle_msd_counts = np.zeros((self.n_lags, self.n_atoms), dtype=int) + self.particle_msd_counts[0, :] = 1 def _parse_msd_type(self): """Sets up the desired dimensionality of the MSD.""" @@ -180,22 +212,45 @@ def _parse_msd_type(self): def run(self, parallel: bool = False) -> tuple: """Run MSD for the current window""" - time_current = self.u.trajectory.ts.time + n = len(self.u.trajectory) # buffer / window might not be full yet positions_current = self.ag.positions[:, self._dim] - - for i in range(0, len(self.u.trajectory) - 1): - ts = self.u.trajectory[i] # set the buffered trajectory frame - delta_t = round(time_current - ts.time, 2) + for i in range(n - 1): + lag = n - 1 - i + _ = self.u.trajectory[i] # set the buffered trajectory frame disp = positions_current - self.ag.positions[:, self._dim] squared_disp = np.sum(disp**2, axis=1) msd = np.mean(squared_disp) - self.msd_dict[delta_t].append(msd) + self.msd_sums[lag] += msd + self.msd_counts[lag] += 1 + if self.show_particle_msds: + self.particle_msd_sums[lag, :] += squared_disp + self.particle_msd_counts[lag, :] += 1 - delta_t_values = sorted(self.msd_dict.keys()) - avg_msds = np.array([np.mean(self.msd_dict[dt]) for dt in delta_t_values]) + # We will have at least 2 frames by the time we are here. + # frame_dt will ensure the delta_t is correct even if we have step + # value (other than 1) configured in the universe configuration + frame_dt = round(self.u.trajectory[1].time - self.u.trajectory[0].time, 2) + delta_t_values = np.arange(n) * frame_dt + avg_msds = self.msd_sums[:n] / self.msd_counts[:n] + msds_by_particle_lines = None + if self.show_particle_msds: + msds_by_particle_array = ( + self.particle_msd_sums[:n, :] / self.particle_msd_counts[:n, :] + ) + msds_by_particle_lines = np.empty((self.n_atoms, n, 2)) + msds_by_particle_lines[:, :, 0] = delta_t_values + msds_by_particle_lines[:, :, 1] = msds_by_particle_array.T return ( delta_t_values, avg_msds, - self.msd_dict if parallel else None, + msds_by_particle_lines, + ( + self.msd_sums, + self.msd_counts, + self.particle_msd_sums, + self.particle_msd_counts, + ) + if parallel + else (None,) * 4, ) diff --git a/mdadash/backend/tests/test_server.py b/mdadash/backend/tests/test_server.py index 52dbf71..3f30970 100644 --- a/mdadash/backend/tests/test_server.py +++ b/mdadash/backend/tests/test_server.py @@ -610,6 +610,7 @@ async def test_widget_run_msd_serial(_client, imd_server): inputs = [ ("selection", "resid 1"), ("custom_title", ""), + ("show_particle_msds", True), ("log_scale", False), ] await check_input_changes(uuid, inputs) @@ -624,6 +625,7 @@ async def test_widget_run_msd_parallel(_client, imd_server): await connect_to_simulation(imd_server, step=1, batch_size=2) inputs = [ ("selection", "resid 1"), + ("show_particle_msds", True), ("_run_mode", "parallel"), ] await check_input_changes(uuid, inputs) diff --git a/mdadash/backend/widgets/base.py b/mdadash/backend/widgets/base.py index 461f9a0..28bea4f 100644 --- a/mdadash/backend/widgets/base.py +++ b/mdadash/backend/widgets/base.py @@ -556,8 +556,12 @@ def _run_parallel_jobs(self, parallel_widgets, parallel_results): func, args, kwargs = widget.get_parallel_job() parallel_jobs.append((self.with_reset_frame, (func,) + args, kwargs)) try: + # without max_nbytes=None, np arrays passed / returned + # are marked read-only in subsequent calls (eg: msd case) results = Parallel( - n_jobs=self.n_jobs, initializer=WidgetManager._patch_IMDReader + n_jobs=self.n_jobs, + max_nbytes=None, + initializer=WidgetManager._patch_IMDReader, )(parallel_jobs) parallel_results.extend(results) # pylint: disable=broad-exception-caught From 284226acb00b6f6d67ad24ec477b16f293fa6149 Mon Sep 17 00:00:00 2001 From: Pardhav Maradani Date: Wed, 15 Jul 2026 20:06:25 +0530 Subject: [PATCH 7/7] - Fix state re-assignment in parallel mode --- mdadash/backend/analyses/msd.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/mdadash/backend/analyses/msd.py b/mdadash/backend/analyses/msd.py index c7afb43..a9988cc 100644 --- a/mdadash/backend/analyses/msd.py +++ b/mdadash/backend/analyses/msd.py @@ -162,8 +162,9 @@ def apply_parallel_results(self, values): # update msd state self.msd.msd_sums = v1 self.msd.msd_counts = v2 - self.msd.particle_msd_sums = v3 - self.msd.particle_msd_counts = v4 + if self.show_particle_msds: + self.msd.particle_msd_sums = v3 + self.msd.particle_msd_counts = v4 class SlidingWindowMSD: @@ -248,8 +249,8 @@ def run(self, parallel: bool = False) -> tuple: ( self.msd_sums, self.msd_counts, - self.particle_msd_sums, - self.particle_msd_counts, + self.particle_msd_sums if self.show_particle_msds else None, + self.particle_msd_counts if self.show_particle_msds else None, ) if parallel else (None,) * 4,