mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d98b9f2f93 | ||
|
|
f6416f737d |
@@ -795,6 +795,11 @@ class Trainer(TrainerIO):
|
||||
for optimizer in self.optimizers:
|
||||
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
|
||||
optimizer.zero_grad()
|
||||
|
||||
|
||||
@@ -22,3 +22,17 @@ class ModelHooks(torch.nn.Module):
|
||||
def on_tng_metrics(self, metrics):
|
||||
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/
|
||||
setup(
|
||||
name="pytorch-lightning",
|
||||
version='0.3.2',
|
||||
version='0.3.3',
|
||||
description="The Keras for ML researchers using PyTorch",
|
||||
author="William Falcon",
|
||||
author_email="waf2107@columbia.edu",
|
||||
|
||||
Reference in New Issue
Block a user