diff --git a/examples/fancy_ptycho.py b/examples/fancy_ptycho.py index ebee68f..b7e6fbf 100644 --- a/examples/fancy_ptycho.py +++ b/examples/fancy_ptycho.py @@ -1,4 +1,5 @@ import cdtools +import torch as t from matplotlib import pyplot as plt filename = 'example_data/lab_ptycho_data.cxi' @@ -12,12 +13,15 @@ model = cdtools.models.FancyPtycho.from_dataset( probe_support_radius=120, # Force the probe to 0 outside a radius of 120 pix 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 + obj_view_crop=-50, # Expands the field of view in the object plot by 50 pix, + exponentiate_obj=False, + panel_plot_mode=True, + plot_level=2, ) -device = 'cuda' -model.to(device=device) -dataset.get_as(device=device) +if t.cuda.is_available(): + model.to(device='cuda') + dataset.get_as(device='cuda') # For this script, we use a slightly different pattern where we explicitly # create a `Reconstructor` class to orchestrate the reconstruction. The @@ -26,13 +30,14 @@ dataset.get_as(device=device) # e.g. estimates of the moments of individual parameters recon = cdtools.reconstructors.AdamReconstructor(model, dataset) + # The learning rate parameter sets the alpha for Adam. # The beta parameters are (0.9, 0.999) by default # The batch size sets the minibatch size for loss in recon.optimize(50, lr=0.02, batch_size=10): print(model.report()) # Plotting is expensive, so we only do it every tenth epoch - if model.epoch % 10 == 0: + if model.epoch % 2 == 0: model.inspect(dataset) # It's common to chain several different reconstruction loops. Here, we diff --git a/examples/simple_ptycho.py b/examples/simple_ptycho.py index 571f24f..467b9af 100644 --- a/examples/simple_ptycho.py +++ b/examples/simple_ptycho.py @@ -8,6 +8,7 @@ more powerful FancyPtycho model and include more information on how to correct for common sources of error. """ import cdtools +import torch as t from matplotlib import pyplot as plt # We load an example dataset from a .cxi file @@ -15,12 +16,12 @@ filename = 'example_data/lab_ptycho_data.cxi' dataset = cdtools.datasets.Ptycho2DDataset.from_cxi(filename) # We create a ptychography model from the dataset -model = cdtools.models.SimplePtycho.from_dataset(dataset, panel_plot_mode=True) +model = cdtools.models.SimplePtycho.from_dataset(dataset) -# We move the model to the GPU -device = 'cuda' -model.to(device=device) -dataset.get_as(device=device) +# We move the model to the GPU, if possible +if t.cuda.is_available(): + model.to(device='cuda') + dataset.get_as(device='cuda') model.inspect(dataset) print('hi') @@ -32,7 +33,6 @@ for loss in model.Adam_optimize(30, dataset, batch_size=10): if model.epoch % 10 == 0: model.inspect(dataset) - # We study the results model.inspect(dataset, replot_all=True) model.compare(dataset) diff --git a/src/cdtools/models/base.py b/src/cdtools/models/base.py index 2dba377..8decc4d 100644 --- a/src/cdtools/models/base.py +++ b/src/cdtools/models/base.py @@ -28,6 +28,7 @@ loss """ +from sympy import Q import torch as t from torch.utils import data as torchdata import matplotlib @@ -604,7 +605,7 @@ class CDIModel(t.nn.Module): # from the panels flat = [] for panel in plot_panel_list: - panel_level = panel.get('plot_level', 0) + panel_level = panel.get('plot_level', 1) for plot in panel['plots']: # We add the plot level from the larger panel flat.append({**plot, 'plot_level': panel_level}) @@ -628,9 +629,14 @@ class CDIModel(t.nn.Module): self ): backend = matplotlib.get_backend().lower() - interactive_bk = matplotlib.backends.backend_registry.list_builtin( - matplotlib.backends.BackendFilter.INTERACTIVE - ) + try: + # matplotlib >= 3.9 + interactive_bk = matplotlib.backends.backend_registry.list_builtin( + matplotlib.backends.BackendFilter.INTERACTIVE + ) + except AttributeError: + # older matplotlib + interactive_bk = matplotlib.rcsetup.interactive_bk return backend in [b.lower() for b in interactive_bk] @@ -656,7 +662,7 @@ class CDIModel(t.nn.Module): for plot in plot_list: # Level filter - if plot.get('plot_level', 0) > self.plot_level: + if plot.get('plot_level', 1) > self.plot_level: continue # Condition check @@ -670,15 +676,17 @@ class CDIModel(t.nn.Module): continue if self.has_inspect_been_called and \ - replot_all == False and \ + not replot_all and \ not plt.fignum_exists(plot['title']): continue if not self.has_inspect_been_called: - fig = plt.figure(plot['title']) + fig = plt.figure(plot['title'], + constrained_layout=True) else: with plt.rc_context({'figure.raise_window': False}): - fig = plt.figure(plot['title']) + fig = plt.figure(plot['title'], + constrained_layout=True) try: plot['plot_func'](self, fig) @@ -711,34 +719,40 @@ class CDIModel(t.nn.Module): rendered = [] - for panel_idx, panel_def in enumerate(plot_panel_list): - panel_level = panel_def.get('plot_level', 0) + for panel_def in plot_panel_list[::-1]: # Flip so first ones show on top + panel_level = panel_def.get('plot_level', 1) if panel_level > self.plot_level: continue # skip entire panel nrows, ncols = panel_def['grid'] figsize = panel_def.get('figure_size', None) - title = panel_def.get('title', '') - + if self.has_inspect_been_called and \ - replot_all == False and \ + not replot_all and \ not plt.fignum_exists(panel_def['title']): continue if not self.has_inspect_been_called: - fig = plt.figure(panel_def['title']) + fig = plt.figure(panel_def['title'], figsize=figsize, + constrained_layout=True) else: with plt.rc_context({'figure.raise_window': False}): - fig = plt.figure(panel_def['title']) + fig = plt.figure(panel_def['title'], + constrained_layout=True) - # Remove all axes and recreate them fresh each update. - # plt.colorbar() shrinks the parent axes to make room for - # itself, so clearing and recreating is simpler than trying - # to undo that resizing. - for ax in list(fig.axes): - ax.remove() + fig.clear() + + fig.get_layout_engine().set( + rect=(0.02, 0.02, 0.96, 0.96), + ) + gs = fig.add_gridspec( + nrows, ncols, + width_ratios=[1]*ncols, + height_ratios=[1]*nrows, + ) + for plot in panel_def['plots']: condition = plot.get('condition', None) if condition is not None: @@ -748,25 +762,22 @@ class CDIModel(t.nn.Module): except TypeError: if not condition(self, dataset): continue - - row, col = plot['subplot'] - position = row * ncols + col + 1 # 1-indexed for matplotlib - - ax_key = (panel_idx, row, col) - ax = fig.add_subplot(nrows, ncols, position) + subfig = fig.add_subfigure(gs[plot['subplot'][0], + plot['subplot'][1]]) try: - plot['plot_func'](self, ax) - ax.set_title(plot['title']) + plot['plot_func'](self, subfig) + plt.gca().set_title(plot['title']) except TypeError: if dataset is not None: try: - plot['plot_func'](self, ax, dataset) - ax.set_title(plot['title']) - except Exception: + plot['plot_func'](self, subfig, dataset) + plt.gca().set_title(plot['title']) + except TypeError:#Exception: pass - except Exception: - pass + #except Exception: + # pass + rendered.append(fig) if self._is_backend_interactive(): @@ -775,6 +786,41 @@ class CDIModel(t.nn.Module): return rendered + def plot_loss_history(self, fig=None, clear_fig=True): + """Plots the loss history on a semilogy axis + + Parameters + ---------- + fig : matplotlib.figure.Figure + Default is a new figure, a matplotlib figure to use to plot + clear_fig : bool + Default is True. Whether to clear the figure before plotting. + + Returns + ------- + used_fig : matplotlib.figure.Figure + The figure object that was actually plotted to. + """ + + if fig is None: + fig = plt.figure() + + if clear_fig: + fig.clear() + + if len(fig.axes) >= 1: + ax = fig.axes[0] + else: + ax = fig.add_subplot(111) + + ax.semilogy(self.loss_history) + plt.title('Loss History') + + ax.set_xlabel('Epoch') + ax.set_ylabel('Loss Metric') + + return fig + def save_figures(self, prefix='', extension='.pdf'): """Saves all currently open inspection figures. diff --git a/src/cdtools/models/fancy_ptycho.py b/src/cdtools/models/fancy_ptycho.py index 0b3d499..fb9d4da 100644 --- a/src/cdtools/models/fancy_ptycho.py +++ b/src/cdtools/models/fancy_ptycho.py @@ -45,9 +45,12 @@ class FancyPtycho(CDIModel): near_field=False, angular_spectrum_propagator=None, inv_angular_spectrum_propagator=None, + panel_plot_mode=False, + plot_level=2, ): - super(FancyPtycho, self).__init__() + super(FancyPtycho, self).__init__(panel_plot_mode=panel_plot_mode, + plot_level=plot_level) self.register_buffer('wavelength', t.as_tensor(wavelength, dtype=dtype)) self.store_detector_geometry(detector_geometry, @@ -252,6 +255,8 @@ class FancyPtycho(CDIModel): obj_view_crop=None, obj_padding=200, near_field=False, + panel_plot_mode=False, + plot_level=2, ): wavelength = dataset.wavelength @@ -517,6 +522,8 @@ class FancyPtycho(CDIModel): near_field=near_field, angular_spectrum_propagator=angular_spectrum_propagator, inv_angular_spectrum_propagator=inv_angular_spectrum_propagator, + panel_plot_mode=panel_plot_mode, + plot_level=plot_level, ) @@ -887,7 +894,22 @@ class FancyPtycho(CDIModel): cmap=cmap, **kwargs), + + def plot_illumination_intensity(self, fig, dataset): + if not hasattr(self, 'weights') or self.weights.ndim != 1: + raise NotImplementedError('Not yet implemented for OPRP') + p.plot_nanomap( + self.corrected_translations(dataset), + self.weights**2, + fig=fig, + cmap='magma', + cmap_label='Intensity (a.u.)', + units=self.units, + convention='probe', + invert_xaxis=True + ) + def plot_translations_and_originals(self, fig, dataset): """Only used to make a plot for the plot list.""" p.plot_translations( @@ -910,62 +932,160 @@ class FancyPtycho(CDIModel): plt.legend() + plot_panel_list = [ + { + 'title': 'Main Results', + 'plot_level': 1, + 'grid': (2,2), + 'figure_size': (9,7), + 'plots': [ + { + 'title': 'Object Phase', + 'subplot': (0,0), + 'plot_func': lambda self, fig: p.plot_phase( + self.obj[self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units), + 'condition': lambda self: not self.exponentiate_obj, + }, + { + 'title': 'Object Amplitude', + 'subplot': (1,0), + 'plot_func': lambda self, fig: p.plot_amplitude( + self.obj[self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units), + 'condition': lambda self: not self.exponentiate_obj, + }, + { + 'title': 'Real Part of T', + 'subplot': (0,0), + 'plot_func': lambda self, fig: p.plot_real( + self.obj[self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units, + cmap='cividis'), + 'condition': lambda self: self.exponentiate_obj, + }, + { + 'title': 'Imaginary Part of T', + 'subplot': (1,0), + 'plot_func': lambda self, fig: p.plot_imag( + self.obj[self.obj_view_slice], + fig=fig, + basis=self.obj_basis, + units=self.units), + 'condition': lambda self: self.exponentiate_obj, + }, + { + 'title': 'Basis Probes, 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, + basis=self.probe_basis, + units=self.units), + }, + { + 'title': 'Basis Probes, 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, + basis=self.probe_basis, + units=self.units), + }, + ], + }, + { + 'title': 'Advanced Monitoring', + 'plot_level': 2, + 'figure_size': (12,7), + 'grid': (2,3), + 'plots': [ + { + 'title': 'Basis Probes, 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 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': 'Illumination Intensity', + 'subplot': (0,1), + 'plot_func': lambda self, fig, dataset: self.plot_illumination_intensity(fig, dataset), + 'condition': lambda self: hasattr(self, 'weights') and self.weights.ndim == 1 + }, + { + '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)'), + }, + { + 'title': 'Corrected Translations', + 'subplot': (0,2), + 'plot_func': lambda self, fig, dataset: self.plot_translations_and_originals(fig, dataset), + }, + { + 'title': 'Loss History', + 'subplot': (1,2), + 'plot_func': lambda self, fig: self.plot_loss_history(fig), + }, + ], + }, + ] + plot_list = [ - ('', - lambda self, fig, dataset: self.plot_wavefront_variation( + {'title': 'Per-Exposure Probe Intensity', + 'plot_level': 3, + 'plot_func': lambda self, fig, dataset: self.plot_wavefront_variation( dataset, fig=fig, mode='root_sum_intensity', image_title='Root Summed Probe Intensities', image_colorbar_title='Square Root of Intensity'), - lambda self: len(self.weights.shape) >= 2), - ('', - lambda self, fig, dataset: self.plot_wavefront_variation( + 'condition': lambda self: len(self.weights.shape) >= 2}, + {'title': 'Per-Exposure Probe Amplitudes', + 'plot_level': 3, + 'plot_func': lambda self, fig, dataset: self.plot_wavefront_variation( dataset, fig=fig, mode='amplitude', image_title='Probe Amplitudes (scroll to view modes)', image_colorbar_title='Probe Amplitude'), - lambda self: len(self.weights.shape) >= 2), - ('', - lambda self, fig, dataset: self.plot_wavefront_variation( + 'condition': lambda self: len(self.weights.shape) >= 2}, + {'title': 'Per-Exposure Probe Phases', + 'plot_level': 3, + 'plot_func': lambda self, fig, dataset: self.plot_wavefront_variation( dataset, fig=fig, mode='phase', image_title='Probe Phases (scroll to view modes)', image_colorbar_title='Probe Phase'), - lambda self: len(self.weights.shape) >= 2), - ('Basis Probe Fourier Space Amplitudes', - lambda self, fig: p.plot_amplitude( - (self.probe if self.fourier_probe - else tools.propagators.far_field(self.probe)), - fig=fig)), - ('Basis Probe Fourier Space Colorized', - lambda self, fig: p.plot_colorized( - (self.probe if self.fourier_probe - else tools.propagators.far_field(self.probe)) - , fig=fig)), - ('Basis Probe Real Space Amplitudes', - lambda self, fig: p.plot_amplitude( - (self.probe if not self.fourier_probe - else tools.propagators.inverse_far_field(self.probe)), - fig=fig, - basis=self.probe_basis, - units=self.units)), - ('Basis Probe Real Space Colorized', - lambda self, fig: p.plot_colorized( - (self.probe if not self.fourier_probe - else tools.propagators.inverse_far_field(self.probe)), - fig=fig, - basis=self.probe_basis, - units=self.units)), - ('Average Weight Matrix Amplitudes', - lambda self, fig: p.plot_amplitude( + 'condition': lambda self: len(self.weights.shape) >= 2}, + {'title': 'Average Weight Matrix Amplitudes', + 'plot_level': 1, + 'plot_func': lambda self, fig: p.plot_amplitude( np.nanmean(np.abs(self.weights.data.cpu().numpy()), axis=0), fig=fig), - lambda self: len(self.weights.shape) >= 2), - ('% of Power in Top Mode', - lambda self, fig, dataset: p.plot_nanomap( + 'condition': lambda self: len(self.weights.shape) >= 2}, + {'title': '% of Power in Top Mode', + 'plot_level': 3, + 'plot_func': lambda self, fig, dataset: p.plot_nanomap( self.corrected_translations(dataset), 100 * t.stack([ analysis.calc_mode_power_fractions( @@ -975,44 +1095,11 @@ class FancyPtycho(CDIModel): ], dim=0), fig=fig, units=self.units), - lambda self: len(self.weights.shape) >= 2), - ('Object Amplitude', - lambda self, fig: p.plot_amplitude( - self.obj[self.obj_view_slice], - fig=fig, - basis=self.obj_basis, - units=self.units), - lambda self: not self.exponentiate_obj), - ('Object Phase', - lambda self, fig: p.plot_phase( - self.obj[self.obj_view_slice], - fig=fig, - basis=self.obj_basis, - units=self.units), - lambda self: not self.exponentiate_obj), - ('Real Part of T', - lambda self, fig: p.plot_real( - self.obj[self.obj_view_slice], - fig=fig, - basis=self.obj_basis, - units=self.units, - cmap='cividis'), - lambda self: self.exponentiate_obj), - ('Imaginary Part of T', - lambda self, fig: p.plot_imag( - self.obj[self.obj_view_slice], - fig=fig, - basis=self.obj_basis, - units=self.units), - lambda self: self.exponentiate_obj), - - ('Corrected Translations', - lambda self, fig, dataset: self.plot_translations_and_originals(fig, dataset)), - ('Background', - lambda self, fig: p.plot_amplitude(self.background**2, fig=fig)), - ('Quantum Efficiency Mask', - lambda self, fig: p.plot_amplitude(self.qe_mask, fig=fig), - lambda self: (hasattr(self, 'qe_mask') and self.qe_mask is not None)) + 'condition': lambda self: len(self.weights.shape) >= 2}, + {'title': 'Quantum Efficiency Mask', + 'plot_level': 3, + 'plot_func': lambda self, fig: p.plot_amplitude(self.qe_mask, fig=fig), + 'condition': lambda self: (hasattr(self, 'qe_mask') and self.qe_mask is not None)}, ] diff --git a/src/cdtools/models/simple_ptycho.py b/src/cdtools/models/simple_ptycho.py index 0460ecb..234ed5d 100644 --- a/src/cdtools/models/simple_ptycho.py +++ b/src/cdtools/models/simple_ptycho.py @@ -15,15 +15,10 @@ class SimplePtycho(CDIModel): probe_guess, obj_guess, min_translation = [0,0], - panel_plot_mode=False, - plot_level=1, ): # We initialize the superclass - super().__init__( - panel_plot_mode=panel_plot_mode, - plot_level=plot_level, - ) + super().__init__() # We register all the constants, like wavelength, as buffers. This # lets the model hook into some nice pytorch features, like using @@ -48,8 +43,7 @@ class SimplePtycho(CDIModel): @classmethod - def from_dataset(cls, dataset,panel_plot_mode=False, - plot_level=1, ): + def from_dataset(cls, dataset): # We get the key geometry information from the dataset wavelength = dataset.wavelength @@ -83,8 +77,6 @@ class SimplePtycho(CDIModel): probe, obj, min_translation=min_translation, - panel_plot_mode=panel_plot_mode, - plot_level=plot_level, ) @@ -115,58 +107,31 @@ class SimplePtycho(CDIModel): # This lists all the plots to display on a call to model.inspect() - plot_panel_list = [ - { - # Title for window - 'title' : 'Probe Results', - # (width, height) in inches - 'figure_size': (7, 7), - # (nrows, ncols) for subplot grid - 'grid': (2, 2), - # A setting to control how many plots are produced - 'plot_level': 1, - # The list of plots to include - 'plots' : [ - { - 'title': 'Probe Amplitude', - 'subplot' : (0, 0), - 'plot_func': lambda self, fig: - p.plot_amplitude(self.probe, fig, - basis=self.probe_basis) - },{ - 'title': 'Probe Phase', - 'subplot' : (0, 1), - 'plot_func': lambda self, fig: - p.plot_phase(self.probe, fig, - basis=self.probe_basis) - } - ] - }, - { - # Title for window - 'title' : 'Object Results', - # (width, height) in inches - 'figure_size': (7, 7), - # (nrows, ncols) for subplot grid - 'grid': (2, 2), - # A setting to control how many plots are produced - 'plot_level': 1, - # The list of plots to include - 'plots' : [ - { - 'title': 'Object Amplitude', - 'subplot' : (1, 0), - 'plot_func': lambda self, fig: - p.plot_amplitude(self.obj, fig, - basis=self.probe_basis) - }, { - 'title': 'Object Phase', - 'subplot' : (1, 1), - 'plot_func': lambda self, fig: - p.plot_phase(self.obj, fig, - basis=self.probe_basis) - }, - ] + plot_list = [ + { + 'title': 'Probe Amplitude', + 'subplot' : (0, 0), + 'plot_func': lambda self, fig: + p.plot_amplitude(self.probe, fig, + basis=self.probe_basis), + }, { + 'title': 'Probe Phase', + 'subplot' : (0, 1), + 'plot_func': lambda self, fig: + p.plot_phase(self.probe, fig, + basis=self.probe_basis) + }, { + 'title': 'Object Amplitude', + 'subplot' : (1, 0), + 'plot_func': lambda self, fig: + p.plot_amplitude(self.obj, fig, + basis=self.probe_basis) + }, { + 'title': 'Object Phase', + 'subplot' : (1, 1), + 'plot_func': lambda self, fig: + p.plot_phase(self.obj, fig, + basis=self.probe_basis) }, ] diff --git a/src/cdtools/tools/plotting/plotting.py b/src/cdtools/tools/plotting/plotting.py index e820764..367c906 100644 --- a/src/cdtools/tools/plotting/plotting.py +++ b/src/cdtools/tools/plotting/plotting.py @@ -92,7 +92,6 @@ def get_units_factor(units): factor=1e12 return factor - def plot_image( im, plot_func=lambda x: x, @@ -164,28 +163,18 @@ def plot_image( else: im = im.detach().cpu().numpy() - # Support passing an Axes object instead of a Figure - ax_mode = isinstance(fig, plt.Axes) - if ax_mode: - ax = fig - fig = ax.get_figure() - elif fig is None: + if fig is None: fig = plt.figure() - ax = fig.add_subplot(111, **kwargs) - # This nukes everything and updates either the appropriate image from the # stack of images, or the only image if only a single image has been # given def make_plot(idx): - plt.figure(fig.number) - if ax_mode: - title = ax.get_title() - ax.cla() - plt.sca(ax) - else: - title = plt.gca().get_title() - fig.clear() - + #plt.figure(fig.number) + #title = plt.gca().get_title() + try: + title = fig.axes[0].get_title() + except IndexError: + title = '' # If im only has two dimensions, this reshape will add a leading # dimension, and update will be called on index 0. If it has 3 or more @@ -194,19 +183,41 @@ def plot_image( s = im.shape reshaped_im = im.reshape(-1,s[-2],s[-1]) num_images = reshaped_im.shape[0] - plot_holder = ax if ax_mode else fig - plot_holder.plot_idx = idx % num_images + fig.plot_idx = idx % num_images - to_plot = plot_func(reshaped_im[plot_holder.plot_idx]) + to_plot = plot_func(reshaped_im[fig.plot_idx]) - mpl_im = plt.imshow( + # By only updating the data, and not redrawing the fig, we + # don't "reset" the home positions of the other + if hasattr(fig, '_current_im'): + print('Just changing data') + fig._current_im.set_data(to_plot) + 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", + # which is super annoying, more so than the reset. + if fig.canvas.toolbar is not None: + fig.canvas.toolbar.home() + fig.canvas.toolbar.update() + # Replace existing mode number + for artist in fig.texts: + artist.set_text(f'Mode {fig.plot_idx}') + + return fig + + fig.clear() + ax = fig.add_subplot(111, **kwargs) + + mpl_im = ax.imshow( to_plot, cmap = cmap, interpolation = interpolation, vmin=vmin, vmax=vmax, ) - plt.gca().set_facecolor('k') + fig._current_im = mpl_im + ax.set_facecolor('k') if basis is not None: # we've closed over basis, so we can't edit it @@ -264,50 +275,51 @@ def plot_image( corners = np.matmul(transform_matrix,corners.transpose()) mins = np.min(corners, axis=1) maxes = np.max(corners, axis=1) - plt.gca().set_xlim([mins[0], maxes[0]]) - plt.gca().set_ylim([mins[1], maxes[1]]) - plt.gca().invert_yaxis() + ax.set_xlim([mins[0], maxes[0]]) + ax.set_ylim([mins[1], maxes[1]]) + ax.invert_yaxis() if show_cbar: - cbar = plt.colorbar() + cbar = fig.colorbar(mpl_im, ax=ax, fraction=0.05, pad=0.05) if cmap_label is not None: cbar.set_label(cmap_label) if basis is not None: - plt.xlabel('X (' + units + ')') - plt.ylabel('Y (' + units + ')') + ax.set_xlabel('X (' + units + ')') + ax.set_ylabel('Y (' + units + ')') else: - plt.xlabel('j (pixels)') - plt.ylabel('i (pixels)') + ax.set_xlabel('j (pixels)') + ax.set_ylabel('i (pixels)') - - plt.title(title) + ax.set_title(title) if len(im.shape) >= 3: - text_transform = ax.transAxes if ax_mode else plt.gcf().transFigure - plt.text(0.03, 0.03, str(plot_holder.plot_idx), fontsize=14, transform=text_transform) + fig.text(0.03, 0.03, f'Mode {fig.plot_idx}', fontsize=14) + + if fig.canvas.toolbar is not None: + fig.canvas.toolbar.update() return fig - plot_holder = ax if ax_mode else fig - if hasattr(plot_holder, 'plot_idx'): - result = make_plot(plot_holder.plot_idx) + if hasattr(fig, 'plot_idx'): + result_fig = make_plot(fig.plot_idx) else: - result = make_plot(0) - + result_fig = make_plot(0) + update = make_plot - def on_action(event): - plot_holder = ax if ax_mode else fig + # Protection for multi-subfigure situation + if event.inaxes not in fig.axes: + return if not hasattr(event, 'button'): event.button = None if not hasattr(event, 'key'): event.key = None if event.key == 'up' or event.button == 'up': - update(plot_holder.plot_idx - 1) + update(fig.plot_idx - 1) elif event.key == 'down' or event.button == 'down': - update(plot_holder.plot_idx + 1) + update(fig.plot_idx + 1) plt.draw() if len(im.shape) >=3: @@ -320,7 +332,7 @@ def plot_image( fig.my_callbacks.append(fig.canvas.mpl_connect('key_press_event',on_action)) fig.my_callbacks.append(fig.canvas.mpl_connect('scroll_event',on_action)) - return result + return result_fig def plot_real(im, fig = None, basis=None, units='$\\mu$m', cmap='viridis', cmap_label='Real Part (a.u.)', **kwargs): @@ -571,17 +583,16 @@ def plot_translations(translations, fig=None, units='$\\mu$m', lines=True, inver factor = get_units_factor(units) - if isinstance(fig, plt.Axes): - ax = fig - fig = ax.get_figure() - plt.sca(ax) - elif fig is None: + if fig is None: fig = plt.figure() - ax = fig.add_subplot(111, **kwargs) + + if clear_fig: + fig.clear() + + if len(fig.axes) >= 1: + ax = fig.axes[0] else: - plt.figure(fig.number) - if clear_fig: - plt.gcf().clear() + ax = fig.add_subplot(111, **kwargs) if isinstance(translations, t.Tensor): translations = translations.detach().cpu().numpy() @@ -590,25 +601,33 @@ def plot_translations(translations, fig=None, units='$\\mu$m', lines=True, inver linestyle = '-' if lines else 'None' linewidth = 1 if lines else 0 - plt.plot(translations[:,0], translations[:,1], - marker=marker, linestyle=linestyle, - label=label, color=color, - linewidth=linewidth) + ax.plot(translations[:,0], translations[:,1], + marker=marker, linestyle=linestyle, + label=label, color=color, + linewidth=linewidth) if invert_xaxis: - ax = plt.gca() x_min, x_max = ax.get_xlim() # Protect against flipping twice if plotting on top of existing graph if x_min <= x_max: ax.invert_xaxis() - plt.xlabel('X (' + units + ')') - plt.ylabel('Y (' + units + ')') + ax.set_xlabel('X (' + units + ')') + ax.set_ylabel('Y (' + units + ')') return fig -def plot_nanomap(translations, values, fig=None, units='$\\mu$m', convention='probe', invert_xaxis=True): +def plot_nanomap( + translations, + values, + fig=None, + cmap='viridis', + cmap_label=None, + units='$\\mu$m', + convention='probe', + invert_xaxis=True +): """Plots a set of nanomap data in a flexible way Parameters @@ -619,6 +638,10 @@ def plot_nanomap(translations, values, fig=None, units='$\\mu$m', convention='pr A length-N object of values associated with the translations fig : matplotlib.figure.Figure Default is a new figure, a matplotlib figure to use to plot + cmap : str + Default is 'viridis', the colormap to plot with + cmap_label : str + Default is no label, what to label the colorbar when plotting. units : str Default is um, units to report in (assuming input in m) convention : str @@ -632,22 +655,14 @@ def plot_nanomap(translations, values, fig=None, units='$\\mu$m', convention='pr The figure object that was actually plotted to. """ - ax_mode = isinstance(fig, plt.Axes) - if ax_mode: - ax = fig - fig = ax.get_figure() - ax.cla() - plt.sca(ax) - elif fig is None: + if fig is None: fig = plt.figure() - else: - plt.figure(fig.number) - plt.gcf().clear() - + + fig.clear() + ax = fig.add_subplot(111) factor = get_units_factor(units) - plot_area = ax if ax_mode else fig - bbox = plot_area.get_window_extent().transformed(fig.dpi_scale_trans.inverted()) + bbox = fig.get_window_extent().transformed(fig.dpi_scale_trans.inverted()) if isinstance(translations, t.Tensor): trans = translations.detach().cpu().numpy() else: @@ -664,14 +679,17 @@ def plot_nanomap(translations, values, fig=None, units='$\\mu$m', convention='pr s = bbox.width * bbox.height / trans.shape[0] * 72**2 #72 is points per inch s /= 4 # A rough value to make the size work out - plt.scatter(factor * trans[:,0],factor * trans[:,1],s=s,c=values) + scatter_plot = ax.scatter( + factor * trans[:,0],factor * trans[:,1],s=s,c=values, cmap=cmap) if invert_xaxis: - plt.gca().invert_xaxis() + ax.invert_xaxis() - plt.gca().set_facecolor('k') - plt.xlabel('Translation x (' + units + ')') - plt.ylabel('Translation y (' + units + ')') - plt.colorbar() + ax.set_facecolor('k') + ax.set_xlabel('Translation x (' + units + ')') + ax.set_ylabel('Translation y (' + units + ')') + cbar = fig.colorbar(scatter_plot, ax=ax, fraction=0.05, pad=0.05) + if cmap_label is not None: + cbar.set_label(cmap_label) return fig