MCPcopy Create free account
hub / github.com/Sense-GVT/DeCLIP / DistributedEpochSampler

Class DistributedEpochSampler

prototype/data/sampler.py:109–169  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

107
108
109class DistributedEpochSampler(Sampler):
110 def __init__(self, dataset, total_iter, batch_size, world_size=None, rank=None, last_iter=0):
111 if world_size is None:
112 world_size = link.get_world_size()
113 if rank is None:
114 rank = link.get_rank()
115 assert rank < world_size
116 self.dataset = dataset
117 self.total_iter = total_iter
118 self.batch_size = batch_size
119 self.world_size = world_size
120 self.rank = rank
121 self.last_iter = last_iter
122
123 self.all_size_single = self.total_iter * self.batch_size
124
125 self.indices = self.gen_new_list()
126 self.call = 0
127
128 def __iter__(self):
129 if self.call == 0:
130 self.call = 1
131 return iter(self.indices[self.last_iter*self.batch_size:])
132 else:
133 raise RuntimeError(
134 "this sampler is not designed to be called more than once!!")
135
136 def get_one_epoch_self_part(self):
137 num = len(self.dataset)
138 indices = np.arange(num)
139 extra_indices = np.random.choice(
140 num, self.extra_per_epoch, replace=False)
141 indices = np.concatenate((indices, extra_indices))
142 np.random.shuffle(indices)
143 assert len(indices) % (self.world_size * self.batch_size) == 0
144 num_single = len(indices) // self.world_size
145 return indices[self.rank*num_single:(self.rank+1)*num_single]
146
147 def gen_new_list(self):
148 np.random.seed(0)
149
150 self.all_num = self.total_iter * self.batch_size * self.world_size
151 iter_per_epoch = (len(self.dataset) -
152 1) // (self.batch_size * self.world_size) + 1
153 self.num_per_epoch = iter_per_epoch * self.batch_size * self.world_size
154 self.extra_per_epoch = self.num_per_epoch - len(self.dataset)
155 repeat = (self.all_num - 1) // self.num_per_epoch + 1
156 indices = []
157 for i in range(repeat):
158 indice = self.get_one_epoch_self_part()
159 indices.append(indice)
160
161 indices = np.concatenate(indices)
162 indices = indices[:self.all_size_single]
163
164 assert len(indices) == self.all_size_single
165
166 return indices

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected