mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-12 12:40:20 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d98b9f2f93 | ||
|
|
f6416f737d |
@@ -795,6 +795,11 @@ class Trainer(TrainerIO):
|
|||||||
for optimizer in self.optimizers:
|
for optimizer in self.optimizers:
|
||||||
optimizer.step()
|
optimizer.step()
|
||||||
|
|
||||||
|
# insert after step hook
|
||||||
|
if self.__is_function_implemented('on_before_zero_grad'):
|
||||||
|
model_ref = self.__get_model()
|
||||||
|
response = model_ref.on_before_zero_grad(optimizer)
|
||||||
|
|
||||||
# clear gradients
|
# clear gradients
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
|
|
||||||
|
|||||||
@@ -22,3 +22,17 @@ class ModelHooks(torch.nn.Module):
|
|||||||
def on_tng_metrics(self, metrics):
|
def on_tng_metrics(self, metrics):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def on_before_zero_grad(self, optimizer):
|
||||||
|
"""
|
||||||
|
Called after optimizer.step() and before optimizer.zero_grad()
|
||||||
|
|
||||||
|
for optimizer in optimizers:
|
||||||
|
optimizer.step()
|
||||||
|
model.on_before_zero_grad(optimizer) # < ---- called here
|
||||||
|
optimizer.zero_grad
|
||||||
|
|
||||||
|
:param optimizer:
|
||||||
|
:return:
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from setuptools import setup, find_packages
|
|||||||
# http://blog.ionelmc.ro/2014/05/25/python-packaging/
|
# http://blog.ionelmc.ro/2014/05/25/python-packaging/
|
||||||
setup(
|
setup(
|
||||||
name="pytorch-lightning",
|
name="pytorch-lightning",
|
||||||
version='0.3.2',
|
version='0.3.3',
|
||||||
description="The Keras for ML researchers using PyTorch",
|
description="The Keras for ML researchers using PyTorch",
|
||||||
author="William Falcon",
|
author="William Falcon",
|
||||||
author_email="waf2107@columbia.edu",
|
author_email="waf2107@columbia.edu",
|
||||||
|
|||||||
Reference in New Issue
Block a user