| 44 | |
| 45 | |
| 46 | class 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: |
nothing calls this directly
no outgoing calls
no test coverage detected