diff --git a/CDTools/models/__init__.py b/CDTools/models/__init__.py index 096a453..fc53ef4 100644 --- a/CDTools/models/__init__.py +++ b/CDTools/models/__init__.py @@ -236,6 +236,7 @@ class CDIModel(t.nn.Module): Args: dataset (CDataset) : A dataset containing the simulated diffraction patterns to compare agains """ + fig, axes = plt.subplots(1,3,figsize=(12,5.3)) fig.tight_layout(rect=[0.02, 0.09, 0.98, 0.96]) axslider = plt.axes([0.15,0.06,0.75,0.03]) @@ -324,8 +325,6 @@ class CDIModel(t.nn.Module): fig.canvas.mpl_connect('scroll_event',on_action) update(0) - - pass diff --git a/examples/example_data/AuBalls_700ms_30nmStep_3_6SS_filter.cxi b/examples/example_data/AuBalls_700ms_30nmStep_3_6SS_filter.cxi new file mode 100644 index 0000000..01c5645 Binary files /dev/null and b/examples/example_data/AuBalls_700ms_30nmStep_3_6SS_filter.cxi differ diff --git a/examples/gold_ball_ptycho.py b/examples/gold_ball_ptycho.py index 8343bfc..3b0ae73 100644 --- a/examples/gold_ball_ptycho.py +++ b/examples/gold_ball_ptycho.py @@ -1,44 +1,25 @@ from __future__ import division, print_function, absolute_import import CDTools -from CDTools.tools.plotting import * -from CDTools.tools.cmath import * -from CDTools.tools import interactions import h5py import numpy as np -from matplotlib import pyplot as plt -import torch as t import pickle +from matplotlib import pyplot as plt -filename = '../../../Downloads/AuBalls_700ms_30nmStep_3_3SS_filter.cxi' -#filename = '/media/Data Bank/CSX_3_19/Processed_CXIs/115195_p.cxi' +filename = 'example_data/AuBalls_700ms_30nmStep_3_6SS_filter.cxi' with h5py.File(filename,'r') as f: dataset = CDTools.datasets.Ptycho_2D_Dataset.from_cxi(f) - #darks = np.array(f['entry_1/instrument_1/detector_1/data_dark']) - -#old_patterns = dataset.patterns.clone() -#dataset.patterns -= t.tensor(np.nanmean(darks,axis=0)) -#dataset.patterns = t.clamp(dataset.patterns,min=0) model = CDTools.models.FancyPtycho.from_dataset(dataset,n_modes=3,randomize_ang=0.1*np.pi) -#dataset.patterns = old_patterns + # default is CPU with 32-bit floats model.to(device='cuda') dataset.get_as(device='cuda') -#model.translation_offsets.requires_grad = False -for i, loss in enumerate(model.Adam_optimize(30, dataset, batch_size=100)): - model.inspect(dataset) - print(i,loss) - -for i, loss in enumerate(model.Adam_optimize(30, dataset, batch_size=100, lr=0.001)): - model.inspect(dataset) - print(i,loss) - -for i, loss in enumerate(model.Adam_optimize(30, dataset, batch_size=100, lr=0.0001)): +for i, loss in enumerate(model.Adam_optimize(3, dataset, batch_size=100)): model.inspect(dataset) print(i,loss) @@ -49,4 +30,3 @@ for i, loss in enumerate(model.Adam_optimize(30, dataset, batch_size=100, lr=0.0 model.inspect(dataset) model.compare(dataset) plt.show() -exit()