(model, data_loader, print_freq=50, args=None)
| 28 | |
| 29 | |
| 30 | def 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 | |
| 73 | def extract_features2(model, data_loader, print_freq=50): |
| 74 | model.eval() |
no test coverage detected