(self, sparse, dp_group, dp_world_size=None)
| 3757 | return sparse_list |
| 3758 | |
| 3759 | def sparse_allreduce(self, sparse, dp_group, dp_world_size=None): |
| 3760 | original_data_type = sparse.values.dtype |
| 3761 | if self.communication_data_type != sparse.values.dtype: |
| 3762 | if self.communication_data_type in (torch.float16, torch.bfloat16): |
| 3763 | indices = sparse.indices.to(torch.int32) |
| 3764 | else: |
| 3765 | indices = sparse.indices |
| 3766 | values = sparse.values.to(self.communication_data_type) |
| 3767 | else: |
| 3768 | indices = sparse.indices |
| 3769 | values = sparse.values |
| 3770 | |
| 3771 | if dp_world_size is None: |
| 3772 | dp_world_size = dist.get_world_size(group=dp_group) |
| 3773 | if self.postscale_gradients(): |
| 3774 | if self.gradient_average: |
| 3775 | values.mul_(self.gradient_predivide_factor() / (dp_world_size)) |
| 3776 | else: |
| 3777 | values.mul_(1. / (dp_world_size)) |
| 3778 | |
| 3779 | indices_device_list = self.sparse_all_gather(indices, dp_group) |
| 3780 | values_device_list = self.sparse_all_gather(values, dp_group) |
| 3781 | |
| 3782 | sparse.indices = torch.cat(indices_device_list).to(torch.long) |
| 3783 | sparse.values = torch.cat(values_device_list).to(original_data_type) |
| 3784 | return sparse |
| 3785 | |
| 3786 | def sparse_all_gather(self, value, dp_group): |
| 3787 | my_size = torch.LongTensor([value.size()[0]]).to(self.device) |
no test coverage detected