added lbfgs support

This commit is contained in:
William Falcon
2019-10-05 10:47:18 -04:00
parent bf09060fef
commit 967957e55c
4 changed files with 62 additions and 23 deletions
+6 -2
View File
@@ -92,16 +92,20 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
"""
raise NotImplementedError
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):
"""
Do something instead of the standard optimizer behavior
:param epoch_nb:
:param batch_nb:
:param optimizer:
:param optimizer_i:
:param second_order_closure: closure for second order methods
:return:
"""
optimizer.step()
if isinstance(optimizer, torch.optim.LBFGS):
optimizer.step(second_order_closure)
else:
optimizer.step()
# clear gradients
optimizer.zero_grad()
@@ -119,7 +119,10 @@ class LightningTestModelBase(LightningModule):
:return: list of optimizers
"""
# try no scheduler for this model (testing purposes)
optimizer = optim.Adam(self.parameters(), lr=self.hparams.learning_rate)
if self.hparams.optimizer == 'lbfgs':
optimizer = optim.LBFGS(self.parameters(), lr=self.hparams.learning_rate)
else:
optimizer = optim.Adam(self.parameters(), lr=self.hparams.learning_rate)
# test returning only 1 list instead of 2
return optimizer
+27 -18
View File
@@ -275,6 +275,8 @@ class Trainer(TrainerIO):
raise ModuleNotFoundError(msg)
def __configure_accumulated_gradients(self, accumulate_grad_batches):
self.accumulate_grad_batches = None
if isinstance(accumulate_grad_batches, dict):
self.accumulation_scheduler = GradientAccumulationScheduler(accumulate_grad_batches)
elif isinstance(accumulate_grad_batches, int):
@@ -1267,27 +1269,34 @@ class Trainer(TrainerIO):
# call training_step once per optimizer
for opt_idx, optimizer in enumerate(self.optimizers):
# forward pass
loss, model_specific_tqdm_metrics = self.__training_forward(batch, batch_nb, opt_idx)
# wrap the forward step in a closure so second order methods work
def optimizer_closure():
# forward pass
closure_loss, model_specific_tqdm_metrics = self.__training_forward(batch, batch_nb, opt_idx)
# track metrics
self.__add_tqdm_metrics(model_specific_tqdm_metrics)
# track metrics
self.__add_tqdm_metrics(model_specific_tqdm_metrics)
# accumulate loss
# (if accumulate_grad_batches = 1 no effect)
loss = loss / self.accumulate_grad_batches
# accumulate loss
# (if accumulate_grad_batches = 1 no effect)
closure_loss = closure_loss / self.accumulate_grad_batches
# backward pass
if self.use_amp:
with amp.scale_loss(loss, optimizer) as scaled_loss:
scaled_loss.backward()
else:
loss.backward()
# backward pass
if self.use_amp:
with amp.scale_loss(closure_loss, optimizer) as scaled_loss:
scaled_loss.backward()
else:
closure_loss.backward()
# insert after step hook
if self.__is_function_implemented('on_after_backward'):
model_ref = self.__get_model()
model_ref.on_after_backward()
# insert after step hook
if self.__is_function_implemented('on_after_backward'):
model_ref = self.__get_model()
model_ref.on_after_backward()
return closure_loss
# calculate loss
loss = optimizer_closure()
# nan grads
if self.print_nan_grads:
@@ -1311,7 +1320,7 @@ class Trainer(TrainerIO):
# calls .step(), .zero_grad()
# override function to modify this behavior
model = self.__get_model()
model.optimizer_step(self.current_epoch, batch_nb, optimizer, opt_idx)
model.optimizer_step(self.current_epoch, batch_nb, optimizer, opt_idx, optimizer_closure)
# calculate running loss for display
self.running_loss.append(self.batch_loss_value)
+25 -2
View File
@@ -44,7 +44,6 @@ def test_default_logger_callbacks_cpu_model():
Test each of the trainer options
:return:
"""
trainer_options = dict(
max_nb_epochs=1,
gradient_clip_val=1.0,
@@ -63,6 +62,28 @@ def test_default_logger_callbacks_cpu_model():
model.unfreeze()
def test_lbfgs_cpu_model():
"""
Test each of the trainer options
:return:
"""
trainer_options = dict(
max_nb_epochs=1,
gradient_clip_val=1.0,
overfit_pct=0.20,
print_nan_grads=True,
show_progress_bar=False,
train_percent_check=0.01,
val_percent_check=0.01
)
model, hparams = get_model(use_test_model=True, lbfgs=True)
run_model_test_no_loggers(trainer_options, model, hparams, on_gpu=False)
# test freeze on cpu
model.freeze()
model.unfreeze()
def test_multi_gpu_model_ddp2():
"""
Make sure DDP2 works
@@ -1447,9 +1468,11 @@ def get_hparams(continue_training=False, hpc_exp_number=0):
return hparams
def get_model(use_test_model=False):
def get_model(use_test_model=False, lbfgs=False):
# set up model with these hyperparams
hparams = get_hparams()
if lbfgs:
hparams.optimizer = 'lbfgs'
if use_test_model:
model = LightningTestModel(hparams)