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

Function do_train_pair

processor/processor.py:13–86  ·  view source on GitHub ↗
(cfg, model, train_loader_pair, optimizer, scheduler, local_rank)

Source from the content-addressed store, hash-verified

11
12
13def do_train_pair(cfg, model, train_loader_pair, optimizer, scheduler, local_rank):
14 log_period = cfg.SOLVER.LOG_PERIOD
15 checkpoint_period = cfg.SOLVER.CHECKPOINT_PERIOD
16
17 device = "cuda"
18 epochs = cfg.SOLVER.MAX_EPOCHS
19
20 logger = logging.getLogger("transreid.train")
21 logger.info("start training")
22 _LOCAL_PROCESS_GROUP = None
23
24 if device:
25 model.to(local_rank)
26 if torch.cuda.device_count() > 1 and cfg.MODEL.DIST_TRAIN:
27 print("Using {} GPUs for training".format(torch.cuda.device_count()))
28 model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank], find_unused_parameters=True)
29
30 loss_meter = AverageMeter()
31 scaler = amp.GradScaler()
32
33 # train pair
34 if cfg.MODEL.PAIR:
35 if torch.cuda.device_count() > 1 and cfg.MODEL.DIST_TRAIN:
36 model.module.train_with_pair()
37 else:
38 model.train_with_pair()
39 for epoch in range(1, epochs + 1):
40 start_time = time.time()
41 loss_meter.reset()
42 scheduler.step(epoch)
43 model.train()
44 if hasattr(train_loader_pair, "sampler") and hasattr(train_loader_pair.sampler, "set_epoch"):
45 train_loader_pair.sampler.set_epoch(epoch)
46 for n_iter, (img, vid, target_cam) in enumerate(train_loader_pair):
47 optimizer.zero_grad()
48 img = img.to(device)
49 target = vid.to(device)
50 target_cam = target_cam.to(device)
51 with amp.autocast(enabled=True):
52 logits_per_sar = model(img, target, cam_label=target_cam)
53 loss = clip_loss(logits_per_sar)
54
55 scaler.scale(loss).backward()
56
57 scaler.step(optimizer)
58 scaler.update()
59
60 loss_meter.update(loss.item(), img.shape[0])
61
62 torch.cuda.synchronize()
63 if (n_iter + 1) % log_period == 0:
64 logger.info(
65 "Epoch[{}] Iteration[{}/{}] Loss: {:.3f}, Base Lr: {:.2e}".format(
66 epoch, (n_iter + 1), len(train_loader_pair), loss_meter.avg, scheduler._get_lr(epoch)[0]
67 )
68 )
69
70 end_time = time.time()

Callers 1

train_pair.pyFile · 0.90

Calls 8

resetMethod · 0.95
updateMethod · 0.95
AverageMeterClass · 0.90
clip_lossFunction · 0.90
train_with_pairMethod · 0.80
stepMethod · 0.80
state_dictMethod · 0.80
_get_lrMethod · 0.45

Tested by

no test coverage detected