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 fro
(data, group=None)
| 154 | |
| 155 | |
| 156 | def all_gather(data, group=None): |
| 157 | """ |
| 158 | Run all_gather on arbitrary picklable data (not necessarily tensors). |
| 159 | |
| 160 | Args: |
| 161 | data: any picklable object |
| 162 | group: a torch process group. By default, will use a group which |
| 163 | contains all ranks on gloo backend. |
| 164 | Returns: |
| 165 | list[data]: list of data gathered from each rank |
| 166 | """ |
| 167 | if get_world_size() == 1: |
| 168 | return [data] |
| 169 | if group is None: |
| 170 | group = _get_global_gloo_group() |
| 171 | if dist.get_world_size(group) == 1: |
| 172 | return [data] |
| 173 | |
| 174 | tensor = _serialize_to_tensor(data, group) |
| 175 | |
| 176 | size_list, tensor = _pad_to_largest_tensor(tensor, group) |
| 177 | max_size = max(size_list) |
| 178 | |
| 179 | # receiving Tensor from all ranks |
| 180 | tensor_list = [ |
| 181 | torch.empty((max_size,), dtype=torch.uint8, device=tensor.device) |
| 182 | for _ in size_list |
| 183 | ] |
| 184 | dist.all_gather(tensor_list, tensor, group=group) |
| 185 | |
| 186 | data_list = [] |
| 187 | for size, tensor in zip(size_list, tensor_list): |
| 188 | buffer = tensor.cpu().numpy().tobytes()[:size] |
| 189 | data_list.append(pickle.loads(buffer)) |
| 190 | |
| 191 | return data_list |
| 192 | |
| 193 | |
| 194 | def gather(data, dst=0, group=None): |
no test coverage detected