Opposite of above function, gather values from model parallel ranks.
(tensor)
| 111 | |
| 112 | |
| 113 | def gather_split_1d_tensor(tensor): |
| 114 | """Opposite of above function, gather values from model parallel ranks.""" |
| 115 | world_size = get_model_parallel_world_size() |
| 116 | numel = torch.numel(tensor) |
| 117 | numel_gathered = world_size * numel |
| 118 | gathered = torch.empty(numel_gathered, dtype=tensor.dtype, |
| 119 | device=torch.cuda.current_device(), |
| 120 | requires_grad=False) |
| 121 | chunks = [gathered[i*numel:(i+1)*numel] for i in range(world_size)] |
| 122 | torch.distributed.all_gather(chunks, tensor, |
| 123 | group=get_model_parallel_group()) |
| 124 | return gathered |
| 125 | |
| 126 | |
| 127 | class CudaRNGStatesTracker: |
no test coverage detected