MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / all_gather

Function all_gather

util/distribute_utils.py:178–223  ·  view source on GitHub ↗

Run all_gather on arbitrary picklable data (not necessarily tensors) Args: data: Any picklable object Returns: data_list(list): List of data gathered from each rank

(data)

Source from the content-addressed store, hash-verified

176
177
178def all_gather(data):
179 """
180 Run all_gather on arbitrary picklable data (not necessarily tensors)
181 Args:
182 data:
183 Any picklable object
184 Returns:
185 data_list(list):
186 List of data gathered from each rank
187 """
188 world_size = get_world_size()
189 if world_size == 1:
190 return [data]
191
192 # serialized to a Tensor
193 buffer = pickle.dumps(data)
194 storage = torch.ByteStorage.from_buffer(buffer)
195 tensor = torch.ByteTensor(storage).to('cuda')
196
197 # obtain Tensor size of each rank
198 local_size = torch.tensor([tensor.numel()], device='cuda')
199 size_list = [torch.tensor([0], device='cuda') for _ in range(world_size)]
200 dist.all_gather(size_list, local_size)
201 size_list = [int(size.item()) for size in size_list]
202 max_size = max(size_list)
203
204 # receiving Tensor from all ranks
205 # we pad the tensor because torch all_gather does not support
206 # gathering tensors of different shapes
207 tensor_list = []
208 for _ in size_list:
209 tensor_list.append(
210 torch.empty((max_size, ), dtype=torch.uint8, device='cuda'))
211 if local_size != max_size:
212 padding = torch.empty(size=(max_size - local_size, ),
213 dtype=torch.uint8,
214 device='cuda')
215 tensor = torch.cat((tensor, padding), dim=0)
216 dist.all_gather(tensor_list, tensor)
217
218 data_list = []
219 for size, tensor in zip(size_list, tensor_list):
220 buffer = tensor.cpu().numpy().tobytes()[:size]
221 data_list.append(pickle.loads(buffer))
222
223 return data_list

Callers

nothing calls this directly

Calls 2

get_world_sizeFunction · 0.70
toMethod · 0.45

Tested by

no test coverage detected