(self, value, dp_group)
| 2373 | return sparse |
| 2374 | |
| 2375 | def sparse_all_gather(self, value, dp_group): |
| 2376 | my_size = torch.LongTensor([value.size()[0]]).to(self.device) |
| 2377 | all_sizes = self.all_gather_scalar(my_size, dp_group) |
| 2378 | max_size = torch.cat(all_sizes).max() |
| 2379 | fill_size = max_size - my_size |
| 2380 | |
| 2381 | assert value.dim() in [1, 2] |
| 2382 | if value.dim() == 1: |
| 2383 | if fill_size > 0: |
| 2384 | value = torch.cat([value, value.new_empty(fill_size)]) |
| 2385 | tensor_list = [value.new_empty(max_size) for _ in range(dist.get_world_size(group=dp_group))] |
| 2386 | else: |
| 2387 | if fill_size > 0: |
| 2388 | value = torch.cat([value, value.new_empty(fill_size, value.size()[1])]) |
| 2389 | tensor_list = [ |
| 2390 | value.new_empty(max_size, |
| 2391 | value.size()[1]) for _ in range(dist.get_world_size(group=dp_group)) |
| 2392 | ] |
| 2393 | |
| 2394 | dist.all_gather(tensor_list, value, group=dp_group) |
| 2395 | tensors = [] |
| 2396 | for dev_idx, t in enumerate(tensor_list): |
| 2397 | size = all_sizes[dev_idx][0] |
| 2398 | tensors.append(t.index_select(0, torch.arange(size, dtype=torch.long, device=self.device))) |
| 2399 | |
| 2400 | return tensors |
| 2401 | |
| 2402 | def all_gather_scalar(self, value, dp_group): |
| 2403 | tensor_list = [value.new_zeros(value.size()) for _ in range(dist.get_world_size(group=dp_group))] |
no test coverage detected