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)
| 463 | |
| 464 | |
| 465 | def 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) |
no test coverage detected