@@ -108,6 +108,8 @@ def make_parser():
108108 parser .add_argument ('--num-workers' , type = int , default = 4 )
109109 parser .add_argument ('--amp' , action = 'store_true' ,
110110 help = 'Whether to enable AMP ops. When false, uses TF32 on A100 and FP32 on V100 GPUS.' )
111+ parser .add_argument ('--log-interval' , type = int , default = 20 ,
112+ help = 'Logging interval.' )
111113 parser .add_argument ('--json-summary' , type = str , default = None ,
112114 help = 'If provided, the json summary will be written to'
113115 'the specified file.' )
@@ -273,15 +275,18 @@ def log_params(logger, args):
273275
274276 if args .mode == 'benchmark-training' :
275277 train_loop_func = benchmark_train_loop
276- logger = BenchLogger ('Training benchmark' , json_output = args .json_summary )
278+ logger = BenchLogger ('Training benchmark' , log_interval = args .log_interval ,
279+ json_output = args .json_summary )
277280 args .epochs = 1
278281 elif args .mode == 'benchmark-inference' :
279282 train_loop_func = benchmark_inference_loop
280- logger = BenchLogger ('Inference benchmark' , json_output = args .json_summary )
283+ logger = BenchLogger ('Inference benchmark' , log_interval = args .log_interval ,
284+ json_output = args .json_summary )
281285 args .epochs = 1
282286 else :
283287 train_loop_func = train_loop
284- logger = Logger ('Training logger' , print_freq = 1 , json_output = args .json_summary )
288+ logger = Logger ('Training logger' , log_interval = args .log_interval ,
289+ json_output = args .json_summary )
285290
286291 log_params (logger , args )
287292
0 commit comments