mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
apply PEP8
This commit is contained in:
@@ -4,6 +4,7 @@ Module to describe gradients
|
||||
|
||||
from torch import nn
|
||||
|
||||
|
||||
class GradInformation(nn.Module):
|
||||
|
||||
def grad_norm(self, norm_type):
|
||||
@@ -17,11 +18,10 @@ class GradInformation(nn.Module):
|
||||
norm = param_norm ** (1 / norm_type)
|
||||
|
||||
results['grad_{}_norm_{}'.format(norm_type, i)] = round(norm.data.cpu().numpy().flatten()[0], 3)
|
||||
except Exception as e:
|
||||
except Exception:
|
||||
# this param had no grad
|
||||
pass
|
||||
|
||||
total_norm = total_norm ** (1. / norm_type)
|
||||
results['grad_{}_norm_total'.format(norm_type)] = round(total_norm.data.cpu().numpy().flatten()[0], 3)
|
||||
return results
|
||||
|
||||
|
||||
@@ -43,4 +43,3 @@ class ModelHooks(torch.nn.Module):
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
@@ -94,7 +94,7 @@ class ModelSummary(object):
|
||||
mods = list(self.model.modules())
|
||||
sizes = []
|
||||
|
||||
for i in range(1,len(mods)):
|
||||
for i in range(1, len(mods)):
|
||||
m = mods[i]
|
||||
p = list(m.parameters())
|
||||
modsz = []
|
||||
@@ -127,7 +127,7 @@ class ModelSummary(object):
|
||||
if self.model.example_input_array is not None:
|
||||
cols.extend(['In_sizes', 'Out_sizes'])
|
||||
|
||||
df = pd.DataFrame(np.zeros( (len(self.layer_names), len(cols))))
|
||||
df = pd.DataFrame(np.zeros((len(self.layer_names), len(cols))))
|
||||
df.columns = cols
|
||||
|
||||
df['Name'] = self.layer_names
|
||||
@@ -152,16 +152,16 @@ class ModelSummary(object):
|
||||
self.make_summary()
|
||||
|
||||
|
||||
def print_mem_stack(): # pragma: no cover
|
||||
def print_mem_stack(): # pragma: no cover
|
||||
for obj in gc.get_objects():
|
||||
try:
|
||||
if torch.is_tensor(obj) or (hasattr(obj, 'data') and torch.is_tensor(obj.data)):
|
||||
print(type(obj), obj.size())
|
||||
except Exception as e:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def count_mem_items(): # pragma: no cover
|
||||
def count_mem_items(): # pragma: no cover
|
||||
nb_params = 0
|
||||
nb_tensors = 0
|
||||
for obj in gc.get_objects():
|
||||
@@ -172,7 +172,7 @@ def count_mem_items(): # pragma: no cover
|
||||
nb_params += 1
|
||||
else:
|
||||
nb_tensors += 1
|
||||
except Exception as e:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return nb_params, nb_tensors
|
||||
|
||||
@@ -129,6 +129,3 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
|
||||
def unfreeze(self):
|
||||
for param in self.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user