MCPcopy Create free account
hub / github.com/XLearning-SCU/2022-CVPR-DART / test

Function test

run.py:411–452  ·  view source on GitHub ↗
(net1, net2)

Source from the content-addressed store, hash-verified

409
410
411def 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

Callers 1

run.pyFile · 0.85

Calls 2

eval_regdbFunction · 0.90
eval_sysuFunction · 0.90

Tested by

no test coverage detected