Skip to content

Commit 266676f

Browse files
Cleanup of widget invocation (#36)
- Remove passing redundant batch_size param - The buffer / batch size is available in trajectory - Reset frame to the latest timestep before each widget invocation - Widget runs can iterate trajectory and hence required
1 parent a840cf4 commit 266676f

6 files changed

Lines changed: 36 additions & 40 deletions

File tree

mdadash/backend/analyses/dssp.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -154,7 +154,7 @@ def _compute_current_frame(self):
154154
self.u.trajectory.ts.data["time"],
155155
)
156156

157-
def _compute_batch(self, _batch_size):
157+
def _compute_batch(self):
158158
"""Compute values for current batch"""
159159
self.dssp.run()
160160
# `run()` can also be invoked with a non-serial backend
@@ -206,14 +206,14 @@ def run_every_frame(self):
206206
"""every-frame run handler"""
207207
self._update_plot(self._compute_current_frame())
208208

209-
def run_batch(self, batch_size):
209+
def run_batch(self):
210210
"""batch run handler"""
211-
self._update_plot(self._compute_batch(batch_size))
211+
self._update_plot(self._compute_batch())
212212

213-
def get_parallel_job(self, batch_size):
213+
def get_parallel_job(self):
214214
"""get parallel job handler"""
215215
if self._run_frequency == "batch":
216-
return delayed(self._compute_batch)(batch_size)
216+
return delayed(self._compute_batch)()
217217
return delayed(self._compute_current_frame)()
218218

219219
def apply_parallel_results(self, values):

mdadash/backend/analyses/rog.py

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -194,10 +194,10 @@ def _compute_current_frame(self):
194194
rog,
195195
)
196196

197-
def _compute_batch(self, batch_size):
197+
def _compute_batch(self):
198198
"""Compute ROG values for current batch"""
199199
values = []
200-
for i in range(batch_size):
200+
for i in range(self.u.trajectory.buffer_size):
201201
_ = self.u.trajectory[i]
202202
values.append(self._compute_current_frame())
203203
return values
@@ -225,14 +225,14 @@ def run_every_frame(self):
225225
"""every-frame run handler"""
226226
self._update_plot(self._compute_current_frame())
227227

228-
def run_batch(self, batch_size):
228+
def run_batch(self):
229229
"""batch run handler"""
230-
self._update_plot(self._compute_batch(batch_size))
230+
self._update_plot(self._compute_batch())
231231

232-
def get_parallel_job(self, batch_size):
232+
def get_parallel_job(self):
233233
"""get parallel job handler"""
234234
if self._run_frequency == "batch":
235-
return delayed(self._compute_batch)(batch_size)
235+
return delayed(self._compute_batch)()
236236
return delayed(self._compute_current_frame)()
237237

238238
def apply_parallel_results(self, values):

mdadash/backend/kernel/core.py

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -65,17 +65,21 @@ class BufferedTrajectory:
6565
6666
"""
6767

68-
def __init__(self, trajectory: mda.Universe.trajectory, batch_size: int):
68+
def __init__(self, trajectory: mda.Universe.trajectory, buffer_size: int):
6969
self._trajectory = trajectory
70-
self._batch_size = batch_size
71-
self._buffer = deque(maxlen=batch_size)
70+
self._buffer_size = buffer_size
71+
self._buffer = deque(maxlen=buffer_size)
7272
self._buffer.append(trajectory.ts.copy())
7373
BufferedTrajectory.next.__doc__ = type(trajectory).next.__doc__
7474

7575
@property
7676
def n_frames(self):
7777
return len(self._buffer)
7878

79+
@property
80+
def buffer_size(self):
81+
return self._buffer_size
82+
7983
def __len__(self):
8084
return len(self._buffer)
8185

@@ -389,7 +393,7 @@ async def _iter_loop(self):
389393
batch_ready = (u.trajectory._frame + step) % (
390394
step * batch_size
391395
) == 0
392-
self._wm.run_widgets(uid, batch_ready, batch_size)
396+
self._wm.run_widgets(uid, batch_ready)
393397
# pylint: disable=broad-exception-caught
394398
except Exception: # pragma: no cover
395399
logger.exception("Trajectory iteration failed for uid %d", uid)

mdadash/backend/tests/test_server.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -234,7 +234,7 @@ def run_every_frame(self):
234234
class _TestWidget4b(WidgetBase):
235235
name = "TestWidget4b"
236236

237-
def run_batch(self, batch_size):
237+
def run_batch(self):
238238
pass
239239

240240
# test duplicate widget name registraion exception

mdadash/backend/widgets/base.py

Lines changed: 14 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -55,6 +55,10 @@ def _set_universe(self, u: mda.Universe):
5555
"""Internal: Set the universe"""
5656
self.u = u
5757

58+
def _reset_frame_latest(self):
59+
"""Internal: Reset frame to latest timestep"""
60+
_ = self.u.trajectory[-1]
61+
5862
def _get_inputs(self):
5963
"""Internal: Get the current instance inputs"""
6064
inputs = getattr(self, "_inputs")
@@ -172,33 +176,24 @@ def run_every_frame(self) -> None:
172176
173177
"""
174178

