some horrible state but I need to stop developing on perlmutter

This commit is contained in:
abe
2022-10-17 16:32:59 -07:00
parent cb954a6067
commit e13b519e22
5 changed files with 24 additions and 13 deletions
+2
View File
@@ -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()
+9 -7
View File
@@ -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__':
+1 -1
View File
@@ -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 -3
View File
@@ -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)