MCPcopy Create free account
hub / github.com/NVIDIA/semantic-segmentation / train

Function train

train.py:465–533  ·  view source on GitHub ↗

Runs the training loop per epoch train_loader: Data loader for train net: thet network optimizer: optimizer curr_epoch: current epoch return:

(train_loader, net, optim, curr_epoch)

Source from the content-addressed store, hash-verified

463
464
465def train(train_loader, net, optim, curr_epoch):
466 """
467 Runs the training loop per epoch
468 train_loader: Data loader for train
469 net: thet network
470 optimizer: optimizer
471 curr_epoch: current epoch
472 return:
473 """
474 net.train()
475
476 train_main_loss = AverageMeter()
477 start_time = None
478 warmup_iter = 10
479
480 for i, data in enumerate(train_loader):
481 if i <= warmup_iter:
482 start_time = time.time()
483 # inputs = (bs,3,713,713)
484 # gts = (bs,713,713)
485 images, gts, _img_name, scale_float = data
486 batch_pixel_size = images.size(0) * images.size(2) * images.size(3)
487 images, gts, scale_float = images.cuda(), gts.cuda(), scale_float.cuda()
488 inputs = {'images': images, 'gts': gts}
489
490 optim.zero_grad()
491 main_loss = net(inputs)
492
493 if args.apex:
494 log_main_loss = main_loss.clone().detach_()
495 torch.distributed.all_reduce(log_main_loss,
496 torch.distributed.ReduceOp.SUM)
497 log_main_loss = log_main_loss / args.world_size
498 else:
499 main_loss = main_loss.mean()
500 log_main_loss = main_loss.clone().detach_()
501
502 train_main_loss.update(log_main_loss.item(), batch_pixel_size)
503 if args.fp16:
504 with amp.scale_loss(main_loss, optim) as scaled_loss:
505 scaled_loss.backward()
506 else:
507 main_loss.backward()
508
509 optim.step()
510
511 if i >= warmup_iter:
512 curr_time = time.time()
513 batches = i - warmup_iter + 1
514 batchtime = (curr_time - start_time) / batches
515 else:
516 batchtime = 0
517
518 msg = ('[epoch {}], [iter {} / {}], [train main loss {:0.6f}],'
519 ' [lr {:0.6f}] [batchtime {:0.3g}]')
520 msg = msg.format(
521 curr_epoch, i + 1, len(train_loader), train_main_loss.avg,
522 optim.param_groups[-1]['lr'], batchtime)

Callers 1

mainFunction · 0.85

Calls 3

updateMethod · 0.95
AverageMeterClass · 0.90
stepMethod · 0.80

Tested by

no test coverage detected