mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-18 08:42:08 +02:00
Start moving the example code to datasets which can fit on the git server
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
|
||||
|
||||
Binary file not shown.
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user