Skip to content

Commit 7310196

Browse files
Add batching support for Energy widgets
1 parent 09340d2 commit 7310196

3 files changed

Lines changed: 71 additions & 13 deletions

File tree

docs/source/built_in_widgets.rst

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -34,7 +34,7 @@ and can run in parallel:
3434
- ✅
3535
* - :mod:`~mdadash.backend.analyses.energies`
3636
- Widgets for various simulation energies
37-
-
37+
-
3838
- —
3939
* - :mod:`~mdadash.backend.analyses.helix_analysis`
4040
- Helix Analysis

mdadash/backend/analyses/energies.py

Lines changed: 56 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -44,10 +44,14 @@ class EnergyWidgetBase:
4444
4545
**Inputs**
4646
47-
Max values
47+
Run frequency
4848
.. compound::
49-
Max values to show in plot
50-
Default: ``100``
49+
The frequency with which the widget is run - `every-frame` or `batch`
50+
Default: ``every-frame``
51+
52+
Max values
53+
Max values to show in plot
54+
Default: ``100``
5155
5256
Title
5357
Title for the plot
@@ -87,6 +91,16 @@ class EnergyWidgetBase:
8791
)
8892

8993
_inputs: ClassVar = [
94+
{
95+
"attribute": "_run_frequency",
96+
"name": "Run frequency",
97+
"description": "The frequency with which the widget is run",
98+
"type": "select",
99+
"items": [
100+
"every-frame",
101+
"batch",
102+
],
103+
},
90104
{
91105
"attribute": "maxlen",
92106
"name": "Max values",
@@ -154,6 +168,10 @@ def on_post_create(self):
154168
self._set_title()
155169
self._reset_plot_values()
156170

171+
def on_post_connect(self):
172+
""":meth:`~mdadash.backend.widgets.base.WidgetBase.on_post_connect` handler"""
173+
self._update_plot(self._compute_current_frame())
174+
157175
def on_input_change(self, attribute, _old_value, new_value):
158176
""":meth:`~mdadash.backend.widgets.base.WidgetBase.on_input_change` handler"""
159177
if attribute == "maxlen":
@@ -165,21 +183,48 @@ def on_input_change(self, attribute, _old_value, new_value):
165183
elif attribute == "x_type":
166184
self._set_x_values()
167185

168-
def run_every_frame(self):
169-
""":meth:`~mdadash.backend.widgets.base.WidgetBase.run_every_frame` handler"""
186+
def _compute_current_frame(self):
187+
"""Compute for current frame"""
170188
ts = self.u.trajectory.ts # pylint: disable=no-member
171-
if self.data_key not in ts.data:
172-
return # pragma: no cover
173-
self.steps.append(ts.data["step"])
174-
self.times.append(ts.data["time"])
175-
self.y_values.append(ts.data[self.data_key])
176-
# update plot
189+
return (
190+
ts.data["step"],
191+
ts.data["time"],
192+
ts.data.get(self.data_key),
193+
)
194+
195+
def _compute_batch(self):
196+
"""Compute for current batch"""
197+
u = self.u # pylint: disable=no-member
198+
values = []
199+
for i in range(u.trajectory.buffer_size):
200+
_ = u.trajectory[i]
201+
values.append(self._compute_current_frame())
202+
return values
203+
204+
def _update_plot(self, values):
205+
"""Append values and update plot"""
206+
if isinstance(values, tuple):
207+
values = [values]
208+
for value in values:
209+
(steps, times, v) = value
210+
self.steps.append(steps)
211+
self.times.append(times)
212+
self.y_values.append(v)
213+
# update plot points
177214
self.plot.set_data(self.x_values, self.y_values)
178215
self.ax.relim()
179216
self.ax.autoscale_view()
180217
self.fig.canvas.draw()
181218
display(self.fig)
182219

220+
def run_every_frame(self):
221+
""":meth:`~mdadash.backend.widgets.base.WidgetBase.run_every_frame` handler"""
222+
self._update_plot(self._compute_current_frame())
223+
224+
def run_batch(self):
225+
""":meth:`~mdadash.backend.widgets.base.WidgetBase.run_batch` handler"""
226+
self._update_plot(self._compute_batch())
227+
183228

184229
class AbsoluteTemperature(EnergyWidgetBase, WidgetBase):
185230
"""Absolute Temperature

mdadash/backend/tests/test_server.py

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -419,7 +419,7 @@ async def test_widget_invalid_inputs(_client, imd_server):
419419
await disconnect_from_simulation()
420420

421421

422-
async def test_widget_run_energies(_client, imd_server):
422+
async def test_widget_run_energies_serial(_client, imd_server):
423423
uuid = await add_widget("Absolute Temperature")
424424
await connect_to_simulation(imd_server)
425425
await resume_simulation(imd_server)
@@ -428,6 +428,19 @@ async def test_widget_run_energies(_client, imd_server):
428428
await disconnect_from_simulation()
429429

430430

431+
async def test_widget_run_energies_batch(_client, imd_server):
432+
uuid = await add_widget("Absolute Temperature")
433+
await connect_to_simulation(imd_server)
434+
inputs = [
435+
("_run_frequency", "batch"),
436+
]
437+
await check_input_changes(uuid, inputs)
438+
await resume_simulation(imd_server)
439+
assert await sio_event_emitted(sio, "widgets:output", n=1)
440+
await remove_widget(uuid)
441+
await disconnect_from_simulation()
442+
443+
431444
async def test_widget_run_com_distance_serial_every_frame(_client, imd_server):
432445
await connect_to_simulation(imd_server)
433446
uuid = await add_widget("COMDistance")

0 commit comments

Comments
 (0)