mirror of
https://github.com/cdtools-developers/cdtools.git
synced 2026-09-09 21:12:42 +02:00
some horrible state but I need to stop developing on perlmutter
This commit is contained in:
@@ -35,6 +35,8 @@ def demo_basic():
|
||||
print(i)
|
||||
optimizer.zero_grad()
|
||||
outputs = ddp_model(torch.randn(20, 10))
|
||||
mat = torch.rand([3,3], dtype=torch.complex64)
|
||||
inv_mat = torch.inverse(mat)
|
||||
labels = torch.randn(20, 5).to(device_id)
|
||||
loss_fn(outputs, labels).backward()
|
||||
optimizer.step()
|
||||
|
||||
@@ -34,7 +34,7 @@ def demo():
|
||||
sampler = torchdata.distributed.DistributedSampler(dataset)
|
||||
# Make a dataloader
|
||||
#data_loader = torchdata.DataLoader(dataset, batch_size=20,
|
||||
# shuffle=True)
|
||||
# shuffle=True, num_workers=0)
|
||||
data_loader = torchdata.DataLoader(dataset, batch_size=20,
|
||||
shuffle=False, sampler=sampler)
|
||||
|
||||
@@ -42,26 +42,29 @@ def demo():
|
||||
dataset.get_as(device=device)
|
||||
|
||||
# Define the optimizer
|
||||
optimizer = MyAdam(model.parameters(), lr=0.01)
|
||||
#optimizer = MyAdam(model.parameters(), lr=0.01)
|
||||
optimizer = t.optim.Adam(model.parameters(), lr=0.01)
|
||||
|
||||
normalization=0
|
||||
for inputs, patterns in data_loader:
|
||||
for inputs, patterns in dataset:#data_loader:
|
||||
normalization += t.sum(patterns)
|
||||
|
||||
dist.all_reduce(normalization)
|
||||
|
||||
for it in range(100):
|
||||
for it in range(10):
|
||||
print(it)
|
||||
sampler.set_epoch(it)
|
||||
loss = 0
|
||||
N = 0
|
||||
t0 = time.time()
|
||||
for inputs, patterns in data_loader:
|
||||
#print(inputs)
|
||||
N += 1
|
||||
def closure():
|
||||
optimizer.zero_grad()
|
||||
|
||||
sim_patterns = model.forward(*inputs)
|
||||
|
||||
|
||||
if hasattr(model, 'mask'):
|
||||
loss = model.loss(patterns,sim_patterns, mask=model.mask)
|
||||
else:
|
||||
@@ -70,12 +73,11 @@ def demo():
|
||||
loss.backward()
|
||||
return loss.detach()
|
||||
|
||||
loss += optimizer.step(closure).detach()#.cpu().numpy()
|
||||
loss += optimizer.step(closure)#.cpu().numpy()
|
||||
#dist.all_reduce(loss)
|
||||
#loss = dist.all_reduce(loss)
|
||||
#print(it, 'time', time.time()-t0)
|
||||
#print(loss / normalization)
|
||||
|
||||
return loss#.cpu().numpy()
|
||||
|
||||
if __name__=='__main__':
|
||||
|
||||
@@ -18,7 +18,7 @@ model = cdtools.models.SimplePtycho.from_dataset(dataset)
|
||||
# def __getattr__(self, name):
|
||||
# return getattr(self.module, name)
|
||||
|
||||
model = t.nn.DataParallel(model, device_ids=[0,1,2,3])
|
||||
#model = t.nn.DataParallel(model, device_ids=[0,1,2,3])
|
||||
|
||||
|
||||
# Make a dataloader
|
||||
|
||||
@@ -10,6 +10,11 @@ import numpy as np
|
||||
|
||||
__all__ = ['SimplePtycho']
|
||||
|
||||
class complexWrapper(t.Tensor):
|
||||
def __new__(base_tensor):
|
||||
return t.view_as_complex(base_tensor)
|
||||
|
||||
|
||||
class SimplePtycho(CDIModel):
|
||||
"""A simple ptychography model for exploring ideas and extensions
|
||||
|
||||
@@ -44,11 +49,14 @@ class SimplePtycho(CDIModel):
|
||||
# object
|
||||
self.register_buffer('probe_norm', t.max(t.abs(probe_guess)))
|
||||
|
||||
self.probe_data = t.nn.Parameter(t.view_as_real(probe_guess / self.probe_norm))
|
||||
self.probe_data = complexParameter(probe_guess/self.probe_norm)
|
||||
#self.probe_data = t.nn.Parameter(t.view_as_real(probe_guess / self.probe_norm))
|
||||
self.obj_data = t.nn.Parameter(t.view_as_real(obj_guess))
|
||||
|
||||
@property
|
||||
def probe(self):
|
||||
probe = t.view_as_complex(self.probe_data)
|
||||
|
||||
return t.view_as_complex(self.probe_data)
|
||||
|
||||
@property
|
||||
@@ -106,7 +114,6 @@ class SimplePtycho(CDIModel):
|
||||
|
||||
|
||||
def interaction(self, index, translations):
|
||||
|
||||
pix_trans = tools.interactions.translations_to_pixel(self.probe_basis,
|
||||
translations,
|
||||
surface_normal=self.surface_normal)
|
||||
@@ -181,7 +188,7 @@ class SimplePtycho(CDIModel):
|
||||
('Object Phase',
|
||||
lambda self, fig: p.plot_phase(self.obj, fig=fig, basis=self.probe_basis))
|
||||
]
|
||||
|
||||
|
||||
|
||||
def ePIE(self, iterations, dataset, beta = 1.0):
|
||||
"""Runs an ePIE reconstruction as described in `Maiden et al. (2017) <https://www.osapublishing.org/optica/abstract.cfm?uri=optica-4-7-736>`_.
|
||||
|
||||
@@ -43,7 +43,6 @@ def translations_to_pixel(basis, translations, surface_normal=t.Tensor([0.,0.,1.
|
||||
pixel_translations : torch.Tensor
|
||||
A Jx2 stack of translations in internal (i,j) pixel-space, or a single translation
|
||||
"""
|
||||
|
||||
projection_1 = t.as_tensor(np.array([[1,0,0],
|
||||
[0,1,0],
|
||||
[0,0,0]]),
|
||||
@@ -53,7 +52,8 @@ def translations_to_pixel(basis, translations, surface_normal=t.Tensor([0.,0.,1.
|
||||
projection_2[2] = t.as_tensor(-surface_normal / surface_normal[2],
|
||||
dtype=projection_2.dtype,
|
||||
device=projection_2.device)
|
||||
projection_2 = t.inverse(projection_2)
|
||||
projection_2 = t.linalg.inv(projection_2)
|
||||
|
||||
|
||||
basis_vectors_inv = t.pinverse(basis).to(device=translations.device,
|
||||
dtype=translations.dtype)
|
||||
|
||||
Reference in New Issue
Block a user