(cfg, model, center_criterion, train_loader, val_loader, optimizer, optimizer_center, scheduler, loss_fn, num_query, local_rank)
| 87 | |
| 88 | |
| 89 | def do_train(cfg, model, center_criterion, train_loader, val_loader, optimizer, optimizer_center, scheduler, loss_fn, num_query, local_rank): |
| 90 | log_period = cfg.SOLVER.LOG_PERIOD |
| 91 | checkpoint_period = cfg.SOLVER.CHECKPOINT_PERIOD |
| 92 | eval_period = cfg.SOLVER.EVAL_PERIOD |
| 93 | |
| 94 | device = "cuda" |
| 95 | epochs = cfg.SOLVER.MAX_EPOCHS |
| 96 | |
| 97 | logger = logging.getLogger("transreid.train") |
| 98 | logger.info("start training") |
| 99 | _LOCAL_PROCESS_GROUP = None |
| 100 | |
| 101 | if device: |
| 102 | model.to(local_rank) |
| 103 | if torch.cuda.device_count() > 1 and cfg.MODEL.DIST_TRAIN: |
| 104 | print("Using {} GPUs for training".format(torch.cuda.device_count())) |
| 105 | model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank], find_unused_parameters=True) |
| 106 | |
| 107 | loss_meter = AverageMeter() |
| 108 | acc_meter = AverageMeter() |
| 109 | |
| 110 | evaluator = R1_mAP_eval(num_query, max_rank=50, feat_norm=cfg.TEST.FEAT_NORM) |
| 111 | scaler = amp.GradScaler() |
| 112 | |
| 113 | # train |
| 114 | if torch.cuda.device_count() > 1 and cfg.MODEL.DIST_TRAIN: |
| 115 | model.module.train_with_single() |
| 116 | else: |
| 117 | model.train_with_single() |
| 118 | for epoch in range(1, epochs + 1): |
| 119 | start_time = time.time() |
| 120 | loss_meter.reset() |
| 121 | acc_meter.reset() |
| 122 | evaluator.reset() |
| 123 | scheduler.step(epoch) |
| 124 | model.train() |
| 125 | for n_iter, (img, vid, target_cam, target_view, img_wh) in enumerate(train_loader): |
| 126 | optimizer.zero_grad() |
| 127 | optimizer_center.zero_grad() |
| 128 | img = img.to(device) |
| 129 | target = vid.to(device) |
| 130 | target_cam = target_cam.to(device) |
| 131 | img_wh = img_wh.to(device) |
| 132 | with amp.autocast(enabled=True): |
| 133 | score, feat = model(img, target, cam_label=target_cam, img_wh=img_wh) |
| 134 | loss = loss_fn(score, feat, target, target_cam) |
| 135 | |
| 136 | scaler.scale(loss).backward() |
| 137 | |
| 138 | scaler.step(optimizer) |
| 139 | scaler.update() |
| 140 | |
| 141 | if "center" in cfg.MODEL.METRIC_LOSS_TYPE: |
| 142 | for param in center_criterion.parameters(): |
| 143 | param.grad.data *= 1.0 / cfg.SOLVER.CENTER_LOSS_WEIGHT |
| 144 | scaler.step(optimizer_center) |
| 145 | scaler.update() |
| 146 | if isinstance(score, list): |
no test coverage detected