(step, G, Cs, dataset_source, dataset_target, save_path,
ova=True)
| 11 | |
| 12 | |
| 13 | def feat_get(step, G, Cs, dataset_source, dataset_target, save_path, |
| 14 | ova=True): |
| 15 | G.eval() |
| 16 | |
| 17 | for batch_idx, data in enumerate(dataset_source): |
| 18 | if batch_idx == 500: |
| 19 | break |
| 20 | with torch.no_grad(): |
| 21 | img_s = data[0] |
| 22 | label_s = data[1] |
| 23 | img_s, label_s = Variable(img_s.cuda()), \ |
| 24 | Variable(label_s.cuda()) |
| 25 | feat_s = G(img_s) |
| 26 | |
| 27 | |
| 28 | if batch_idx == 0: |
| 29 | feat_all_s = feat_s.data.cpu().numpy() |
| 30 | label_all_s = label_s.data.cpu().numpy() |
| 31 | else: |
| 32 | feat_s = feat_s.data.cpu().numpy() |
| 33 | label_s = label_s.data.cpu().numpy() |
| 34 | feat_all_s = np.r_[feat_all_s, feat_s] |
| 35 | label_all_s = np.r_[label_all_s, label_s] |
| 36 | for batch_idx, data in enumerate(dataset_target): |
| 37 | if batch_idx == 500: |
| 38 | break |
| 39 | with torch.no_grad(): |
| 40 | img_t = data[0] |
| 41 | label_t = data[1] |
| 42 | img_t, label_t = Variable(img_t.cuda()), \ |
| 43 | Variable(label_t.cuda()) |
| 44 | feat_t = G(img_t) |
| 45 | |
| 46 | out_t = Cs[0](feat_t) |
| 47 | pred = out_t.data.max(1)[1] |
| 48 | out_t = F.softmax(out_t) |
| 49 | if ova: |
| 50 | out_open = Cs[1](feat_t) |
| 51 | out_open = F.softmax(out_open.view(out_t.size(0), 2, -1), 1) |
| 52 | tmp_range = torch.range(0, out_t.size(0) - 1).long().cuda() |
| 53 | pred_unk = out_open[tmp_range, 0, pred] |
| 54 | weights_open = Cs[1].module.fc.weight.data.cpu().numpy() |
| 55 | else: |
| 56 | pred_unk = -torch.sum(out_t * torch.log(out_t), 1) |
| 57 | |
| 58 | if batch_idx == 0: |
| 59 | feat_all = feat_t.data.cpu().numpy() |
| 60 | label_all = label_t.data.cpu().numpy() |
| 61 | unk_all = pred_unk.data.cpu().numpy() |
| 62 | pred_all = pred.data.cpu().numpy() |
| 63 | pred_all_soft = out_t.data.cpu().numpy() |
| 64 | else: |
| 65 | feat_t = feat_t.data.cpu().numpy() |
| 66 | label_t = label_t.data.cpu().numpy() |
| 67 | pred_unk = pred_unk.data.cpu().numpy() |
| 68 | feat_all = np.r_[feat_all, feat_t] |
| 69 | label_all = np.r_[label_all, label_t] |
| 70 | unk_all = np.r_[unk_all, pred_unk] |
nothing calls this directly
no outgoing calls
no test coverage detected