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

Class RandomIdentitySampler

project_utils/sampler.py:19–43  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

17
18
19class RandomIdentitySampler(Sampler):
20 def __init__(self, data_source, num_instances):
21 self.data_source = data_source
22 self.num_instances = num_instances
23 self.index_dic = defaultdict(list)
24 for index, (_, pid, _) in enumerate(data_source):
25 self.index_dic[pid].append(index)
26 self.pids = list(self.index_dic.keys())
27 self.num_samples = len(self.pids)
28
29 def __len__(self):
30 return self.num_samples * self.num_instances
31
32 def __iter__(self):
33 indices = torch.randperm(self.num_samples).tolist()
34 ret = []
35 for i in indices:
36 pid = self.pids[i]
37 t = self.index_dic[pid]
38 if len(t) >= self.num_instances:
39 t = np.random.choice(t, size=self.num_instances, replace=False)
40 else:
41 t = np.random.choice(t, size=self.num_instances, replace=True)
42 ret.extend(t)
43 return iter(ret)
44
45
46class RandomMultipleGallerySampler(Sampler):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected