-
Notifications
You must be signed in to change notification settings - Fork 21
Expand file tree
/
Copy pathNFC.py
More file actions
36 lines (29 loc) · 1.04 KB
/
Copy pathNFC.py
File metadata and controls
36 lines (29 loc) · 1.04 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
import torch
def pairwise_distance(query_features, gallery_features):
x = query_features
y = gallery_features
m, n = x.size(0), y.size(0)
x = x.view(m, -1)
y = y.view(n, -1)
dist = torch.pow(x, 2).sum(dim=1, keepdim=True).expand(m, n) + \
torch.pow(y, 2).sum(dim=1, keepdim=True).expand(n, m).t()
dist.addmm_(1, -2, x, y.t())
# dist = dist.clamp(min=1e-12).sqrt()
return dist
def NFC(feat: torch.tensor, k1=2, k2=2):
feat = feat.clone()
dist = pairwise_distance(feat.to('cuda'), feat.to('cuda')).to('cpu')
eye = torch.eye(dist.size(0)).to(dist.device)
dist[eye == 1] = 1000
val, rank = dist.topk(k1, largest=False)
mutual_topk_list = []
for i in range(rank.size(0)):
mutual_list = []
for j in rank[i]:
if i in rank[j][:k2]:
mutual_list.append(j.item())
mutual_topk_list.append(mutual_list)
feat_copy = feat.clone()
for i in range(rank.size(0)):
feat[i] += feat_copy[mutual_topk_list[i]].sum(dim=0)
return feat