Run all_gather on arbitrary picklable data (not necessarily tensors). Args: data: any picklable object group: a torch process group. By default, will use a group which contains all ranks on gloo backend. Returns: list[data]: list of data gathered from
(data, group=None)
| 321 | |
| 322 | |
| 323 | def all_gather(data, group=None): |
| 324 | """ |
| 325 | Run all_gather on arbitrary picklable data (not necessarily tensors). |
| 326 | Args: |
| 327 | data: any picklable object |
| 328 | group: a torch process group. By default, will use a group which |
| 329 | contains all ranks on gloo backend. |
| 330 | Returns: |
| 331 | list[data]: list of data gathered from each rank |
| 332 | """ |
| 333 | if get_world_size() == 1: |
| 334 | return [data] |
| 335 | if group is None: |
| 336 | group = _get_global_gloo_group() |
| 337 | if dist.get_world_size(group) == 1: |
| 338 | return [data] |
| 339 | |
| 340 | tensor = _serialize_to_tensor(data, group) |
| 341 | |
| 342 | size_list, tensor = _pad_to_largest_tensor(tensor, group) |
| 343 | max_size = max(size_list) |
| 344 | |
| 345 | # receiving Tensor from all ranks |
| 346 | tensor_list = [ |
| 347 | torch.empty((max_size, ), dtype=torch.uint8, device=tensor.device) |
| 348 | for _ in size_list |
| 349 | ] |
| 350 | dist.all_gather(tensor_list, tensor, group=group) |
| 351 | |
| 352 | data_list = [] |
| 353 | for size, tensor in zip(size_list, tensor_list): |
| 354 | buffer = tensor.cpu().numpy().tobytes()[:size] |
| 355 | data_list.append(pickle.loads(buffer)) |
| 356 | |
| 357 | return data_list |
| 358 | |
| 359 | |
| 360 | def is_on_same_device(model: torch.nn.Module) -> bool: |
no test coverage detected
searching dependent graphs…