| 90 | |
| 91 | |
| 92 | class DistributedFlopBalanceSampler(Sampler): |
| 93 | |
| 94 | def __init__( |
| 95 | self, |
| 96 | dataset, |
| 97 | dp_rank: int, |
| 98 | dp_size: int, |
| 99 | global_seed: int = 0, |
| 100 | bucket_config_type: str = 'DefaultBucketConfigNotExact', |
| 101 | call_set_epoch: bool = True, |
| 102 | ) -> None: |
| 103 | self.dataset = dataset |
| 104 | self.global_seed = global_seed |
| 105 | self.dp_rank = dp_rank |
| 106 | self.dp_size = dp_size |
| 107 | |
| 108 | assert bucket_config_type in [ |
| 109 | 'DefaultBucketConfigNotExact', 'DefaultBucketConfig', |
| 110 | 'BucketConfigHardCoded1', 'BucketConfigHardCoded2', |
| 111 | 'BucketConfigHardCoded3', 'DefaultBucketConfig3ARNotExact' |
| 112 | ] |
| 113 | bucket_config = BucketConfig.from_class_name(bucket_config_type, {}) |
| 114 | ori_size_list = self.get_ori_size_list() |
| 115 | # ensure the global_seed is the same across different processes |
| 116 | self.rnd_state = np.random.RandomState(global_seed) |
| 117 | self.bucket_factory = BucketFactory( |
| 118 | ori_size_list, |
| 119 | dp_size=self.dp_size, |
| 120 | rnd_state=self.rnd_state, |
| 121 | bucket_config=bucket_config, |
| 122 | ) |
| 123 | if call_set_epoch: |
| 124 | self.set_epoch(0) |
| 125 | |
| 126 | def get_ori_size_list(self, ): |
| 127 | ori_size_list = [None] * len(self.dataset.data_list) |
| 128 | for i, data in enumerate(self.dataset.data_list): |
| 129 | ori_size_list[i] = (data['length'], data['height'], data['width']) |
| 130 | return ori_size_list |
| 131 | |
| 132 | def bucket_prepare(self, ): |
| 133 | print('----------prepare bucket-----------') |
| 134 | self.final_idx_list, self.flop_list, self.final_bucket_key_list = self.bucket_factory( |
| 135 | ) |
| 136 | assert len(self.final_idx_list) >= len(self.flop_list) |
| 137 | assert len(self.final_idx_list) % self.dp_size == 0 |
| 138 | print('----------bucket prepared-----------') |
| 139 | |
| 140 | def set_epoch(self, epoch: int) -> None: |
| 141 | """different from DistributedSampler, you don't need to call this |
| 142 | function at the start of each epoch, because the return Iterator of |
| 143 | __iter__ will be different for each epoch originally.""" |
| 144 | self.epoch_count = epoch |
| 145 | self.rnd_state.seed(epoch + self.global_seed) |
| 146 | self.bucket_prepare() |
| 147 | |
| 148 | def __len__(self, ): |
| 149 | return len(self.final_idx_list) // self.dp_size |