MCPcopy Create free account
hub / github.com/TPCD/DCCL / RandomMultipleGallerySampler

Class RandomMultipleGallerySampler

project_utils/sampler.py:46–106  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

44
45
46class RandomMultipleGallerySampler(Sampler):
47 def __init__(self, data_source, num_instances=4):
48 super().__init__(data_source)
49 self.data_source = data_source
50 self.index_pid = defaultdict(int)
51 self.pid_cam = defaultdict(list)
52 self.pid_index = defaultdict(list)
53 self.num_instances = num_instances
54
55 for index, (_, pid, cam) in enumerate(data_source):
56 if pid < 0:
57 continue
58 self.index_pid[index] = pid
59 self.pid_cam[pid].append(cam)
60 self.pid_index[pid].append(index)
61
62 self.pids = list(self.pid_index.keys())
63 self.num_samples = len(self.pids)
64
65 def __len__(self):
66 return self.num_samples * self.num_instances
67
68 def __iter__(self):
69 indices = torch.randperm(len(self.pids)).tolist()
70 ret = []
71
72 for kid in indices:
73 i = random.choice(self.pid_index[self.pids[kid]])
74
75 _, i_pid, i_cam = self.data_source[i]
76
77 ret.append(i)
78
79 pid_i = self.index_pid[i]
80 cams = self.pid_cam[pid_i]
81 index = self.pid_index[pid_i]
82 select_cams = No_index(cams, i_cam)
83
84 if select_cams:
85
86 if len(select_cams) >= self.num_instances:
87 cam_indexes = np.random.choice(select_cams, size=self.num_instances-1, replace=False)
88 else:
89 cam_indexes = np.random.choice(select_cams, size=self.num_instances-1, replace=True)
90
91 for kk in cam_indexes:
92 ret.append(index[kk])
93
94 else:
95 select_indexes = No_index(index, i)
96 if not select_indexes:
97 continue
98 if len(select_indexes) >= self.num_instances:
99 ind_indexes = np.random.choice(select_indexes, size=self.num_instances-1, replace=False)
100 else:
101 ind_indexes = np.random.choice(select_indexes, size=self.num_instances-1, replace=True)
102
103 for kk in ind_indexes:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected