"""
Helix Analysis
"""
import logging
import warnings
from collections import deque
from typing import ClassVar
import matplotlib.pyplot as plt
from IPython.display import display
from joblib import delayed
from MDAnalysis.analysis.helix_analysis import HELANAL
from mdadash.backend.widgets.base import WidgetBase
logger = logging.getLogger(__name__)
[docs]
class HelixAnalysis(WidgetBase):
"""
**Helix Analysis**
This widget uses `MDAnalysis.analysis.helix_analysis.HELANAL`_ to perform helix
analysis and plot different computed properties as per the MDAnalysis
`Helix analysis`_ examples.
.. _MDAnalysis.analysis.helix_analysis.HELANAL: https://docs.mdanalysis.org/
stable/documentation_pages/analysis/helix_analysis.html
#MDAnalysis.analysis.helix_analysis.HELANAL
.. _Helix analysis: https://userguide.mdanalysis.org/stable/examples/
analysis/structure/helanal.html
"""
name = "Helix Analysis"
description = "Helix analysis using HELANAL"
_inputs: ClassVar = [
{
"attribute": "_run_frequency",
"name": "Run frequency",
"description": "The frequency with which the widget is run",
"type": "select",
"items": [
"every-frame",
"batch",
],
},
{
"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",
"validations": ["required"],
},
{
"attribute": "property",
"name": "Property",
"description": "Computed property to plot",
"type": "select",
"items": [
"local_twists",
"local_nres_per_turn",
"local_bends",
"local_heights",
"local_screw_angles",
],
},
{
"attribute": "custom_title",
"name": "Custom title",
"description": "Custom title for the plot",
"type": "str",
},
{
"attribute": "maxlen",
"name": "Max values",
"description": "Max values to show in plot",
"type": "int",
},
{
"attribute": "x_type",
"name": "X-axis",
"type": "toggle",
"options": [
{"name": "Time", "value": "time"},
{"name": "Step", "value": "step"},
],
},
]
def __init__(self):
super().__init__()
self.selection = "resid 1:10"
self.property = "local_twists"
self.ha = None
self.title = "Helix Analysis"
self.custom_title = None
self.default_maxlen = 100
self.maxlen = self.default_maxlen
self.x_type = "time"
self.x_values = None
self.y_labels = {
"local_twists": "Average local twist (degrees)",
"local_nres_per_turn": "Average residues per turn",
"local_bends": "Average local bends (degrees)",
"local_heights": "Average rise of each local helix (Å)",
"local_screw_angles": "Average local screw angle (degrees)",
}
self._setup_plot()
self._reset_plot_values()
def _setup_plot(self):
"""Setup matplotlib plot"""
self.fig, self.ax = plt.subplots()
(self.plot,) = self.ax.plot([], [])
self.ax.grid(True)
self._set_title()
def _reset_plot_values(self):
"""Reset plot values"""
self.steps = deque(maxlen=self.maxlen)
self.times = deque(maxlen=self.maxlen)
self.y_values = deque(maxlen=self.maxlen)
self.ax.set_ylabel(self.y_labels[self.property])
self._set_x_values()
def _set_title(self):
"""Set plot title"""
self.ax.set_title(
self.custom_title.replace("\\n", "\n") if self.custom_title else self.title
)
def _set_x_values(self):
"""Set the values for the x-axis"""
if self.x_type == "step":
x_label = "Step"
self.x_values = self.steps
else:
x_label = "Time (ps)"
self.x_values = self.times
self.ax.set_xlabel(x_label)
def _create_ha(self):
"""Update atom groups when selection phrases change"""
self.ha = HELANAL(self.u, select=self.selection)
self.title = f"Helix analysis of '{self.selection}'"
self._set_title()
[docs]
def on_post_create(self):
"""on_post_create handler"""
self._set_title()
self._reset_plot_values()
[docs]
def on_post_connect(self):
"""on_post_connect handler"""
self._create_ha()
def _compute_current_frame(self):
"""Compute values for current frame"""
with warnings.catch_warnings():
warnings.simplefilter("ignore", category=RuntimeWarning)
self.ha.run(frames=[self.u.trajectory.frame])
results = getattr(self.ha.results, self.property)
mean_values = results.mean(axis=1)
return (
self.u.trajectory.ts.data["step"],
self.u.trajectory.ts.data["time"],
mean_values[0],
)
def _compute_batch(self):
"""Compute values for current batch"""
with warnings.catch_warnings():
warnings.simplefilter("ignore", category=RuntimeWarning)
self.ha.run()
results = getattr(self.ha.results, self.property)
mean_values = results.mean(axis=1)
values = []
for i, v in enumerate(mean_values):
_ = self.u.trajectory[i]
values.append(
(
self.u.trajectory.ts.data["step"],
self.u.trajectory.ts.data["time"],
v,
)
)
return values
def _update_plot(self, values):
"""Append values and update plot"""
if isinstance(values, tuple):
values = [values]
# update plot points
for value in values:
(steps, times, v) = value
self.steps.append(steps)
self.times.append(times)
self.y_values.append(v)
# update plot
self.plot.set_data(self.x_values, self.y_values)
self.ax.relim()
self.ax.autoscale_view()
self.fig.canvas.draw()
display(self.fig)
[docs]
def run_every_frame(self):
"""every-frame run handler"""
self._update_plot(self._compute_current_frame())
[docs]
def run_batch(self):
"""batch run handler"""
self._update_plot(self._compute_batch())
[docs]
def get_parallel_job(self):
"""get parallel job handler"""
if self._run_frequency == "batch":
return delayed(self._compute_batch)()
return delayed(self._compute_current_frame)()
[docs]
def apply_parallel_results(self, values):
"""apply parallel results handler"""
self._update_plot(values)