mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
added lbfgs support
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user