(net1, net2)
| 409 | |
| 410 | |
| 411 | def test(net1, net2): |
| 412 | # switch to evaluation mode |
| 413 | net1.eval() |
| 414 | net2.eval() |
| 415 | ptr = 0 |
| 416 | gall_feat_att = np.zeros((ngall, 2048)) |
| 417 | |
| 418 | with torch.no_grad(): |
| 419 | for batch_idx, (input, label) in enumerate(gall_loader): |
| 420 | batch_num = input.size(0) |
| 421 | input = input.cuda() |
| 422 | _, feat_att1 = net1(input, input, test_mode[0]) |
| 423 | _, feat_att2 = net2(input, input, test_mode[0]) |
| 424 | feat_att = (feat_att1 + feat_att2) / 2. |
| 425 | gall_feat_att[ptr:ptr + batch_num, :] = feat_att.detach().cpu().numpy() |
| 426 | ptr = ptr + batch_num |
| 427 | |
| 428 | # switch to evaluation |
| 429 | net1.eval() |
| 430 | net2.eval() |
| 431 | ptr = 0 |
| 432 | query_feat_att = np.zeros((nquery, 2048)) |
| 433 | with torch.no_grad(): |
| 434 | for batch_idx, (input, label) in enumerate(query_loader): |
| 435 | batch_num = input.size(0) |
| 436 | input = input.cuda() |
| 437 | _, feat_att1 = net1(input, input, test_mode[1]) |
| 438 | _, feat_att2 = net2(input, input, test_mode[1]) |
| 439 | feat_att = (feat_att1 + feat_att2) / 2. |
| 440 | query_feat_att[ptr:ptr + batch_num, :] = feat_att.detach().cpu().numpy() |
| 441 | ptr = ptr + batch_num |
| 442 | |
| 443 | # compute the similarity |
| 444 | distmat_att = np.matmul(query_feat_att, np.transpose(gall_feat_att)) |
| 445 | |
| 446 | # evaluation |
| 447 | if dataset == 'regdb': |
| 448 | cmc_att, mAP_att, mINP_att = eval_regdb(-distmat_att, query_label, gall_label) |
| 449 | elif dataset == 'sysu': |
| 450 | cmc_att, mAP_att, mINP_att = eval_sysu(-distmat_att, query_label, gall_label, query_cam, gall_cam) |
| 451 | |
| 452 | return cmc_att, mAP_att, mINP_att |
| 453 | |
| 454 | |
| 455 | # training |
no test coverage detected