Skip to content

Commit d6cd6b8

Browse files
shakandrewnv-kkudrynski
authored andcommitted
[SSD/PyT] Improved logging
1 parent 01bbec9 commit d6cd6b8

2 files changed

Lines changed: 11 additions & 6 deletions

File tree

PyTorch/Detection/SSD/main.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -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

PyTorch/Detection/SSD/src/logger.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -54,12 +54,12 @@ def update_epoch(self, epoch):
5454

5555

5656
class Logger:
57-
def __init__(self, name, json_output=None, print_freq=20):
57+
def __init__(self, name, json_output=None, log_interval=20):
5858
self.name = name
5959
self.train_loss_logger = IterationAverageMeter("Training loss")
6060
self.train_epoch_time_logger = EpochMeter("Training 1 epoch time")
6161
self.val_acc_logger = EpochMeter("Validation accuracy")
62-
self.print_freq = print_freq
62+
self.log_interval = log_interval
6363

6464
backends = [ DLLogger.StdOutBackend(DLLogger.Verbosity.DEFAULT) ]
6565
if json_output:
@@ -95,7 +95,7 @@ def log_summary(self):
9595
def update_iter(self, epoch, iteration, loss):
9696
self.train_iter = iteration
9797
self.train_loss_logger.update_iter(loss)
98-
if iteration % self.print_freq == 0:
98+
if iteration % self.log_interval == 0:
9999
self.log('loss', loss)
100100

101101
def update_epoch(self, epoch, acc):

0 commit comments

Comments
 (0)