Broadcasts the inputs to all ranks. Arguments: inputs : Any objects that can be serialized by pickle. src (int): Source rank. Returns: Each rank returns the same value as src.
(inputs, src)
| 217 | |
| 218 | |
| 219 | def broadcast(inputs, src): |
| 220 | """ |
| 221 | Broadcasts the inputs to all ranks. |
| 222 | |
| 223 | Arguments: |
| 224 | inputs : Any objects that can be serialized by pickle. |
| 225 | src (int): Source rank. |
| 226 | Returns: |
| 227 | Each rank returns the same value as src. |
| 228 | """ |
| 229 | rank = dist.get_rank() |
| 230 | shape_tensor = torch.tensor([0], device='cuda') |
| 231 | |
| 232 | if rank == src: |
| 233 | inputs_tensor = torch.tensor( |
| 234 | bytearray(pickle.dumps(inputs)), dtype=torch.uint8, device='cuda') |
| 235 | shape_tensor = torch.tensor(inputs_tensor.shape, device='cuda') |
| 236 | |
| 237 | dist.barrier() |
| 238 | dist.broadcast(shape_tensor, src) |
| 239 | |
| 240 | if rank != src: |
| 241 | inputs_tensor = torch.full((shape_tensor.item(), ), |
| 242 | 0, |
| 243 | dtype=torch.uint8, |
| 244 | device='cuda') |
| 245 | |
| 246 | dist.barrier() |
| 247 | dist.broadcast(inputs_tensor, src) |
| 248 | |
| 249 | return pickle.loads(inputs_tensor.cpu().numpy().tobytes()) |
| 250 | |
| 251 | |
| 252 | def set_random_seed(seed): |
no test coverage detected
searching dependent graphs…