diff --git a/CDTools/models/base.py b/CDTools/models/base.py index 9d8e9eb..bca2e3a 100644 --- a/CDTools/models/base.py +++ b/CDTools/models/base.py @@ -451,16 +451,17 @@ class CDIModel(t.nn.Module): """ - print('base models inspect: checking the object') - a = self.obj.detach() - def saveobj(a, filename): - a = np.abs(a) - plt.imshow(a) - plt.savefig(filename) - f = ['base_a.png', 'base_b.png', 'base_c.png', 'base_d.png'] - comp = [a[i, j, :, :] for i, j in zip([0, 0, 1, 1], [0, 1, 0, 1])] - for i in range(4): - saveobj(comp[i], f[i]) + #print('base models inspect: checking the object') + #a = self.obj.detach() + #def saveobj(a, filename): + # a = np.abs(a) + # plt.imshow(a) + # plt.savefig(filename) + #f = ['base_a.png', 'base_b.png', 'base_c.png', 'base_d.png'] + #comp = [a[i, j, :, :] for i, j in zip([0, 0, 1, 1], [0, 1, 0, 1])] + #for i in range(4): + # saveobj(comp[i], f[i]) + first_update = False if update and hasattr(self, 'figs') and self.figs: figs = self.figs diff --git a/CDTools/models/polarized_fancy_ptycho.py b/CDTools/models/polarized_fancy_ptycho.py index 93c37f7..99dff60 100644 --- a/CDTools/models/polarized_fancy_ptycho.py +++ b/CDTools/models/polarized_fancy_ptycho.py @@ -63,7 +63,23 @@ class PolarizedFancyPtycho(FancyPtycho): def from_dataset(cls, dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, dm_rank=None, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, restrict_obj=-1, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um', left_polarized=True): # When using this method, remember to pass through the inputs - model = FancyPtycho.from_dataset(dataset, probe_size=None, randomize_ang=0, padding=0, n_modes=1, dm_rank=None, translation_scale = 1, saturation=None, probe_support_radius=None, propagation_distance=None, scattering_mode=None, oversampling=1, auto_center=False, opt_for_fft=False, loss='amplitude mse', units='um') + model = FancyPtycho.from_dataset( + dataset, + probe_size=probe_size, + randomize_ang=randomize_ang, + padding=padding, + n_modes=n_modes, + dm_rank=dm_rank, + translation_scale=translation_scale, + saturation=saturation, + probe_support_radius=probe_support_radius, + propagation_distance=propagation_distance, + scattering_mode=scattering_mode, + oversampling=oversampling, + auto_center=auto_center, + opt_for_fft=opt_for_fft, + loss=loss, + units=units) # Mutate the class to its subclass @@ -79,9 +95,9 @@ class PolarizedFancyPtycho(FancyPtycho): probe_max = t.max(t.abs(probe)) probe_stack = [0.01 * probe_max * t.rand(probe.shape, dtype=probe.dtype) for i in range(n_modes - 1)] probe = t.stack([probe, ] + probe_stack) - print('probe', type(probe), probe.shape) + #print('probe', type(probe), probe.shape) model.probe.data = probe - print(model.probe.shape) + #print(model.probe.shape) # obj = t.stack((model.obj.data, model.obj.data), dim=-3) # model.obj.data = t.stack((obj.data, obj.data), dim=-4) # obj = t.exp(1j * randomize_ang * (t.rand(obj_size)-0.5)) @@ -90,14 +106,14 @@ class PolarizedFancyPtycho(FancyPtycho): # initialization (e.g. ((obj,0*obj),(0*obj,obj)) obj = t.stack((obj, obj), dim=-3) obj = t.stack((obj, obj), dim=-4) - print('object', type(obj), obj.shape) + #print('object', type(obj), obj.shape) model.obj.data = obj - print('polarized fancy ptycho from datset obj') + #print('polarized fancy ptycho from datset obj') a = obj.detach() - plt.imshow(np.real(a[0, 0, :, :])) - plt.show() - plt.imshow(np.real(a[0, 1, :, :])) - plt.show() + #plt.imshow(np.real(a[0, 0, :, :])) + #plt.figure() + #plt.imshow(np.real(a[0, 1, :, :])) + #plt.show() # tensor vs tensor.data return model diff --git a/CDTools/tools/plotting/plotting.py b/CDTools/tools/plotting/plotting.py index f1451b2..a712b65 100644 --- a/CDTools/tools/plotting/plotting.py +++ b/CDTools/tools/plotting/plotting.py @@ -17,11 +17,11 @@ from matplotlib import ticker, patheffects __all__ = ['colorize', 'plot_amplitude', 'plot_phase', 'plot_colorized', 'plot_translations', 'get_units_factor', 'plot_nanomap', 'plot_real', 'plot_imag', - 'plot_nanomap_with_images', - 'polarized_plot_component_amplitudes', - 'polarized_plot_phase_ret', - 'polarized_plot_global_phases', - 'polarized_plot_ellipses'] + 'plot_nanomap_with_images']#, + #'polarized_plot_component_amplitudes', + #'polarized_plot_phase_ret', + #'polarized_plot_global_phases', + #'polarized_plot_ellipses'] def colorize(z):