| 5 | |
| 6 | |
| 7 | def visualize(config, base, loaders): |
| 8 | |
| 9 | base.set_eval() |
| 10 | |
| 11 | # meters |
| 12 | query_features_meter, query_pids_meter, query_cids_meter = CatMeter(), CatMeter(), CatMeter() |
| 13 | gallery_features_meter, gallery_pids_meter, gallery_cids_meter = CatMeter(), CatMeter(), CatMeter() |
| 14 | |
| 15 | # init dataset |
| 16 | if config.visualize_dataset == 'market': |
| 17 | _datasets = [loaders.market_query_samples, loaders.market_gallery_samples] |
| 18 | _loaders = [loaders.market_query_loader, loaders.market_gallery_loader] |
| 19 | elif config.visualize_dataset == 'duke': |
| 20 | _datasets = [loaders.duke_query_samples, loaders.duke_gallery_samples] |
| 21 | _loaders = [loaders.duke_query_loader, loaders.duke_gallery_loader] |
| 22 | elif config.visualize_dataset == 'customed': |
| 23 | _datasets = [loaders.query_samples, loaders.gallery_samples] |
| 24 | _loaders = [loaders.query_loader, loaders.gallery_loader] |
| 25 | |
| 26 | # compute query and gallery features |
| 27 | with torch.no_grad(): |
| 28 | for loader_id, loader in enumerate(_loaders): |
| 29 | for data in loader: |
| 30 | # compute feautres |
| 31 | images, pids, cids = data |
| 32 | images = images.cuda() |
| 33 | features = base.model(images) |
| 34 | # save as query features |
| 35 | if loader_id == 0: |
| 36 | query_features_meter.update(features.data) |
| 37 | query_pids_meter.update(pids) |
| 38 | query_cids_meter.update(cids) |
| 39 | # save as gallery features |
| 40 | elif loader_id == 1: |
| 41 | gallery_features_meter.update(features.data) |
| 42 | gallery_pids_meter.update(pids) |
| 43 | gallery_cids_meter.update(cids) |
| 44 | |
| 45 | # compute distance |
| 46 | query_features = query_features_meter.get_val() |
| 47 | gallery_features = gallery_features_meter.get_val() |
| 48 | |
| 49 | if config.test_metric is 'cosine': |
| 50 | distance = tensor_cosine_dist(query_features, gallery_features).data.cpu().numpy() |
| 51 | |
| 52 | elif config.test_metric is 'euclidean': |
| 53 | distance = tensor_euclidean_dist(query_features, gallery_features).data.cpu().numpy() |
| 54 | |
| 55 | # visualize |
| 56 | visualize_ranked_results(distance, _datasets, config.visualize_output_path, mode=config.visualize_mode, only_show=config.visualize_mode_onlyshow) |