From 9d19ab5850ebbf8f0053348aa78a11a09a00c25a Mon Sep 17 00:00:00 2001 From: William Falcon Date: Thu, 16 May 2019 15:45:56 -0400 Subject: [PATCH] added amp level option --- pytorch_lightning/models/trainer.py | 4 +++- pytorch_lightning/utils/arg_parse.py | 2 ++ 2 files changed, 5 insertions(+), 1 deletion(-) diff --git a/pytorch_lightning/models/trainer.py b/pytorch_lightning/models/trainer.py index 120688e1..7d94f636 100644 --- a/pytorch_lightning/models/trainer.py +++ b/pytorch_lightning/models/trainer.py @@ -33,6 +33,7 @@ class Trainer(TrainerIO): log_save_interval=1, add_log_row_interval=1, lr_scheduler_milestones=None, use_amp=False, + amp_level='O2', nb_sanity_val_steps=5): # Transfer params @@ -58,6 +59,7 @@ class Trainer(TrainerIO): self.nb_sanity_val_steps = nb_sanity_val_steps self.lr_scheduler_milestones = [] if lr_scheduler_milestones is None else [int(x.strip()) for x in lr_scheduler_milestones.split(',')] self.lr_schedulers = [] + self.amp_level = amp_level # training state self.optimizers = None @@ -222,7 +224,7 @@ class Trainer(TrainerIO): if self.use_amp: # An example self.model, optimizer = amp.initialize( - self.model, self.optimizers[0], opt_level="O2", + self.model, self.optimizers[0], opt_level=self.amp_level, ) self.optimizers[0] = optimizer model.trainer = self diff --git a/pytorch_lightning/utils/arg_parse.py b/pytorch_lightning/utils/arg_parse.py index fc971395..e466301c 100644 --- a/pytorch_lightning/utils/arg_parse.py +++ b/pytorch_lightning/utils/arg_parse.py @@ -51,6 +51,8 @@ def add_default_args(parser, root_dir, rand_seed=None, possible_model_names=None parser.add_argument('--disable_cuda', dest='disable_cuda', action='store_true') parser.add_argument('--default_tensor_type', default='torch.cuda.FloatTensor', type=str) parser.add_argument('--use_amp', dest='use_amp', action='store_true') + parser.add_argument('--amp_level', dest='O2', action='store_true') + # run on hpc parser.add_argument('--on_cluster', dest='on_cluster', action='store_true')