diff --git a/examples/demo.py b/examples/demo.py index 5d7bf0a..a5fa260 100644 --- a/examples/demo.py +++ b/examples/demo.py @@ -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() diff --git a/examples/multi_gpu_simple_ptycho.py b/examples/multi_gpu_simple_ptycho.py index ad7e646..94990f2 100644 --- a/examples/multi_gpu_simple_ptycho.py +++ b/examples/multi_gpu_simple_ptycho.py @@ -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__': diff --git a/examples/simple_ptycho.py b/examples/simple_ptycho.py index 2140f66..04f4218 100644 --- a/examples/simple_ptycho.py +++ b/examples/simple_ptycho.py @@ -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 diff --git a/src/cdtools/models/simple_ptycho.py b/src/cdtools/models/simple_ptycho.py index c4d85be..79c0df4 100644 --- a/src/cdtools/models/simple_ptycho.py +++ b/src/cdtools/models/simple_ptycho.py @@ -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) `_. diff --git a/src/cdtools/tools/interactions/interactions.py b/src/cdtools/tools/interactions/interactions.py index 0b56d4f..a092bf3 100644 --- a/src/cdtools/tools/interactions/interactions.py +++ b/src/cdtools/tools/interactions/interactions.py @@ -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)