88import mxnet as mx
99from mxnet import gluon , autograd
1010from mxnet .gluon .data .vision import transforms
11+ from mxnet .contrib import amp
1112
1213import gluoncv
1314gluoncv .utils .check_version ('0.6.0' )
@@ -93,6 +94,11 @@ def parse_args():
9394 # synchronized Batch Normalization
9495 parser .add_argument ('--syncbn' , action = 'store_true' , default = False ,
9596 help = 'using Synchronized Cross-GPU BatchNorm' )
97+ # performance related
98+ parser .add_argument ('--amp' , action = 'store_true' ,
99+ help = 'Use MXNet AMP for mixed precision training.' )
100+ parser .add_argument ('--auto-layout' , action = 'store_true' ,
101+ help = 'Add layout optimization to AMP. Must be used in addition of `--amp`.' )
96102 # the parser
97103 args = parser .parse_args ()
98104
@@ -200,7 +206,11 @@ def __init__(self, args, logger):
200206 v .wd_mult = 0.0
201207
202208 self .optimizer = gluon .Trainer (self .net .module .collect_params (), 'sgd' ,
203- optimizer_params , kvstore = kv )
209+ optimizer_params , kvstore = (False if args .amp else None ))
210+
211+
212+ if args .amp :
213+ amp .init_trainer (trainer )
204214 # evaluation metrics
205215 self .metric = gluoncv .utils .metrics .SegmentationMetric (trainset .num_class )
206216
@@ -212,7 +222,11 @@ def training(self, epoch):
212222 outputs = self .net (data .astype (args .dtype , copy = False ))
213223 losses = self .criterion (outputs , target )
214224 mx .nd .waitall ()
215- autograd .backward (losses )
225+ if args .amp :
226+ with amp .scale_loss (losses , self .optimizer ) as scaled_losses :
227+ autograd .backward (scaled_losses )
228+ else :
229+ autograd .backward (losses )
216230 self .optimizer .step (self .args .batch_size )
217231 for loss in losses :
218232 train_loss += np .mean (loss .asnumpy ()) / len (losses )
@@ -252,7 +266,10 @@ def save_checkpoint(net, args, epoch, mIoU, is_best=False):
252266
253267if __name__ == "__main__" :
254268 args = parse_args ()
269+ assert not args .auto_layout or args .amp , "--auto-layout needs to be used with --amp"
255270
271+ if args .amp :
272+ amp .init (layout_optimization = args .auto_layout )
256273 # build logger
257274 filehandler = logging .FileHandler (os .path .join (args .save_dir , args .logging_file ))
258275 streamhandler = logging .StreamHandler ()
0 commit comments