MCPcopy Create free account
hub / github.com/Alioth2000/Hoss-ReID / do_train

Function do_train

processor/processor.py:89–212  ·  view source on GitHub ↗
(cfg, model, center_criterion, train_loader, val_loader, optimizer, optimizer_center, scheduler, loss_fn, num_query, local_rank)

Source from the content-addressed store, hash-verified

87
88
89def 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):

Callers 1

train.pyFile · 0.90

Calls 11

resetMethod · 0.95
resetMethod · 0.95
updateMethod · 0.95
updateMethod · 0.95
computeMethod · 0.95
AverageMeterClass · 0.90
R1_mAP_evalClass · 0.90
train_with_singleMethod · 0.80
stepMethod · 0.80
state_dictMethod · 0.80
_get_lrMethod · 0.45

Tested by

no test coverage detected