From 4e1c90d892b95dbe71a9140095f8306fd69a6eaf Mon Sep 17 00:00:00 2001 From: William Falcon Date: Sat, 5 Oct 2019 20:50:40 -0400 Subject: [PATCH] cleaning up docs --- docs/Trainer/hooks.md | 6 +++--- tests/test_models.py | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/docs/Trainer/hooks.md b/docs/Trainer/hooks.md index 726b5f5a..6fd8df3d 100644 --- a/docs/Trainer/hooks.md +++ b/docs/Trainer/hooks.md @@ -65,12 +65,12 @@ You can override this method to adjust how you do the optimizer step for each op Called once per optimizer ```python # DEFAULT -def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i): +def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i, second_order_closure=None): optimizer.step() optimizer.zero_grad() # Alternating schedule for optimizer steps (ie: GANs) -def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i): +def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i, second_order_closure=None): # update generator opt every 2 steps if optimizer_i == 0: if batch_nb % 2 == 0 : @@ -91,7 +91,7 @@ This step allows you to do a lot of non-standard training tricks such as learnin ```python # learning rate warm-up -def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i): +def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i, second_order_closure=None): # warm up lr if self.trainer.global_step < 500: lr_scale = min(1., float(self.trainer.global_step + 1) / 500.) diff --git a/tests/test_models.py b/tests/test_models.py index ed38dc5e..581469cb 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -446,7 +446,7 @@ def test_gradient_accumulation_scheduling(): assert Trainer(accumulate_grad_batches={1: 2.5, 3: 5}) # test optimizer call freq matches scheduler - def optimizer_step(self, epoch_nb, batch_nb, optimizer, optimizer_i): + def optimizer_step(self, epoch_nb, batch_nb, optimizer, optimizer_i, second_order_closure=None): # only test the first 12 batches in epoch if batch_nb < 12: if epoch_nb == 0: