| 4 | |
| 5 | |
| 6 | class SeqInferenceSampler(Sampler): |
| 7 | def __init__(self, seq_name_list: List, size: int): |
| 8 | """ |
| 9 | Args: |
| 10 | size (int): the total number of data of the underlying dataset to sample from |
| 11 | """ |
| 12 | self._size = size |
| 13 | assert size > 0 |
| 14 | self._rank = comm.get_rank() |
| 15 | self._world_size = comm.get_world_size() |
| 16 | self.idx_per_seq = self.build_idx_per_seq(seq_name_list) |
| 17 | self.sequence_names = list(self.idx_per_seq.keys()) |
| 18 | self.pad_idx_per_rank = self.build_idx_per_rank() |
| 19 | self._local_indices = self.pad_idx_per_rank[self._rank] |
| 20 | |
| 21 | def build_idx_per_seq(self, seq_name_list): |
| 22 | idx_per_seq = {} |
| 23 | for idx, seq_name in enumerate(seq_name_list): |
| 24 | if seq_name not in idx_per_seq: |
| 25 | idx_per_seq[seq_name] = [] |
| 26 | idx_per_seq[seq_name].append(idx) |
| 27 | return idx_per_seq |
| 28 | |
| 29 | def build_idx_per_rank(self): |
| 30 | total_num_sequence = len(self.sequence_names) |
| 31 | num_seq_per_rank = (total_num_sequence - 1) // self._world_size + 1 |
| 32 | idx_per_rank = {} |
| 33 | for rank_id in range(self._world_size): |
| 34 | begin_seq_idx = rank_id * num_seq_per_rank |
| 35 | end_seq_idx = (rank_id + 1) * num_seq_per_rank |
| 36 | idx_list = [] |
| 37 | for seq_name in self.sequence_names[begin_seq_idx:end_seq_idx]: |
| 38 | idx_list = [*idx_list, *self.idx_per_seq[seq_name]] |
| 39 | idx_per_rank[rank_id] = idx_list |
| 40 | |
| 41 | pad_idx_per_rank = self.pad_idx(idx_per_rank) |
| 42 | return pad_idx_per_rank |
| 43 | |
| 44 | def pad_idx(self, idx_per_rank): |
| 45 | len_per_rank = [len(x) for x in idx_per_rank.values()] |
| 46 | max_len = max(len_per_rank) |
| 47 | new_idx_per_rank = {} |
| 48 | for rank, indicies in idx_per_rank.items(): |
| 49 | pad_size = max_len - len(indicies) |
| 50 | if pad_size > 0: |
| 51 | pad_indices = indicies + indicies[:pad_size] |
| 52 | else: |
| 53 | pad_indices = indicies |
| 54 | new_idx_per_rank[rank] = pad_indices |
| 55 | return new_idx_per_rank |
| 56 | |
| 57 | def __iter__(self): |
| 58 | yield from self._local_indices |
| 59 | |
| 60 | def __len__(self): |
| 61 | return len(self._local_indices) |
no outgoing calls
no test coverage detected