Source code for mdadash.backend.widgets.base

"""
Base Class for Widgets and Widget Manager
"""

import inspect
import logging
from abc import ABC
from contextlib import contextmanager
from threading import Thread
from typing import TYPE_CHECKING, Any, ClassVar
from uuid import uuid1

import IPython
import MDAnalysis as mda
from joblib import Parallel
from matplotlib_inline.backend_inline import InlineBackend

if TYPE_CHECKING:
    from mdadash.backend.kernel.core import CommHandler, UniverseManager

logger = logging.getLogger(__name__)
InlineBackend.instance().figure_formats = {"jpeg"}


[docs] class WidgetBase(ABC): """WidgetBase This is the base class for all widgets. """ _run_frequency = "every-frame" _run_mode = "serial" def __init_subclass__(cls, **kwargs): """Register any derived class with the WidgetManager""" super().__init_subclass__(**kwargs) WidgetManager.register_class(cls) def __init__(self): self.uid = None self.u = None self.uuid = None self._wm: WidgetManager = None self._input_errors = {} def __getstate__(self): state = self.__dict__.copy() del state["_wm"] return state def __setstate__(self, state): self.__dict__.update(state) self._wm = None def _set_universe(self, u: mda.Universe): """Internal: Set the universe""" self.u = u def _reset_frame_latest(self): """Internal: Reset frame to latest timestep""" _ = self.u.trajectory[-1] def _get_inputs(self): """Internal: Get the current instance inputs""" inputs = getattr(self, "_inputs", []) if inputs is not None: for _input in inputs: # set the value and error states _input["value"] = getattr(self, _input["attribute"], None) _input["error"] = self._input_errors.get(_input["attribute"], None) return inputs def _set_input_state(self, attribute: str, error: str | None = None): """Internal: Set input attribute validation state""" if error is not None: self._input_errors[attribute] = error else: if attribute in self._input_errors: del self._input_errors[attribute] def _get_notes(self): """Internal: Get the current instance notes""" return getattr(self, "_notes", None) def _get_doclink(self): """Internal: Get the current instance doclink""" return getattr(self, "_doclink", None) def _get_tsinfo(self) -> dict: """Internal: Get the current timestep info""" return { "frame": self.u.trajectory.frame, "time": self.u.trajectory.ts.data.get("time"), "step": self.u.trajectory.ts.data.get("step"), } def _run_code(self, code: str): """Internal: Run user-defined code""" if self._wm is not None: return self._wm._run_cell(code) return None # pragma: no cover
[docs] def alert(self, message: str) -> None: """Create an alert A timestamp based on the current timestep is automatically prepended to the message. Parameters ---------- message: str The string message used for the alert """ if self._wm is not None and self._wm._comms is not None: self._wm._comms.send( {"alert": {"tsinfo": self._get_tsinfo(), "message": message}} )
[docs] def pause_simulation(self) -> None: """ Pause simulation and add an alert that this Widget (name) triggered the pause. """ if self._wm is not None and self._wm._comms is not None: self._wm._comms.send( { "pause_simulation": { "tsinfo": self._get_tsinfo(), "message": f"Pause triggered by: {getattr(self, 'name', None)}", } } )
[docs] def on_post_create(self) -> None: """ This handler is called after the widget instance is created and after all the inputs are set (Widget create, duplicate, refresh and re-create from state cases). Since the inputs are set by the time this handler is invoked, further initializations (that are not possible in ``__init__``) can be handled here. .. note:: The Universe (``self.u``) may or may not be available at this stage as that depends on the simulation connected state. Use the :meth:`on_post_connect` handler if you need the Universe to exist. """
[docs] def on_post_connect(self) -> None: """ This handler is called everytime after connecting to the simulation. The Universe (``self.u``) will be available and any ``AtomGroup`` selections or MDAnalysis analysis class instances that depend on the Universe can be created here. """
[docs] def on_post_disconnect(self) -> None: """ This handler is called everytime after disconnect from the simulation. """
[docs] def on_post_pause(self) -> None: """ This handler is called everytime after the simulation is paused. The pause could have been triggered by user clickling on the 'Pause' button on the dashboard or by any other Widget triggering a pause using :meth:`pause_simulation` or :meth:`mdadash.backend.kernel.utils.pause_simulation`. """
[docs] def on_pre_resume(self) -> None: """ This handler is called before trajectory iteration is resumed. The resume is usually triggered by user clicking on the 'Resume' button on the dashboard. """
[docs] def on_input_change(self, attribute: str, old_value: Any, new_value: Any) -> None: """ This handler is called everytime a widget input changes from the dashboard UI. Validations can be performed in this handler and any exceptions raised with messages will show up as errors in the UI. Parameters ---------- attribute: str The input attribute that changed old_value: Any The previous value held by this attribute new_value: Any The current value of this attribute .. note:: This handler is **not** invoked when a Widget inputs are set when it is duplicated or when it is recreated from the state file. Only changes from the UI trigger this handler. """
[docs] def run_every_frame(self) -> None: """ This handler is called everytime after the trajectory iterates forward if the run frequency is set to ``every-frame`` (``_run_frequency='every-frame'``). The trajectory timestep in the handler will be the current timestep. """
[docs] def run_batch(self) -> None: """ This handler is called every time after a new batch of timesteps is full and ready to be run if the run frequency is set to ``batch`` (``_run_frequency='batch'``). ``self.u.trajectory.buffer_size`` is the size of the buffer / batch (N) that can be used by the widget class to iterate the last N frames. """
[docs] def get_parallel_job(self) -> Any: """ This handler is called if the run mode is set to ``parallel`` (``_run_mode='parallel'``) to retrieve the parallel job to run from the Widget class. The Widget class must return a ``joblib``'s ``delayed`` function as the return value. Returns ------- job: Any A joblib's delayed function .. note:: As the parallel job executes in a separate process, everything needed by the Widget class must be explicitly returned by the job. See :meth:`apply_parallel_results` on how these results are available back to the Widget class. Any console outputs (like ``print``) or direct plot outputs will not be captured by the widget run. They will have to be returned and handled in :meth:`apply_parallel_results`. """
[docs] def apply_parallel_results(self, values: Any) -> None: """ This handler is called with the results of the parallel job execution. This is invoked when the run mode is set to ``parallel`` (``_run_mode='parallel'``) after the parallel job completes. Any update to the Widget class state or output plot creation etc will have to happen in this handler. Parameters ---------- values: Any The results returned by the parallel job run """
[docs] class WidgetManager: """WidgetManager This is the manager that manager all widgets. """ _instance: ClassVar = None _widget_classes: ClassVar = {} _widget_instances: ClassVar = {} def __new__(cls, *args, **kwargs): if cls._instance is None: cls._instance = super().__new__(cls) return cls._instance def __init__(self, comms: "CommHandler"): if hasattr(self, "_initialized"): return self._comms = comms self._um: UniverseManager = None self.n_jobs = 2 self._patch_IMDReader() self._initialized = True
[docs] @classmethod def register_class(cls, widget_class: WidgetBase) -> None: """Register widget class Parameters ---------- widget_class A widget class that is derived from WidgetBase """ cls._validate_widget_class(widget_class) WidgetManager._widget_classes[widget_class.name] = widget_class if WidgetManager._instance is not None: # refresh any existing instances of this class name WidgetManager._instance._refresh_instances(widget_class.name)
@classmethod def _validate_widget_class(cls, widget_class: WidgetBase) -> None: """Internal: Method to validate a widget class""" if not issubclass(widget_class, WidgetBase): raise TypeError(f"{widget_class} is not a widget class") if not hasattr(widget_class, "name"): raise ValueError("name not specified in widget class") widget_name = widget_class.name if widget_name in WidgetManager._widget_classes: if hasattr(widget_class, "_override_name") and widget_class._override_name: logger.warning("Overriding widget class for '%s'", widget_name) else: raise ValueError( f"Widget name '{widget_name}' already registered. " f"Use `_override_name` attribute set to `True` to force registration" ) # check for one of the run methods to exist with correct params run_methods = { "run_every_frame": 1, "run_batch": 1, } has_valid_run_method = False for run_method, n_params in run_methods.items(): method = getattr(widget_class, run_method) if method == getattr(WidgetBase, run_method): continue if not callable(method): continue signature = inspect.signature(method) has_valid_run_method = len(signature.parameters.values()) == n_params break if not has_valid_run_method: raise ValueError("run method not found in class") # TODO: add more validations def _invoke_widget_lifecyle_method(self, widget: WidgetBase, method: str) -> None: """Internal: Invoke the lifecycle method if implemented""" if widget._input_errors: # lifecycle methods not invoked when # there are input errors return if hasattr(widget, method): handler = getattr(widget, method) if callable(handler): try: handler() # pylint: disable=broad-exception-caught except Exception: # pragma: no cover logger.exception( "Failed to invoke lifecycle method %s for widget %s", method, widget.uuid, ) def _set_widget_universe( self, widget: WidgetBase, uid: int, u: mda.Universe ) -> None: """Internal: Set the universe for instance""" if widget.uid == uid: widget._set_universe(u) # invoke the on_post_connect handler self._invoke_widget_lifecyle_method(widget, "on_post_connect") def _set_universe(self, uid: int, u: mda.Universe, uuid: str | None = None) -> None: """Internal: Set the universe for all or given widget""" if uuid is None: for widget in WidgetManager._widget_instances.values(): self._set_widget_universe(widget, uid, u) else: widget = WidgetManager._widget_instances[uuid] self._set_widget_universe(widget, uid, u) def _invoke_lifecycle_method(self, method: str) -> None: """Internal: Invoke given lifecycle method for all instances""" for widget in WidgetManager._widget_instances.values(): self._invoke_widget_lifecyle_method(widget, method) def _get_inputs_state(self, inputs): """Internal: Get all the input values and any errors""" return [ {k: i[k] for k in ("attribute", "value", "error") if k in i} for i in inputs ]
[docs] def get_available_widgets(self, _data: dict) -> None: """Get available widgets Sends a dict containing name and description of all available widgets to the client. """ widgets = [ { "name": c.name, "description": getattr(c, "description", None), } for c in sorted( WidgetManager._widget_classes.values(), key=lambda c: c.name.lower() ) ] self._comms.send({"widgets": widgets})
[docs] def recreate_instances(self, data: dict) -> None: """Recreate widget instances Recreate widget instances with data from state file Parameters ---------- data: dict Data of the instances that need to be recreated """ ret = self._recreate_instances(data) self._comms.send({"status": "ok" if ret else "error"})
[docs] def add_widget_instance(self, data: dict) -> dict: """Add widget instance based on registered widget name""" uid = data["uid"] widget_name = data["name"] uuid, details = self._add_widget_instance(uid, widget_name) if uuid is not None: self._comms.send( { "status": "ok", "uuid": uuid, "details": details, } ) else: self._comms.send( { "status": "error", "message": f"Failed to add widget instance for {widget_name}", } )
[docs] def duplicate_widget_instance(self, data: dict) -> None: """Duplicate widget instance based on instance uuid""" uid = data["uid"] new_uuid, details = self._duplicate_widget_instance(uid, data["uuid"]) self._comms.send( { "status": "ok", "uuid": new_uuid, "details": details, } )
[docs] def remove_widget_instance(self, data: dict) -> None: """Remove widget instance Remove widget instance based on uuid returned during the instance creation using :meth:`add_widget_instance` Parameters ---------- data: dict Dict that has the following keys: uuid: str The uuid of the instance """ uuid = self._remove_widget_instance(data["uuid"]) if uuid is not None: self._comms.send({"status": "ok"}) else: self._comms.send( { "status": "error", "message": f"Failed to remove widget instance with uuid {uuid}", } )
[docs] def get_widget_inputs(self, data: dict) -> None: """Get inputs Send a dict containing the inputs and notes for a given widget uuid. Parameters ---------- data: dict Dict that has the following keys: uuid: str The uuid of the instance """ uuid = data["uuid"] self._comms.send( { "status": "ok", "inputs": self._get_widget_inputs(uuid), "notes": self._get_widget_notes(uuid), "doclink": self._get_widget_doclink(uuid), } )
[docs] def set_widget_input(self, data: dict) -> None: """Set input Parameters ---------- data: dict Dict that has the following keys: uuid: str The uuid of the instance attribute: str The input attribute to set value: Any The value to set for the attribute """ ret = self._set_widget_input(data["uuid"], data["attribute"], data["value"]) self._comms.send({"status": "ok" if ret else "error"})
def _add_widget_instance( self, uid: int, widget_name: str ) -> tuple[str, dict] | None: """Add widget instance Add a widget instance based on the widget name already registered with the manager. Parameters ---------- uid: int Universe ID (index into universes array) widget_name: str Name of the widget class registered Returns ------- uuid of instance added and input details or None, None """ if widget_name in WidgetManager._widget_classes: widget_class = WidgetManager._widget_classes[widget_name] uuid = str(uuid1()) instance = widget_class() instance.uid = uid instance.uuid = uuid instance._wm = self WidgetManager._widget_instances[uuid] = instance details = { "uid": uid, "class_name": widget_name, "inputs": self._get_inputs_state(instance._get_inputs()), } # invoke the on_post_create handler self._invoke_widget_lifecyle_method(instance, "on_post_create") # set the universe for the new widget instance if self._um._connected: self._set_universe(uid, self._um._universes[uid], uuid) return uuid, details return None, None def _duplicate_widget_instance(self, uid: int, uuid: str) -> tuple[str, dict]: """Duplicate widget instance Duplicate widget instance based on existing instance uuid Parameters ---------- uid: int Universe ID (index into universes array) uuid: str The uuid of the instance to be duplicated Returns ------- uuid of new instance created and input details """ # get existing instance instance = WidgetManager._widget_instances[uuid] # duplicate instance widget_class = instance.__class__ new_instance = widget_class() new_instance.uid = uid new_instance._wm = self # set inputs for new instance inputs = instance._get_inputs() for _input in inputs: attribute = _input["attribute"] setattr(new_instance, attribute, _input["value"]) if _input["error"] is not None: new_instance._set_input_state(attribute, _input["error"]) # add new instance to instances list new_uuid = str(uuid1()) new_instance.uuid = new_uuid WidgetManager._widget_instances[new_uuid] = new_instance details = { "uid": uid, "class_name": widget_class.name, "inputs": self._get_inputs_state(inputs), } # invoke the on_post_create handler self._invoke_widget_lifecyle_method(new_instance, "on_post_create") # set the universe for the new widget instance if self._um._connected: self._set_universe(uid, self._um._universes[uid], new_uuid) return new_uuid, details def _recreate_instances(self, data: dict) -> None: """Internal: Recreate widget instances""" ret = True for widget_uuid, widget in data.items(): try: widget_class = WidgetManager._widget_classes[widget["class_name"]] instance = widget_class() instance.uid = widget["uid"] instance.uuid = widget_uuid instance._wm = self inputs = widget["inputs"] for _input in inputs: attribute = _input["attribute"] setattr(instance, attribute, _input["value"]) if _input["error"] is not None: instance._set_input_state(attribute, _input["error"]) WidgetManager._widget_instances[widget_uuid] = instance # invoke the on_post_create handler self._invoke_widget_lifecyle_method(instance, "on_post_create") except KeyError: logger.exception("Key error while recreating widget instances") ret = False return ret def _remove_widget_instance(self, uuid: str) -> str | None: """Internal: Remove a widget instance""" if uuid in WidgetManager._widget_instances: del WidgetManager._widget_instances[uuid] return uuid return None def _refresh_instances(self, class_name: str) -> None: """Internal: Recreate widget instances when class is updated""" widget_class = WidgetManager._widget_classes[class_name] for instance in WidgetManager._widget_instances.values(): if instance.name != class_name: continue # create new instance new_instance = widget_class() uid = instance.uid new_instance.uid = uid new_instance._wm = self # set inputs for new instance inputs = instance._get_inputs() for _input in inputs: attribute = _input["attribute"] setattr(new_instance, attribute, _input["value"]) if _input["error"] is not None: # pragma: no cover new_instance._set_input_state(attribute, _input["error"]) # update new instance in instances list uuid = instance.uuid new_instance.uuid = uuid WidgetManager._widget_instances[uuid] = new_instance # invoke the on_post_create handler self._invoke_widget_lifecyle_method(new_instance, "on_post_create") # set the universe for the new widget instance if self._um._connected: self._set_universe(uid, self._um._universes[uid], uuid) def _get_widget_inputs(self, uuid: str) -> list: """Internal: Get a widget inputs""" widget = WidgetManager._widget_instances[uuid] return widget._get_inputs() def _get_widget_notes(self, uuid: str) -> str: """Internal: Get notes for widget instance""" widget = WidgetManager._widget_instances[uuid] return widget._get_notes() def _get_widget_doclink(self, uuid: str) -> str: """Internal: Get doclink for widget instance""" widget = WidgetManager._widget_instances[uuid] return widget._get_doclink() def _set_widget_input(self, uuid: str, attribute: str, value: Any) -> bool: """Internal: Set a widget input""" widget = WidgetManager._widget_instances[uuid] old_value = getattr(widget, attribute, value) old_type = type(old_value) # set input using the same existing type setattr(widget, attribute, value if old_value is None else old_type(value)) try: widget.on_input_change(attribute, old_value, value) widget._set_input_state(attribute) return True except Exception as e: # pylint: disable=broad-exception-caught # noqa: BLE001 widget._set_input_state(attribute, str(e)) return False
[docs] def update_n_jobs(self, data: dict) -> None: """Update n_jobs for ``joblib.Parallel`` Parameters ---------- data: dict Dict that has the following keys: n_jobs: int The number of parallel jobs """ self.n_jobs = data["n_jobs"]
@staticmethod def _patch_IMDReader(): """Internal: Patch `IMDReader` to make it serializable""" # pylint: disable=import-outside-toplevel from MDAnalysis.coordinates.IMD import IMDReader def custom_getstate(self): state = self.__dict__.copy() del state["_imdclient"] return state def custom_setstate(self, state): self.__dict__.update(state) self._imdclient = None IMDReader.__setstate__ = custom_setstate IMDReader.__getstate__ = custom_getstate @staticmethod def _with_reset_frame(func, *args, **kwargs): """Internal: Reset frame to the most recent one""" instance = func.__self__ instance._reset_frame_latest() return func(*args, **kwargs) def _run_parallel_jobs(self, parallel_widgets, parallel_results): """Internal: Run parallel jobs using joblib.Parallel""" parallel_jobs = [] for widget in parallel_widgets: 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, max_nbytes=None, initializer=WidgetManager._patch_IMDReader, )(parallel_jobs) parallel_results.extend(results) # pylint: disable=broad-exception-caught except Exception: # pragma: no cover logger.exception("Parallel run failed for jobs %s", parallel_jobs) # pylint: disable=too-many-branches
[docs] def run_widgets(self, uid: int, batch_ready: bool) -> None: """Run widget instances Parameters ---------- uid: int Universe ID (index into universes array) batch_ready: bool Flag indicating if a batch of timesteps is full """ # collect widgets that need to be run parallel_widgets = [] serial_widgets = [] for widget in WidgetManager._widget_instances.values(): # only run widget if there are no input errors if widget.uid != uid or widget._input_errors: continue if widget._run_mode == "parallel": if widget._run_frequency == "every-frame" or batch_ready: parallel_widgets.append(widget) else: serial_widgets.append(widget) # run parallel widgets in separate thread if parallel_widgets: parallel_results = [] parallel_thread = Thread( target=self._run_parallel_jobs, args=( parallel_widgets, parallel_results, ), ) parallel_thread.start() # run serial widgets for widget in serial_widgets: widget._reset_frame_latest() widget_outputs = None with _capture_outputs() as captured_outputs: try: if widget._run_frequency == "every-frame": widget_outputs = widget.run_every_frame() elif batch_ready: widget_outputs = widget.run_batch() # pylint: disable=broad-exception-caught except Exception: # pragma: no cover logger.exception("Serial run failed for widget %s", widget.uuid) if widget_outputs is not None: # custom code widget returns outputs directly self._comms.send( {"widget_outputs": {"uuid": widget.uuid, "outputs": widget_outputs}} ) elif captured_outputs: self._comms.send( { "widget_outputs": { "uuid": widget.uuid, "outputs": captured_outputs, } } ) # apply parallel results back if parallel_widgets: # wait for all parallel jobs to be done parallel_thread.join() for i, widget in enumerate(parallel_widgets): with _capture_outputs() as captured_outputs: widget.apply_parallel_results(parallel_results[i]) if captured_outputs: self._comms.send( { "widget_outputs": { "uuid": widget.uuid, "outputs": captured_outputs, } } )
def _run_cell(self, code: str) -> list: """Internal: Run code and return all outputs""" outputs = [] with _capture_outputs() as capture: result = IPython.get_ipython().run_cell(code) if result.error_before_exec: outputs.append({"type": "error", "content": str(result.error_before_exec)}) if result.error_in_exec: outputs.append({"type": "error", "content": str(result.error_in_exec)}) if result.result is not None: outputs.append({"type": "text", "content": str(result.result)}) outputs.extend(capture) return outputs
[docs] def execute_code(self, data: dict) -> None: """Execute code in the kernel Parameters ---------- data: dict Dict that has the following keys: code: str The code to execute in the kernel """ outputs = self._run_cell(data["code"]) self._comms.send({"outputs": outputs})
@contextmanager def _capture_outputs(): """Internal: Context manager to capture outputs of code execution""" outputs = [] with IPython.utils.capture.capture_output() as capture: yield outputs if capture.stdout: outputs.append({"type": "text", "content": capture.stdout}) if capture.stderr: outputs.append({"type": "text", "content": capture.stderr}) for out in capture.outputs: data = out.data if "image/jpeg" in data: outputs.append({"type": "image", "content": data["image/jpeg"]}) elif "text/plain" in data: # pragma: no cover # skip text repr of a matplotlib image if image exists outputs.append({"type": "text", "content": data["text/plain"]})