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)
| 176 | |
| 177 | |
| 178 | def 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 |
nothing calls this directly
no test coverage detected