MCPcopy Create free account
hub / github.com/MotrixLab/ViMoGen / DistributedFlopBalanceSampler

Class DistributedFlopBalanceSampler

datasets/multi_resolution_sampler.py:92–164  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

90
91
92class 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

Callers 1

get_dataloaderFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected