mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
allow user to control optimizer step for every optimizer
* added custom hook for user defined optimizer step * refactored to allow multiple optimizers different training_step * refactored to allow multiple optimizers different training_step * refactored to allow multiple optimizers different training_step * refactored to allow multiple optimizers different training_step * refactored to allow multiple optimizers different training_step * pep8
This commit is contained in:
@@ -66,6 +66,20 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def optimizer_step(self, epoch_nb, batch_nb, optimizer, optimizer_i):
|
||||
"""
|
||||
Do something instead of the standard optimizer behavior
|
||||
:param epoch_nb:
|
||||
:param batch_nb:
|
||||
:param optimizer:
|
||||
:param optimizer_i:
|
||||
:return:
|
||||
"""
|
||||
optimizer.step()
|
||||
|
||||
# clear gradients
|
||||
optimizer.zero_grad()
|
||||
|
||||
@data_loader
|
||||
def tng_dataloader(self):
|
||||
"""
|
||||
|
||||
Reference in New Issue
Block a user