(self, device=None)
| 666 | return tensor_part, partition |
| 667 | |
| 668 | def full(self, device=None): |
| 669 | if device is None: |
| 670 | device = self.orig_device |
| 671 | |
| 672 | # Allocate the full tensor as a flat buffer. |
| 673 | full_numel = prod(self.full_size()) |
| 674 | flat_tensor = torch.zeros([full_numel], dtype=self.local_data.dtype, device=device) |
| 675 | |
| 676 | # Prepare all-gather buffer |
| 677 | partition_tensors = [] |
| 678 | for part_id in range(self.num_parts): |
| 679 | part_size = self.partition[part_id + 1] - self.partition[part_id] |
| 680 | buf = flat_tensor.narrow(0, start=self.partition[part_id], length=part_size) |
| 681 | if part_id == self.rank: |
| 682 | buf.copy_(self.local_data) |
| 683 | partition_tensors.append(buf) |
| 684 | |
| 685 | # Collect the full tensor |
| 686 | dist.all_gather(partition_tensors, partition_tensors[self.rank], group=self.group) |
| 687 | |
| 688 | for i in range(len(partition_tensors)): |
| 689 | partition_tensors[i].data = torch.zeros(1) |
| 690 | partition_tensors[i] = None |
| 691 | |
| 692 | return flat_tensor.view(self.full_size()).clone().detach() |
| 693 | |
| 694 | def to_meta(self): |
| 695 | """Returns a torch.LongTensor that encodes partitioning information. |
no test coverage detected