added amp level option

This commit is contained in:
William Falcon
2019-05-16 15:45:56 -04:00
parent 92f9b3e062
commit 9d19ab5850
2 changed files with 5 additions and 1 deletions
+3 -1
View File
@@ -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
+2
View File
@@ -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')