Start moving the example code to datasets which can fit on the git server

This commit is contained in:
Abe Levitan
2019-06-26 17:47:09 -04:00
parent c246a5ed96
commit f769f9e7ab
3 changed files with 5 additions and 26 deletions
+1 -2
View File
@@ -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
+4 -24
View File
@@ -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()