From 91eaa35eeb14f907a7a8cd9785325a22fdbb234f Mon Sep 17 00:00:00 2001 From: Sebbe Blokhuizen Date: Fri, 2 Oct 2026 11:45:48 +0200 Subject: [PATCH 1/3] add option for heatmap --- .../gui/shape_editor/nice_plotter.py | 146 +++++++++++++++++- .../gui/shape_editor/settings_modal.py | 19 +++ 2 files changed, 164 insertions(+), 1 deletion(-) diff --git a/waveform_editor/gui/shape_editor/nice_plotter.py b/waveform_editor/gui/shape_editor/nice_plotter.py index 387f99d1..b22abc99 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( @@ -87,6 +92,20 @@ def __init__(self, **params): colorbar_opts={"title": "Poloidal flux [Wb]"}, show_legend=False, ) + 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._cached_equilibrium = None + self._cached_heatmap_grid = None 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 +117,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 +375,122 @@ 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 ( + hv.Image([], kdims=["r", "z"], vdims=["psi"]) + .opts(self.HEATMAP_OPTS) + .opts(alpha=0.0, colorbar=False) + ) + + 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" + ) + return ( + hv.Image([], kdims=["r", "z"], vdims=["psi"]) + .opts(self.HEATMAP_OPTS) + .opts(alpha=0.0, colorbar=False) + ) + + if ( + self._cached_heatmap_grid is not None + and self._cached_equilibrium is equilibrium + ): + grid_r, grid_z, psi_grid = self._cached_heatmap_grid + else: + grid_r, grid_z, psi_grid = self._calc_heatmap(r, z, psi) + self._cached_equilibrium = equilibrium + self._cached_heatmap_grid = (grid_r, grid_z, psi_grid) + + if grid_r is None: + return ( + hv.Image([], kdims=["r", "z"], vdims=["psi"]) + .opts(self.HEATMAP_OPTS) + .opts(alpha=0.0, colorbar=False) + ) + + return ( + hv.Image( + (grid_r, grid_z, psi_grid), + kdims=["r", "z"], + vdims=["psi"], + ) + .opts(self.HEATMAP_OPTS) + .opts(alpha=self.heatmap_alpha) + ) + + def _calc_heatmap(self, r, z, psi, resolution=250): + """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. + resolution: Number of grid points along each axis. + + Returns: + Tuple of (grid_r, grid_z, psi_grid) or (None, None, None). + """ + try: + r_arr = np.asarray(r, dtype=float) + z_arr = np.asarray(z, dtype=float) + psi_arr = np.asarray(psi, dtype=float) + + # Filter non-finite values + valid = np.isfinite(r_arr) & np.isfinite(z_arr) & np.isfinite(psi_arr) + r_arr, z_arr, psi_arr = r_arr[valid], z_arr[valid], psi_arr[valid] + + if len(r_arr) < 3 or len(z_arr) < 3 or len(psi_arr) < 3: + return None, None, None + + # Deduplicate coordinates so identical points don't cause degeneracies + coords = np.column_stack([r_arr, z_arr]) + _, unique_indices = np.unique(coords, axis=0, return_index=True) + r_clean = r_arr[unique_indices] + z_clean = z_arr[unique_indices] + psi_clean = psi_arr[unique_indices] + + if len(np.unique(r_clean)) < 2 or len(np.unique(z_clean)) < 2: + return None, None, None + + grid_r = np.linspace(r_clean.min(), r_clean.max(), resolution) + grid_z = np.linspace(z_clean.min(), z_clean.max(), resolution) + grid_r_mesh, grid_z_mesh = np.meshgrid(grid_r, grid_z) + + try: + psi_grid = interp.griddata( + (r_clean, z_clean), + psi_clean, + (grid_r_mesh, grid_z_mesh), + method="linear", + ) + except Exception: + psi_grid = interp.griddata( + (r_clean, z_clean), + psi_clean, + (grid_r_mesh, grid_z_mesh), + method="nearest", + ) + + return grid_r, grid_z, psi_grid + except Exception as e: + logger.warning(f"Failed to calculate heatmap interpolation: {e}") + return None, None, None + + @pn.depends("communicator.equilibrium", "show_contour", "levels", "show_heatmap") def _plot_contours(self): """Generates contour plot for poloidal flux. @@ -368,6 +503,15 @@ def _plot_contours(self): else: contours = self._calc_contours(equilibrium, self.levels) + if self.show_heatmap: + return contours.opts( + color="white", + alpha=0.6, + line_width=1, + colorbar=False, + tools=["hover"], + show_legend=False, + ) return contours.opts(self.CONTOUR_OPTS) def _calc_contours(self, equilibrium, levels): diff --git a/waveform_editor/gui/shape_editor/settings_modal.py b/waveform_editor/gui/shape_editor/settings_modal.py index f037ac01..5eae5f5c 100644 --- a/waveform_editor/gui/shape_editor/settings_modal.py +++ b/waveform_editor/gui/shape_editor/settings_modal.py @@ -202,6 +202,23 @@ def _build_modal(self): ["show_contour"], ) + self._heatmap_detail = pn.Column( + self._section_label("Heatmap Detail"), + self._settings_section( + pn.Param( + self.nice_plotter.param, + parameters=["heatmap_alpha"], + show_name=False, + ), + ), + visible=self.nice_plotter.show_heatmap, + sizing_mode="stretch_width", + ) + self.nice_plotter.param.watch( + lambda e: setattr(self._heatmap_detail, "visible", e.new), + ["show_heatmap"], + ) + display_content = pn.Column( self._section_label("Visibility"), self._settings_section( @@ -209,6 +226,7 @@ def _build_modal(self): self.nice_plotter.param, parameters=[ "show_contour", + "show_heatmap", "show_coils", "show_wall", "show_vacuum_vessel", @@ -225,6 +243,7 @@ def _build_modal(self): ), ), self._contour_detail, + self._heatmap_detail, sizing_mode="stretch_width", scroll=True, ) From a9f135ba8d95458029e5411e22fe905f7b6ffb6c Mon Sep 17 00:00:00 2001 From: Sebbe Blokhuizen Date: Tue, 6 Oct 2026 10:02:03 +0200 Subject: [PATCH 2/3] deduplication and clean up of heatmap calculations --- .../gui/shape_editor/nice_plotter.py | 127 +++++------------- .../gui/shape_editor/settings_modal.py | 54 ++++---- 2 files changed, 57 insertions(+), 124 deletions(-) diff --git a/waveform_editor/gui/shape_editor/nice_plotter.py b/waveform_editor/gui/shape_editor/nice_plotter.py index b22abc99..a0b09934 100644 --- a/waveform_editor/gui/shape_editor/nice_plotter.py +++ b/waveform_editor/gui/shape_editor/nice_plotter.py @@ -71,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) @@ -92,6 +93,15 @@ 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, @@ -104,8 +114,6 @@ def __init__(self, **params): colorbar_opts={"title": "Poloidal flux [Wb]"}, show_legend=False, ) - self._cached_equilibrium = None - self._cached_heatmap_grid = None 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. @@ -384,111 +392,51 @@ def _plot_heatmap(self): """ equilibrium = self.communicator.equilibrium if not self.show_heatmap or equilibrium is None: - return ( - hv.Image([], kdims=["r", "z"], vdims=["psi"]) - .opts(self.HEATMAP_OPTS) - .opts(alpha=0.0, colorbar=False) - ) + return self._empty_heatmap() 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" ) - return ( - hv.Image([], kdims=["r", "z"], vdims=["psi"]) - .opts(self.HEATMAP_OPTS) - .opts(alpha=0.0, colorbar=False) - ) - - if ( - self._cached_heatmap_grid is not None - and self._cached_equilibrium is equilibrium - ): - grid_r, grid_z, psi_grid = self._cached_heatmap_grid - else: - grid_r, grid_z, psi_grid = self._calc_heatmap(r, z, psi) - self._cached_equilibrium = equilibrium - self._cached_heatmap_grid = (grid_r, grid_z, psi_grid) - - if grid_r is None: - return ( - hv.Image([], kdims=["r", "z"], vdims=["psi"]) - .opts(self.HEATMAP_OPTS) - .opts(alpha=0.0, colorbar=False) - ) + return self._empty_heatmap() + if self._heatmap_cache[0] is not equilibrium: + self._heatmap_cache = (equilibrium, self._calc_heatmap(r, z, psi)) return ( - hv.Image( - (grid_r, grid_z, psi_grid), - kdims=["r", "z"], - vdims=["psi"], - ) + hv.Image(self._heatmap_cache[1], kdims=["r", "z"], vdims=["psi"]) .opts(self.HEATMAP_OPTS) .opts(alpha=self.heatmap_alpha) ) - def _calc_heatmap(self, r, z, psi, resolution=250): + 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. - resolution: Number of grid points along each axis. Returns: - Tuple of (grid_r, grid_z, psi_grid) or (None, None, None). + Tuple of (grid_r, grid_z, psi_grid). """ - try: - r_arr = np.asarray(r, dtype=float) - z_arr = np.asarray(z, dtype=float) - psi_arr = np.asarray(psi, dtype=float) - - # Filter non-finite values - valid = np.isfinite(r_arr) & np.isfinite(z_arr) & np.isfinite(psi_arr) - r_arr, z_arr, psi_arr = r_arr[valid], z_arr[valid], psi_arr[valid] - - if len(r_arr) < 3 or len(z_arr) < 3 or len(psi_arr) < 3: - return None, None, None - - # Deduplicate coordinates so identical points don't cause degeneracies - coords = np.column_stack([r_arr, z_arr]) - _, unique_indices = np.unique(coords, axis=0, return_index=True) - r_clean = r_arr[unique_indices] - z_clean = z_arr[unique_indices] - psi_clean = psi_arr[unique_indices] - - if len(np.unique(r_clean)) < 2 or len(np.unique(z_clean)) < 2: - return None, None, None - - grid_r = np.linspace(r_clean.min(), r_clean.max(), resolution) - grid_z = np.linspace(z_clean.min(), z_clean.max(), resolution) - grid_r_mesh, grid_z_mesh = np.meshgrid(grid_r, grid_z) - - try: - psi_grid = interp.griddata( - (r_clean, z_clean), - psi_clean, - (grid_r_mesh, grid_z_mesh), - method="linear", - ) - except Exception: - psi_grid = interp.griddata( - (r_clean, z_clean), - psi_clean, - (grid_r_mesh, grid_z_mesh), - method="nearest", - ) - - return grid_r, grid_z, psi_grid - except Exception as e: - logger.warning(f"Failed to calculate heatmap interpolation: {e}") - return None, None, None + 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): @@ -503,16 +451,9 @@ def _plot_contours(self): else: contours = self._calc_contours(equilibrium, self.levels) - if self.show_heatmap: - return contours.opts( - color="white", - alpha=0.6, - line_width=1, - colorbar=False, - tools=["hover"], - show_legend=False, - ) - return contours.opts(self.CONTOUR_OPTS) + return contours.opts( + self.CONTOUR_ON_HEATMAP_OPTS if self.show_heatmap else self.CONTOUR_OPTS + ) def _calc_contours(self, equilibrium, levels): """Calculates the contours of the psi grid of an equilibrium IDS. diff --git a/waveform_editor/gui/shape_editor/settings_modal.py b/waveform_editor/gui/shape_editor/settings_modal.py index 5eae5f5c..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,38 +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.nice_plotter.param.watch( - lambda e: setattr(self._contour_detail, "visible", e.new), - ["show_contour"], - ) - - self._heatmap_detail = pn.Column( - self._section_label("Heatmap Detail"), - self._settings_section( - pn.Param( - self.nice_plotter.param, - parameters=["heatmap_alpha"], - show_name=False, - ), - ), - visible=self.nice_plotter.show_heatmap, - 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._heatmap_detail, "visible", e.new), - ["show_heatmap"], + self._heatmap_detail = self._detail_section( + "Heatmap Detail", ["heatmap_alpha"], self.nice_plotter.param.show_heatmap ) display_content = pn.Column( From 28be6011664e6dad89c28118aa6f6b250d5ca5b6 Mon Sep 17 00:00:00 2001 From: Sebbe Blokhuizen Date: Tue, 6 Oct 2026 16:29:29 +0200 Subject: [PATCH 3/3] deduplication --- .../gui/shape_editor/nice_plotter.py | 44 ++++++++++--------- 1 file changed, 24 insertions(+), 20 deletions(-) diff --git a/waveform_editor/gui/shape_editor/nice_plotter.py b/waveform_editor/gui/shape_editor/nice_plotter.py index a0b09934..897070f4 100644 --- a/waveform_editor/gui/shape_editor/nice_plotter.py +++ b/waveform_editor/gui/shape_editor/nice_plotter.py @@ -394,18 +394,12 @@ def _plot_heatmap(self): if not self.show_heatmap or equilibrium is None: return self._empty_heatmap() - 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 self._empty_heatmap() if self._heatmap_cache[0] is not equilibrium: - self._heatmap_cache = (equilibrium, self._calc_heatmap(r, z, psi)) + 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) @@ -455,6 +449,24 @@ def _plot_contours(self): 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. @@ -466,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):