Skip to content

Commit d3eb0e5

Browse files
Merge branch 'main' into built-in-widgets-docs
2 parents 23c36a4 + 578534f commit d3eb0e5

3 files changed

Lines changed: 119 additions & 14 deletions

File tree

CHANGELOG.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@ The rules for this file:
3030
<!-- New added features -->
3131

3232
- Added built-in widgets documentation (PR #63)
33+
- Added batching and parallel support for com distance widget (PR #64)
3334

3435
### Fixed
3536

mdadash/backend/analyses/com_distance.py

Lines changed: 74 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
import matplotlib.pyplot as plt
1010
from IPython.display import display
11+
from joblib import delayed
1112
from MDAnalysis.exceptions import NoDataError
1213
from MDAnalysis.lib.distances import calc_bonds
1314

@@ -79,6 +80,26 @@ class COMDistance(WidgetBase):
7980
description = "Distance between two COMs"
8081

8182
_inputs: ClassVar = [
83+
{
84+
"attribute": "_run_frequency",
85+
"name": "Run frequency",
86+
"description": "The frequency with which the widget is run",
87+
"type": "select",
88+
"items": [
89+
"every-frame",
90+
"batch",
91+
],
92+
},
93+
{
94+
"attribute": "_run_mode",
95+
"name": "Run mode",
96+
"description": "The mode in which the widget is run",
97+
"type": "select",
98+
"items": [
99+
"serial",
100+
"parallel",
101+
],
102+
},
82103
{
83104
"attribute": "selection1",
84105
"name": "Selection 1",
@@ -234,27 +255,68 @@ def on_input_change(self, attribute, _old_value, new_value):
234255
if reset_plot:
235256
self._reset_plot_values()
236257

237-
def run_every_frame(self):
238-
"""every-frame run handler"""
258+
def _compute_current_frame(self):
259+
"""Compute for current frame"""
239260
try:
240261
com1 = self.ag1.center_of_mass(unwrap=True)
241262
com2 = self.ag2.center_of_mass(unwrap=True)
242263
except NoDataError: # pragma: no cover
243264
# unwrap can fail if there is no bonds info
244265
com1 = self.ag1.center_of_mass()
245266
com2 = self.ag2.center_of_mass()
246-
dist = calc_bonds(com1, com2, box=self.u.dimensions)
247-
self.y_values.append(dist)
248-
self.steps.append(self.u.trajectory.ts.data["step"])
249-
self.times.append(self.u.trajectory.ts.data["time"])
250-
# update plot
267+
return (
268+
self.u.trajectory.ts.data["step"],
269+
self.u.trajectory.ts.data["time"],
270+
calc_bonds(com1, com2, box=self.u.dimensions),
271+
)
272+
273+
def _compute_batch(self):
274+
"""Compute for current batch"""
275+
values = []
276+
for i in range(self.u.trajectory.buffer_size):
277+
_ = self.u.trajectory[i]
278+
values.append(self._compute_current_frame())
279+
return values
280+
281+
def _update_plot(self, values):
282+
"""Append values and update plot"""
283+
if isinstance(values, tuple):
284+
values = [values]
285+
alerted = False
286+
paused = False
287+
for value in values:
288+
(steps, times, dist) = value
289+
self.steps.append(steps)
290+
self.times.append(times)
291+
self.y_values.append(dist)
292+
if dist > self.max_distance:
293+
if self.max_distance_alert and not alerted:
294+
self.alert(f"Distance between '{self.title}' > {self.max_distance}")
295+
alerted = True
296+
if self.max_distance_pause and not paused:
297+
self.pause_simulation()
298+
paused = True
299+
# update plot points
251300
self.plot.set_data(self.x_values, self.y_values)
252301
self.ax.relim()
253302
self.ax.autoscale_view()
254303
self.fig.canvas.draw()
255304
display(self.fig)
256-
if dist > self.max_distance:
257-
if self.max_distance_alert:
258-
self.alert(f"Distance between '{self.title}' > {self.max_distance}")
259-
if self.max_distance_pause:
260-
self.pause_simulation()
305+
306+
def run_every_frame(self):
307+
"""every-frame run handler"""
308+
self._update_plot(self._compute_current_frame())
309+
310+
def run_batch(self):
311+
"""batch run handler"""
312+
self._update_plot(self._compute_batch())
313+
314+
def get_parallel_job(self):
315+
"""get parallel job handler"""
316+
if self._run_frequency == "batch":
317+
return delayed(self._compute_batch)()
318+
return delayed(self._compute_current_frame)()
319+
320+
def apply_parallel_results(self, values):
321+
"""apply parallel results handler"""
322+
self._update_plot(values)

mdadash/backend/tests/test_server.py

Lines changed: 44 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -428,9 +428,9 @@ async def test_widget_run_energies(_client, imd_server):
428428
await disconnect_from_simulation()
429429

430430

431-
async def test_widget_run_com_distance(_client, imd_server):
432-
uuid = await add_widget("COMDistance")
431+
async def test_widget_run_com_distance_serial_every_frame(_client, imd_server):
433432
await connect_to_simulation(imd_server)
433+
uuid = await add_widget("COMDistance")
434434
inputs = [
435435
("selection1", "resid 1"),
436436
("selection2", "resid 2"),
@@ -446,6 +446,48 @@ async def test_widget_run_com_distance(_client, imd_server):
446446
await disconnect_from_simulation()
447447

448448

449+
async def test_widget_run_com_distance_serial_batch(_client, imd_server):
450+
uuid = await add_widget("COMDistance")
451+
await connect_to_simulation(imd_server)
452+
inputs = [
453+
("_run_frequency", "batch"),
454+
]
455+
await check_input_changes(uuid, inputs)
456+
await resume_simulation(imd_server)
457+
assert await sio_event_emitted(sio, "widgets:output", n=1)
458+
await remove_widget(uuid)
459+
await disconnect_from_simulation()
460+
461+
462+
async def test_widget_run_com_distance_parallel_every_frame(_client, imd_server):
463+
uuid = await add_widget("COMDistance")
464+
inputs = [
465+
("_run_mode", "parallel"),
466+
]
467+
await check_input_changes(uuid, inputs)
468+
await connect_to_simulation(imd_server)
469+
await resume_simulation(imd_server)
470+
timeout = 30 if sys.platform == "win32" else 20
471+
assert await sio_event_emitted(sio, "widgets:output", n=1, timeout=timeout)
472+
await remove_widget(uuid)
473+
await disconnect_from_simulation()
474+
475+
476+
async def test_widget_run_com_distance_parallel_batch(_client, imd_server):
477+
uuid = await add_widget("COMDistance")
478+
inputs = [
479+
("_run_frequency", "batch"),
480+
("_run_mode", "parallel"),
481+
]
482+
await check_input_changes(uuid, inputs)
483+
await connect_to_simulation(imd_server)
484+
await resume_simulation(imd_server)
485+
timeout = 30 if sys.platform == "win32" else 20
486+
assert await sio_event_emitted(sio, "widgets:output", n=1, timeout=timeout)
487+
await remove_widget(uuid)
488+
await disconnect_from_simulation()
489+
490+
449491
async def test_widget_run_com_distance_alert_pause(_client, imd_server):
450492
uuid = await add_widget("COMDistance")
451493
await connect_to_simulation(imd_server)

0 commit comments

Comments
 (0)