diff --git a/waveform_editor/gui/shape_editor/nice_plotter.py b/waveform_editor/gui/shape_editor/nice_plotter.py index 387f99d1..897070f4 100644 --- a/waveform_editor/gui/shape_editor/nice_plotter.py +++ b/waveform_editor/gui/shape_editor/nice_plotter.py @@ -6,6 +6,7 @@ import numpy as np import panel as pn import param +import scipy.interpolate as interp from bokeh.models import HoverTool, PointDrawTool from imas.ids_toplevel import IDSToplevel from panel.viewable import Viewer @@ -45,6 +46,10 @@ class NicePlotter(Viewer): levels = param.Integer( default=20, bounds=(1, 200), label="Number of contour levels" ) + show_heatmap = param.Boolean(default=False, label="Show heatmap") + heatmap_alpha = param.Number( + default=0.7, bounds=(0.0, 1.0), step=0.05, label="Heatmap opacity" + ) show_coils = param.Boolean(default=True, label="Show coils") show_wall = param.Boolean(default=True, label="Show limiter and divertor") show_vacuum_vessel = param.Boolean( @@ -66,6 +71,7 @@ class NicePlotter(Viewer): FRAME_WIDTH = round( FRAME_HEIGHT * (R_RANGE[1] - R_RANGE[0]) / (Z_RANGE[1] - Z_RANGE[0]) ) + HEATMAP_RESOLUTION = 250 def __init__(self, **params): super().__init__(**params) @@ -87,6 +93,27 @@ def __init__(self, **params): colorbar_opts={"title": "Poloidal flux [Wb]"}, show_legend=False, ) + self.CONTOUR_ON_HEATMAP_OPTS = hv.opts.Contours( + color="white", + alpha=0.6, + line_width=1, + colorbar=False, + tools=["hover"], + show_legend=False, + ) + self._heatmap_cache = (None, None) + self.HEATMAP_OPTS = hv.opts.Image( + cmap="viridis", + colorbar=True, + tools=["hover"], + hover_tooltips=[ + ("r", "$x{0.00} m"), + ("z", "$y{0.00} m"), + ("psi", "@image{0.000} Wb"), + ], + colorbar_opts={"title": "Poloidal flux [Wb]"}, + show_legend=False, + ) self.DESIRED_SHAPE_OPTS = hv.opts.Curve(color="blue") # Static on purpose. The draw tool binds to this element's renderer, and # redrawing it either breaks the binding or feeds edits back into redraws. @@ -98,6 +125,7 @@ def __init__(self, **params): hooks=[self._capture_points_renderer], ) flux_map_elements = [ + hv.DynamicMap(self._plot_heatmap), hv.DynamicMap(self._plot_contours), hv.DynamicMap(self._plot_separatrix), hv.DynamicMap(self._plot_xo_points), @@ -355,7 +383,56 @@ def _plot_coil_rectangles(self): ) return rects * paths - @pn.depends("communicator.equilibrium", "show_contour", "levels") + @pn.depends("communicator.equilibrium", "show_heatmap", "heatmap_alpha") + def _plot_heatmap(self): + """Generates heatmap plot for poloidal flux. + + Returns: + Holoviews Image containing the poloidal flux field. + """ + equilibrium = self.communicator.equilibrium + if not self.show_heatmap or equilibrium is None: + return self._empty_heatmap() + + flux_map = self._flux_map(equilibrium) + if flux_map is None: + return self._empty_heatmap() + + if self._heatmap_cache[0] is not equilibrium: + self._heatmap_cache = (equilibrium, self._calc_heatmap(*flux_map)) + return ( + hv.Image(self._heatmap_cache[1], kdims=["r", "z"], vdims=["psi"]) + .opts(self.HEATMAP_OPTS) + .opts(alpha=self.heatmap_alpha) + ) + + def _empty_heatmap(self): + return ( + hv.Image([], kdims=["r", "z"], vdims=["psi"]) + .opts(self.HEATMAP_OPTS) + .opts(alpha=0.0, colorbar=False) + ) + + def _calc_heatmap(self, r, z, psi): + """Interpolates psi onto a regular grid for heatmap display. + + Args: + r: Radial coordinates of the mesh nodes. + z: Height coordinates of the mesh nodes. + psi: Poloidal flux values at the mesh nodes. + + Returns: + Tuple of (grid_r, grid_z, psi_grid). + """ + r, z = np.asarray(r, dtype=float), np.asarray(z, dtype=float) + grid_r = np.linspace(r.min(), r.max(), self.HEATMAP_RESOLUTION) + grid_z = np.linspace(z.min(), z.max(), self.HEATMAP_RESOLUTION) + psi_grid = interp.griddata( + (r, z), np.asarray(psi, dtype=float), tuple(np.meshgrid(grid_r, grid_z)) + ) + return grid_r, grid_z, psi_grid + + @pn.depends("communicator.equilibrium", "show_contour", "levels", "show_heatmap") def _plot_contours(self): """Generates contour plot for poloidal flux. @@ -368,7 +445,27 @@ def _plot_contours(self): else: contours = self._calc_contours(equilibrium, self.levels) - return contours.opts(self.CONTOUR_OPTS) + return contours.opts( + self.CONTOUR_ON_HEATMAP_OPTS if self.show_heatmap else self.CONTOUR_OPTS + ) + + def _flux_map(self, equilibrium): + """The poloidal flux on the GGD that NICE fills. + + Args: + equilibrium: The equilibrium IDS to read the flux from. + + Returns: + Tuple of (r, z, psi) at the mesh nodes, or None if NICE did not fill them. + """ + eqggd = equilibrium.time_slice[0].ggd[0] + r, z, psi = eqggd.r[0].values, eqggd.z[0].values, eqggd.psi[0].values + if not r or not z or not psi: + pn.state.notifications.error( + "NICE did not produce a valid poloidal flux field" + ) + return None + return r, z, psi def _calc_contours(self, equilibrium, levels): """Calculates the contours of the psi grid of an equilibrium IDS. @@ -381,19 +478,11 @@ def _calc_contours(self, equilibrium, levels): Returns: Holoviews contours object """ - - eqggd = equilibrium.time_slice[0].ggd[0] - r = eqggd.r[0].values - z = eqggd.z[0].values - psi = eqggd.psi[0].values - - if not r or not z or not psi: - pn.state.notifications.error( - "NICE did not produce a valid poloidal flux field" - ) + flux_map = self._flux_map(equilibrium) + if flux_map is None: return hv.Contours(([0], [0], 0), vdims="psi") - trics = plt.tricontour(r, z, psi, levels=levels) + trics = plt.tricontour(*flux_map, levels=levels) return hv.Contours(self._extract_contour_segments(trics), vdims="psi") def _extract_contour_segments(self, tricontour): diff --git a/waveform_editor/gui/shape_editor/settings_modal.py b/waveform_editor/gui/shape_editor/settings_modal.py index f037ac01..82779e68 100644 --- a/waveform_editor/gui/shape_editor/settings_modal.py +++ b/waveform_editor/gui/shape_editor/settings_modal.py @@ -43,6 +43,25 @@ def _settings_section(self, *items): sizing_mode="stretch_width", ) + def _detail_section(self, label, parameters, toggle): + """A section of plot parameters, shown while the toggle that enables them is. + + Args: + label: The title of the section. + parameters: Names of the parameters of the plotter in the section. + toggle: The boolean parameter that shows the section. + """ + return pn.Column( + self._section_label(label), + self._settings_section( + pn.Param( + self.nice_plotter.param, parameters=parameters, show_name=False + ) + ), + visible=toggle, + sizing_mode="stretch_width", + ) + def _form_row(self, label, widget, warning=None): items = [ pn.pane.HTML( @@ -185,21 +204,11 @@ def _build_modal(self): ) # --- Display tab --- - self._contour_detail = pn.Column( - self._section_label("Contour Detail"), - self._settings_section( - pn.Param( - self.nice_plotter.param, - parameters=["levels"], - show_name=False, - ), - ), - visible=self.nice_plotter.show_contour, - sizing_mode="stretch_width", + self._contour_detail = self._detail_section( + "Contour Detail", ["levels"], self.nice_plotter.param.show_contour ) - self.nice_plotter.param.watch( - lambda e: setattr(self._contour_detail, "visible", e.new), - ["show_contour"], + self._heatmap_detail = self._detail_section( + "Heatmap Detail", ["heatmap_alpha"], self.nice_plotter.param.show_heatmap ) display_content = pn.Column( @@ -209,6 +218,7 @@ def _build_modal(self): self.nice_plotter.param, parameters=[ "show_contour", + "show_heatmap", "show_coils", "show_wall", "show_vacuum_vessel", @@ -225,6 +235,7 @@ def _build_modal(self): ), ), self._contour_detail, + self._heatmap_detail, sizing_mode="stretch_width", scroll=True, )