I'm using PyTorch Lightning to write a simple trainer, but when I try to run the trainer, for some reason, 9 out of 10 times it returns "CUDA error: device-side assert." Simply printing a newline before it somehow seems to make it work. Any ideas?
My code:
class Elementwise(nn.ModuleList):
"""
A simple network container.
Parameters are a list of modules.
Inputs are a 3d Tensor whose last dimension is the same length
as the list.
Outputs are the result of applying modules to inputs elementwise.
An optional merge parameter allows the outputs to be reduced to a
single Tensor.
"""
def __init__(self, merge=None, *args):
assert merge in [None, 'first', 'concat', 'sum', 'mlp']
self.merge = merge
super(Elementwise, self).__init__(*args)
def forward(self, inputs):
inputs_ = [feat.squeeze(1) for feat in inputs.split(1, dim=1)]
for i, j in enumerate(inputs_):
inp = torch.tensor(j).to(device).long()
inputs_[i] = inp
# this does not work
outputs = [f(x) for i, (f, x) in enumerate(zip(self, inputs_))]
if self.merge == 'first':
return outputs[0]
elif self.merge == 'concat' or self.merge == 'mlp':
return torch.cat(outputs, 1)
elif self.merge == 'sum':
return sum(outputs)
else:
return outputs
but somehow magically this works:
class Elementwise(nn.ModuleList):
"""
A simple network container.
Parameters are a list of modules.
Inputs are a 3d Tensor whose last dimension is the same length
as the list.
Outputs are the result of applying modules to inputs elementwise.
An optional merge parameter allows the outputs to be reduced to a
single Tensor.
"""
def __init__(self, merge=None, *args):
assert merge in [None, 'first', 'concat', 'sum', 'mlp']
self.merge = merge
super(Elementwise, self).__init__(*args)
def forward(self, inputs):
inputs_ = [feat.squeeze(1) for feat in inputs.split(1, dim=1)]
for i, j in enumerate(inputs_):
inp = torch.tensor(j).to(device).long()
inputs_[i] = inp
print("")
outputs = [f(x) for i, (f, x) in enumerate(zip(self, inputs_))]
if self.merge == 'first':
return outputs[0]
elif self.merge == 'concat' or self.merge == 'mlp':
return torch.cat(outputs, 1)
elif self.merge == 'sum':
return sum(outputs)
else:
return outputs
Any idea as to how this error gets fixed by simply printing to output?
Edit: This error is only raised when using PyTorch Lightning for training abstraction, using plain PyTorch makes it work fine.