175-
def run_batch(self, batch_size: int) -> None:
179+
def run_batch(self) -> None:
176180
"""run_batch handler
177181
178182
This handler is called every time a new batch of timesteps is full
179183
and ready to be run if the run frequency is set to `batch`
180184
(`_run_frequency='batch'`).
181185
182-
Parameters
183-
----------
184-
batch_size: int
185-
The batch size, which indicates the number of buffered timesteps
186-
available in the buffer is passed
186+
`self.u.trajectory.buffer_size` is the size of the buffer / batch
187+
that can be used by the widget class.
187188
188189
"""
189190

190-
def get_parallel_job(self, batch_size: int) -> Any:
191+
def get_parallel_job(self) -> Any:
191192
"""get_parallel_job handler
192193
193194
This handler is called if the run mode is set to `parallel`
194195
(`_run_mode='parallel'`) to get the parallel job to run.
195196
196-
Parameters
197-
----------
198-
batch_size: int
199-
The configured batch size to use in the parallel jobs
200-
if their run frequency is `batch`
201-
202197
Returns
203198
-------
204199
job: Any
@@ -272,7 +267,7 @@ def _validate_widget_class(cls, widget_class: WidgetBase) -> None:
272267
# check for one of the run methods to exist with correct params
273268
run_methods = {
274269
"run_every_frame": 1,
275-
"run_batch": 2,
270+
"run_batch": 1,
276271
}
277272
has_valid_run_method = False
278273
for run_method, n_params in run_methods.items():
@@ -546,11 +541,11 @@ def custom_setstate(self, state):
546541
IMDReader.__setstate__ = custom_setstate
547542
IMDReader.__getstate__ = custom_getstate
548543

549-
def _run_parallel_jobs(self, parallel_widgets, batch_size, parallel_results):
544+
def _run_parallel_jobs(self, parallel_widgets, parallel_results):
550545
"""Internal: Run parallel jobs using joblib.Parallel"""
551546
parallel_jobs = []
552547
for widget in parallel_widgets:
553-
parallel_jobs.append(widget.get_parallel_job(batch_size))
548+
parallel_jobs.append(widget.get_parallel_job())
554549
try:
555550
results = Parallel(
556551
n_jobs=self.n_jobs, initializer=WidgetManager._patch_IMDReader
@@ -560,7 +555,7 @@ def _run_parallel_jobs(self, parallel_widgets, batch_size, parallel_results):
560555
except Exception: # pragma: no cover
561556
logger.exception("Parallel run failed for jobs %s", parallel_jobs)
562557

563-
def run_widgets(self, uid: int, batch_ready: bool, batch_size: int) -> None:
558+
def run_widgets(self, uid: int, batch_ready: bool) -> None:
564559
"""Run widget instances
565560
566561
Parameters
@@ -571,9 +566,6 @@ def run_widgets(self, uid: int, batch_ready: bool, batch_size: int) -> None:
571566
batch_ready: bool
572567
Flag indicating if a batch of timesteps is full
573568
574-
batch_size: int
575-
Size of the batch
576-
577569
"""
578570
# collect widgets that need to be run
579571
parallel_widgets = []
@@ -594,19 +586,19 @@ def run_widgets(self, uid: int, batch_ready: bool, batch_size: int) -> None:
594586
target=self._run_parallel_jobs,
595587
args=(
596588
parallel_widgets,
597-
batch_size,
598589
parallel_results,
599590
),
600591
)
601592
parallel_thread.start()
602593
# run serial widgets
603594
for widget in serial_widgets:
595+
widget._reset_frame_latest()
604596
with _widget_uuid_in_metadata(widget.uuid):
605597
try:
606598
if widget._run_frequency == "every-frame":
607599
widget.run_every_frame()
608600
elif batch_ready:
609-
widget.run_batch(batch_size)
601+
widget.run_batch()
610602
# pylint: disable=broad-exception-caught
611603
except Exception: # pragma: no cover
612604
logger.exception("Serial run failed for widget %s", widget.uuid)

mdadash/frontend/src/views/SettingsView.vue

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -163,12 +163,12 @@
163163

164164
<v-number-input
165165
class="mb-4"
166-
label="Batch size"
166+
label="Buffer / batch size"
167167
variant="outlined"
168168
v-model="settings.universe_configs[0].batch_size"
169169
:min="1"
170170
control-variant="hidden"
171-
hint="Number of timesteps to batch for a batch run"
171+
hint="Number of timesteps to buffer for a batch run"
172172
persistent-hint
173173
></v-number-input>
174174

0 commit comments

Comments
 (0)