diff --git a/examples/simple_ptycho.py b/examples/simple_ptycho.py index 217d96b..571f24f 100644 --- a/examples/simple_ptycho.py +++ b/examples/simple_ptycho.py @@ -15,21 +15,25 @@ 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) +model = cdtools.models.SimplePtycho.from_dataset(dataset, panel_plot_mode=True) # We move the model to the GPU device = 'cuda' model.to(device=device) dataset.get_as(device=device) +model.inspect(dataset) +print('hi') # We run the reconstruction -for loss in model.Adam_optimize(100, dataset, batch_size=10): +for loss in model.Adam_optimize(30, dataset, batch_size=10): # We print a quick report of the optimization status print(model.report()) # And liveplot the updates to the model as they happen - model.inspect(dataset) + if model.epoch % 10 == 0: + model.inspect(dataset) + # We study the results -model.inspect(dataset) +model.inspect(dataset, replot_all=True) model.compare(dataset) plt.show() diff --git a/src/cdtools/models/base.py b/src/cdtools/models/base.py index f2823b7..2dba377 100644 --- a/src/cdtools/models/base.py +++ b/src/cdtools/models/base.py @@ -30,6 +30,7 @@ loss import torch as t from torch.utils import data as torchdata +import matplotlib from matplotlib import pyplot as plt from matplotlib.widgets import Slider from matplotlib import ticker @@ -66,6 +67,7 @@ class CDIModel(t.nn.Module): self.epoch = 0 self.panel_plot_mode = panel_plot_mode self.plot_level = plot_level + self.has_inspect_been_called = False def from_dataset(self, dataset): raise NotImplementedError() @@ -588,29 +590,56 @@ class CDIModel(t.nn.Module): plot_panel_list = getattr(self, 'plot_panel_list', None) or [] plot_list = getattr(self, 'plot_list', None) or [] - if self.panel_plot_mode and plot_panel_list: - self._inspect_panel(dataset=dataset, replot_all=replot_all) + if self.panel_plot_mode: + # First we plot all the panels + panel_figs = self._inspect_panel( + plot_panel_list, dataset=dataset, replot_all=replot_all) + # And then we plot all the individual figures + individual_figs = self._inspect_individual_figures( + plot_list, dataset=dataset, replot_all=replot_all + ) + self.figs = panel_figs + individual_figs else: - # Flatten plot_panel_list, assigning each subplot the panel's plot_level, - # then prepend to plot_list + # If not in panel plot mode, we first flatten the figures + # from the panels flat = [] for panel in plot_panel_list: panel_level = panel.get('plot_level', 0) for plot in panel['plots']: + # We add the plot level from the larger panel flat.append({**plot, 'plot_level': panel_level}) + all_plots = flat + list(plot_list) - if not hasattr(self, '_flat_fig_map'): - self._flat_fig_map = {} + # We make sure to keep a reference to the open figs around + self.figs = self._inspect_individual_figures( + all_plots, dataset=dataset, replot_all=replot_all) - self.figs = self._do_inspect(all_plots, self._flat_fig_map, - dataset=dataset, - replot_all=replot_all) + if not self.has_inspect_been_called or replot_all: + # Somehow, this is needed for new figures to appear + if self._is_backend_interactive(): + plt.pause(0.05 * len(self.figs)) + for fig in self.figs: + fig.canvas.flush_events() + self.has_inspect_been_called = True - plt.pause(0.05) + + def _is_backend_interactive( + self + ): + backend = matplotlib.get_backend().lower() + interactive_bk = matplotlib.backends.backend_registry.list_builtin( + matplotlib.backends.BackendFilter.INTERACTIVE + ) + return backend in [b.lower() for b in interactive_bk] + - - def _do_inspect(self, plot_list, fig_map, dataset=None, replot_all=False): + def _inspect_individual_figures( + self, + plot_list, + dataset=None, + replot_all=False + ): """Core one-figure-per-plot rendering logic. fig_map is a dict {title: figure} owned by the caller and updated @@ -622,8 +651,6 @@ class CDIModel(t.nn.Module): Returns the list of figures that were rendered this call. """ - if not plot_list: - return [] rendered = [] @@ -642,45 +669,38 @@ class CDIModel(t.nn.Module): if not condition(self, dataset): continue - title = plot['title'] - fig = fig_map.get(title) + if self.has_inspect_been_called and \ + replot_all == False and \ + not plt.fignum_exists(plot['title']): + continue - if fig is not None and not plt.fignum_exists(fig.number): - # Figure was closed by the user - if replot_all: - fig = None - del fig_map[title] - else: - continue # leave it closed - - if fig is None: - fig = plt.figure(num=title) - fig._panel_label = title - fig_map[title] = fig + if not self.has_inspect_been_called: + fig = plt.figure(plot['title']) + else: + with plt.rc_context({'figure.raise_window': False}): + fig = plt.figure(plot['title']) try: plot['plot_func'](self, fig) - plt.title(title) + plt.title(plot['title']) except TypeError: if dataset is not None: try: plot['plot_func'](self, fig, dataset) - plt.title(title) + plt.title(plot['title']) except Exception: pass except Exception: pass rendered.append(fig) - try: - fig.canvas.draw_idle() - except Exception: - pass + if self._is_backend_interactive(): + plt.draw() return rendered - def _inspect_panel(self, dataset=None, replot_all=False): + def _inspect_panel(self, plot_panel_list, dataset=None, replot_all=False): """Multi-subplot panel rendering. Creates one figure per plot_panel_list entry, placing each subplot's @@ -688,22 +708,9 @@ class CDIModel(t.nn.Module): on subsequent calls unless replot_all=True. Standalone plot_list entries are then rendered via _do_inspect and appended to self.figs. """ - plot_panel_list = getattr(self, 'plot_panel_list', None) or [] - plot_list = getattr(self, 'plot_list', None) or [] - n_panels = len(plot_panel_list) - - # _panel_figs: list of figures (or None if never created / closed). - # _panel_axes: dict keyed by (panel_idx, row, col) → Axes. - # _standalone_fig_map: dict {title: figure} for standalone plot_list. - first_call = not hasattr(self, '_panel_figs') - if first_call: - self._panel_figs = [None] * n_panels - self._panel_axes = {} - self._standalone_fig_map = {} - - if not hasattr(self, '_standalone_fig_map'): - self._standalone_fig_map = {} + rendered = [] + for panel_idx, panel_def in enumerate(plot_panel_list): panel_level = panel_def.get('plot_level', 0) if panel_level > self.plot_level: @@ -713,30 +720,24 @@ class CDIModel(t.nn.Module): figsize = panel_def.get('figure_size', None) title = panel_def.get('title', '') - fig = self._panel_figs[panel_idx] - # Detect if a previously open figure was closed by the user. - if fig is not None and not plt.fignum_exists(fig.number): - self._panel_figs[panel_idx] = None - for k in [k for k in self._panel_axes if k[0] == panel_idx]: - del self._panel_axes[k] - fig = None + if self.has_inspect_been_called and \ + replot_all == False and \ + not plt.fignum_exists(panel_def['title']): + continue - if fig is None: - if not first_call and not replot_all: - continue # was closed; leave it closed - fig = plt.figure(num=title, figsize=figsize) - fig._panel_label = title - self._panel_figs[panel_idx] = fig + if not self.has_inspect_been_called: + fig = plt.figure(panel_def['title']) else: + with plt.rc_context({'figure.raise_window': False}): + fig = plt.figure(panel_def['title']) + # 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() - for k in [k for k in self._panel_axes if k[0] == panel_idx]: - del self._panel_axes[k] for plot in panel_def['plots']: condition = plot.get('condition', None) @@ -753,7 +754,6 @@ class CDIModel(t.nn.Module): ax_key = (panel_idx, row, col) ax = fig.add_subplot(nrows, ncols, position) - self._panel_axes[ax_key] = ax try: plot['plot_func'](self, ax) @@ -767,20 +767,12 @@ class CDIModel(t.nn.Module): pass except Exception: pass + rendered.append(fig) + + if self._is_backend_interactive(): + plt.draw() - try: - fig.canvas.draw_idle() - except Exception: - pass - - # Rebuild self.figs from open panel figures + rendered standalone figures. - panel_figs = [f for f in self._panel_figs if f is not None] - standalone_rendered = self._do_inspect( - list(plot_list), self._standalone_fig_map, - dataset=dataset, replot_all=replot_all, - ) - self.figs = panel_figs + standalone_rendered - + return rendered def save_figures(self, prefix='', extension='.pdf'): diff --git a/src/cdtools/models/simple_ptycho.py b/src/cdtools/models/simple_ptycho.py index bcc02d2..0460ecb 100644 --- a/src/cdtools/models/simple_ptycho.py +++ b/src/cdtools/models/simple_ptycho.py @@ -16,7 +16,7 @@ class SimplePtycho(CDIModel): obj_guess, min_translation = [0,0], panel_plot_mode=False, - plot_level=0, + plot_level=1, ): # We initialize the superclass @@ -49,7 +49,7 @@ class SimplePtycho(CDIModel): @classmethod def from_dataset(cls, dataset,panel_plot_mode=False, - plot_level=0, ): + plot_level=1, ): # We get the key geometry information from the dataset wavelength = dataset.wavelength @@ -142,7 +142,7 @@ class SimplePtycho(CDIModel): } ] }, - { + { # Title for window 'title' : 'Object Results', # (width, height) in inches diff --git a/src/cdtools/reconstructors/base.py b/src/cdtools/reconstructors/base.py index 4beab9f..a9698ee 100644 --- a/src/cdtools/reconstructors/base.py +++ b/src/cdtools/reconstructors/base.py @@ -364,12 +364,9 @@ class Reconstructor: None, ) if open_fig is not None: - try: - open_fig.canvas.start_event_loop(0.01) - except Exception: - time.sleep(0.01) - else: - time.sleep(0.01) + open_fig.canvas.flush_events() + # We need a low value for smooth figure responses + time.sleep(0.001) except KeyboardInterrupt as e: stop_event.set()