From 29656fd78e35f9096b10ff6eb74c01a9feb68bc5 Mon Sep 17 00:00:00 2001 From: allevitan Date: Sat, 21 Mar 2026 21:27:16 +0100 Subject: [PATCH] Improve the colorized plotting further, and add a colorbar which will be useful for publishing figures now that it's not just a simple hsv lookup. Also fix a bug with nonresponsive windows when all model plots are closed, but the dataset plots are still showing. --- examples/fancy_ptycho.py | 1 + src/cdtools/datasets/ptycho_2d_dataset.py | 2 - src/cdtools/models/base.py | 2 +- src/cdtools/models/fancy_ptycho.py | 22 ++++---- src/cdtools/reconstructors/base.py | 14 ++--- src/cdtools/tools/plotting/plotting.py | 62 +++++++++++++++++------ 6 files changed, 65 insertions(+), 38 deletions(-) diff --git a/examples/fancy_ptycho.py b/examples/fancy_ptycho.py index eb9de73..51d0e0f 100644 --- a/examples/fancy_ptycho.py +++ b/examples/fancy_ptycho.py @@ -14,6 +14,7 @@ model = cdtools.models.FancyPtycho.from_dataset( propagation_distance=5e-3, # Propagate the initial probe guess by 5 mm units='mm', # Set the units for the live plots obj_view_crop=-50, # Expands the field of view in the object plot by 50 pix, + panel_plot_mode=False ) if t.cuda.is_available(): diff --git a/src/cdtools/datasets/ptycho_2d_dataset.py b/src/cdtools/datasets/ptycho_2d_dataset.py index adfe866..5a84ca1 100644 --- a/src/cdtools/datasets/ptycho_2d_dataset.py +++ b/src/cdtools/datasets/ptycho_2d_dataset.py @@ -245,8 +245,6 @@ class Ptycho2DDataset(CDataset): return np.log10((meas_data * mask) + log_offset) else: return meas_data * mask - - translations = self.translations.detach().cpu().numpy() # This takes about twice as long as it would to just do it all at # once, but it avoids creating another self.patterns-sized array diff --git a/src/cdtools/models/base.py b/src/cdtools/models/base.py index 36d9156..bc7f73d 100644 --- a/src/cdtools/models/base.py +++ b/src/cdtools/models/base.py @@ -812,7 +812,7 @@ class CDIModel(t.nn.Module): except KeyboardInterrupt: raise except Exception: - pass + raise rendered.append(fig) diff --git a/src/cdtools/models/fancy_ptycho.py b/src/cdtools/models/fancy_ptycho.py index ccb12ef..dbcfc63 100644 --- a/src/cdtools/models/fancy_ptycho.py +++ b/src/cdtools/models/fancy_ptycho.py @@ -924,7 +924,7 @@ class FancyPtycho(CDIModel): self.corrected_translations(dataset), self.get_probe_intensities(), fig=fig, - cmap='magma', + cmap='viridis', cmap_label='Intensity (a.u.)', units=self.units, convention='probe', @@ -1009,25 +1009,26 @@ class FancyPtycho(CDIModel): 'condition': lambda self: self.exponentiate_obj, }, { - 'title': 'Basis Probes, Colorized', + 'title': 'Probe Modes, Colorized', 'subplot': (0,1), 'plot_func': lambda self, fig: p.plot_colorized( (self.probe if not self.fourier_probe else tools.propagators.inverse_far_field(self.probe)), fig=fig, - title='Basis Probes', + title='Probe Modes, Real Space', basis=self.probe_basis, additional_axis_labels=['Mode #',], + amplitude_scaling=np.sqrt, units=self.units), }, { - 'title': 'Basis Probes, Amplitude', + 'title': 'Probe Modes, Amplitude', 'subplot': (1,1), 'plot_func': lambda self, fig: p.plot_amplitude( (self.probe if not self.fourier_probe else tools.propagators.inverse_far_field(self.probe)), fig=fig, - title='Basis Probes', + title='Probe Modes, Real Space', basis=self.probe_basis, additional_axis_labels=['Mode #',], units=self.units), @@ -1041,24 +1042,25 @@ class FancyPtycho(CDIModel): 'grid': (2,3), 'plots': [ { - 'title': 'Basis Probes, Fourier Colorized', + 'title': 'Probe Modes, Fourier Colorized', 'subplot': (0,0), 'plot_func': lambda self, fig: p.plot_colorized( (self.probe if self.fourier_probe else tools.propagators.far_field(self.probe)), fig=fig, - title='Basis Probes, Fourier', + title='Probe Modes, Fourier Space', additional_axis_labels=['Mode #',], + amplitude_scaling = np.sqrt, ), }, { - 'title': 'Basis Probes, Fourier Amplitude', + 'title': 'Probe Modes, Fourier Amplitude', 'subplot': (1,0), 'plot_func': lambda self, fig: p.plot_amplitude( (self.probe if self.fourier_probe else tools.propagators.far_field(self.probe)), fig=fig, - title='Basis Probes, Fourier', + title='Probe Modes, Fourier Space', additional_axis_labels=['Mode #',], ), }, @@ -1070,7 +1072,7 @@ class FancyPtycho(CDIModel): { 'title': 'Detector Background', 'subplot': (1,1), - 'plot_func': lambda self, fig: p.plot_amplitude(self.background**2, fig=fig, cmap='magma', cmap_label='Intensity (detector units)'), + 'plot_func': lambda self, fig: p.plot_amplitude(self.background**2, fig=fig, cmap='viridis', cmap_label='Intensity (detector units)'), }, { 'title': 'Corrected Translations', diff --git a/src/cdtools/reconstructors/base.py b/src/cdtools/reconstructors/base.py index a9698ee..90a5287 100644 --- a/src/cdtools/reconstructors/base.py +++ b/src/cdtools/reconstructors/base.py @@ -233,7 +233,7 @@ class Reconstructor: def optimize(self, iterations: int, batch_size: int = 1, - custom_data_loader: torch.utils.data.DataLoader = None, + custom_data_loader: t.utils.data.DataLoader = None, regularization_factor: Union[float, List[float]] = None, thread: bool = True, calculation_width: int = 10, @@ -358,14 +358,10 @@ class Reconstructor: try: calc.start() while calc.is_alive(): - figs = getattr(self.model, 'figs', []) - open_fig = next( - (f for f in figs if plt.fignum_exists(f.number)), - None, - ) - if open_fig is not None: - open_fig.canvas.flush_events() - # We need a low value for smooth figure responses + open_figs = plt.get_fignums() + with plt.rc_context({'figure.raise_window': False}): + for fignum in open_figs: + plt.figure(fignum).canvas.flush_events() time.sleep(0.001) except KeyboardInterrupt as e: diff --git a/src/cdtools/tools/plotting/plotting.py b/src/cdtools/tools/plotting/plotting.py index bec5122..ca8e53f 100644 --- a/src/cdtools/tools/plotting/plotting.py +++ b/src/cdtools/tools/plotting/plotting.py @@ -31,7 +31,7 @@ __all__ = [ ] -def colorize(z, use_cmocean=True): +def colorize(z, use_cmocean=False, amplitude_scaling=lambda x: x): """ Returns RGB values for a complex color plot given a complex array This function returns a set of RGB values that can be used directly in a call to imshow based on an input complex numpy array (not a @@ -41,6 +41,8 @@ def colorize(z, use_cmocean=True): ---------- z : array A complex-valued array + use_cmocean : bool + If true, uses the cmocean_phase colormap instead of hue Returns ------- rgb : list(array) @@ -48,24 +50,25 @@ def colorize(z, use_cmocean=True): """ amp = np.abs(z) - scaled_amp = amp / np.max(amp) + scaled_amp = amplitude_scaling(amp / np.max(amp)) ph = np.angle(z, deg=1) - if not use_cmocean: - # HSV are values in range [0,1] - h = ((ph + 90) % 360) / 360 - s = 0.85 * np.ones_like(h) - v = scaled_amp - return hsv_to_rgb(np.dstack((h,s,v))) - else: + if use_cmocean: base_rgb_values = [] for channel in range(3): - base_rgb_values.append(np.interp(ph%360, + base_rgb_values.append(np.interp((ph + 180)%360, np.linspace(0, 360, cm_data.shape [0]), cm_data[:,channel])) base_rgb_values = np.dstack(base_rgb_values) rgb_values = base_rgb_values * scaled_amp[...,None] return rgb_values + else: + # HSV are values in range [0,1] + h = ((ph + 90) % 360) / 360 + s = 0.85 * np.ones_like(h) + v = scaled_amp + return hsv_to_rgb(np.dstack((h,s,v))) + @@ -116,6 +119,7 @@ def plot_image( interpolation=None, title=None, additional_axis_labels=None, + updateable_colorbar=True, **kwargs ): """Plots an image with a colorbar and on an appropriate spatial grid @@ -219,7 +223,8 @@ def plot_image( # don't "reset" the home positions of the toolbar if hasattr(fig, '_current_im'): fig._current_im.set_data(to_plot) - fig._current_im.autoscale() + if updateable_colorbar: + fig._current_im.autoscale() # We need to go to the "home" position before updating it # to include the new data, because otherwise it will store # other axes (potentially zoomed in) positions as "home", @@ -344,6 +349,8 @@ def plot_image( if show_cbar: cbar = fig.colorbar(mpl_im, ax=ax, fraction=0.05, pad=0.05, location='right') + if not updateable_colorbar: + cbar.ax.set_navigate(False) if cmap_label is not None: cbar.set_label(cmap_label) @@ -619,7 +626,7 @@ def plot_phase( def plot_amplitude_surfacenorm(): pass -def plot_colorized(im, fig=None, basis=None, units='$\\mu$m', title=None, **kwargs): +def plot_colorized(im, fig=None, basis=None, units='$\\mu$m', title=None, amplitude_scaling=lambda x: x, **kwargs): """ Plots the colorized version of a complex array with dimensions NxM The darkness corresponds to the intensity of the image, and the color @@ -649,11 +656,34 @@ def plot_colorized(im, fig=None, basis=None, units='$\\mu$m', title=None, **kwar used_fig : matplotlib.figure.Figure The figure object that was actually plotted to. """ - plot_func = lambda x: colorize(x) - return plot_image(im, plot_func=plot_func, fig=fig, basis=basis, + plot_func = lambda x: colorize(x, use_cmocean=True, + amplitude_scaling=amplitude_scaling) + plot_fig = plot_image(im, plot_func=plot_func, fig=fig, basis=basis, cmap=cmocean_phase, vmin=-np.pi, vmax=np.pi, cmap_label='Phase (rad)', - units=units, show_cbar=True, title=title, **kwargs) + units=units, show_cbar=True, title=title, + updateable_colorbar=False, **kwargs) + + # Find the colorbar - this is a bit hacky + cbar_ax = [ax for ax in plot_fig.get_axes() if hasattr(ax, '_colorbar')][0] + + # --- Replace the colorbar image --- + # The internal image is a QuadMesh living on cbar.ax + #qm = cbar.ax.collections + # Build a 2D array to match the colorbar's range, here (0,1) in x + # and (pi, pi) in y + yg = np.linspace(-np.pi, np.pi, 256) # colormap values + xg = np.linspace(0, 1, 64) # second dimension + YY, XX = np.meshgrid(yg, xg, indexing='ij') + dummy_im = XX * np.exp(1j*YY) + cbar_im = plot_func(dummy_im) + for artist in list(cbar_ax.get_children()): + if 'QuadMesh' in type(artist).__name__ or \ + 'AxesImage' in type(artist).__name__: + artist.remove() + + cbar_ax.imshow(cbar_im, origin='lower', aspect='auto', + extent=[xg[0], xg[-1], -np.pi, np.pi]) def plot_translations(translations, fig=None, units='$\\mu$m', lines=True, invert_xaxis=True, clear_fig=True, label=None, color=None, marker='.', **kwargs): @@ -880,7 +910,7 @@ def plot_nanomap_with_images(translations, get_image_func, values=None, mask=Non nanomap_units_factor = get_units_factor(nanomap_units) nanomap = axes[0].scatter(nanomap_units_factor * translations[:,0], nanomap_units_factor * translations[:,1], - s=s,c=values, picker=True) + s=s,c=values, picker=True, cmap=cmap) axes[0].invert_xaxis() axes[0].set_facecolor('k')