MCPcopy Create free account
hub / github.com/TPCD/DCCL / extract_features

Function extract_features

project_utils/contrastive_utils.py:30–71  ·  view source on GitHub ↗
(model, data_loader, print_freq=50, args=None)

Source from the content-addressed store, hash-verified

28
29
30def extract_features(model, data_loader, print_freq=50, args=None):
31 model.eval()
32 batch_time = AverageMeter()
33 data_time = AverageMeter()
34
35 features = []
36 labels = []
37 indexs = []
38
39 end = time.time()
40 with torch.no_grad():
41 for i, _item in enumerate(data_loader):
42 imgs = _item[0]
43 targets = _item[1]
44 # uq_idx = _item[2]
45 if_train = _item[3][:, 0].bool()
46 data_time.update(time.time() - end)
47 if args is not None:
48 imgs = to_torch(imgs).to(args.device)
49 else:
50 imgs = to_torch(imgs).cuda()
51 outputs = model(imgs)
52 outputs = outputs.data.cpu()
53
54 features.append(outputs)
55 labels.append(targets)
56 # indexs.append(uq_idx)
57 indexs.append(if_train)
58
59
60 batch_time.update(time.time() - end)
61 end = time.time()
62
63 if (i + 1) % print_freq == 0:
64 print('Extract Features: [{}/{}]\t'
65 'Time {:.3f} ({:.3f})\t'
66 'Data {:.3f} ({:.3f})\t'
67 .format(i + 1, len(data_loader),
68 batch_time.val, batch_time.avg,
69 data_time.val, data_time.avg))
70
71 return features, labels, indexs
72
73def extract_features2(model, data_loader, print_freq=50):
74 model.eval()

Callers 1

trainFunction · 0.90

Calls 3

updateMethod · 0.95
to_torchFunction · 0.85
AverageMeterClass · 0.70

Tested by

no test coverage detected