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

Function extract_features2

project_utils/contrastive_utils.py:73–102  ·  view source on GitHub ↗
(model, data_loader, print_freq=50)

Source from the content-addressed store, hash-verified

71 return features, labels, indexs
72
73def extract_features2(model, data_loader, print_freq=50):
74 model.eval()
75 batch_time = AverageMeter()
76 data_time = AverageMeter()
77
78 features = OrderedDict()
79 labels = OrderedDict()
80
81 end = time.time()
82 with torch.no_grad():
83 for i, (imgs, targets, uq_idx, attribute) in enumerate(data_loader):
84 data_time.update(time.time() - end)
85
86 outputs = extract_cnn_feature(model, imgs)
87 for fname, output, pid in zip(uq_idx, outputs, targets):
88 features[fname] = output
89 labels[fname] = pid
90
91 batch_time.update(time.time() - end)
92 end = time.time()
93
94 if (i + 1) % print_freq == 0:
95 print('Extract Features: [{}/{}]\t'
96 'Time {:.3f} ({:.3f})\t'
97 'Data {:.3f} ({:.3f})\t'
98 .format(i + 1, len(data_loader),
99 batch_time.val, batch_time.avg,
100 data_time.val, data_time.avg))
101
102 return features, labels
103# Ensure that all operations are deterministic on GPU (if used) for reproducibility
104def accuracy(output, target, topk=(1,)):
105 """Computes the accuracy over the k top predictions for the specified values of k"""

Callers

nothing calls this directly

Calls 2

updateMethod · 0.95
AverageMeterClass · 0.70

Tested by

no test coverage detected