(self, sparse, dp_group)
| 2348 | return sparse_list |
| 2349 | |
| 2350 | def sparse_allreduce(self, sparse, dp_group): |
| 2351 | original_data_type = sparse.values.dtype |
| 2352 | if self.communication_data_type != sparse.values.dtype: |
| 2353 | if self.communication_data_type in (torch.float16, torch.bfloat16): |
| 2354 | indices = sparse.indices.to(torch.int32) |
| 2355 | else: |
| 2356 | indices = sparse.indices |
| 2357 | values = sparse.values.to(self.communication_data_type) |
| 2358 | else: |
| 2359 | indices = sparse.indices |
| 2360 | values = sparse.values |
| 2361 | |
| 2362 | if self.postscale_gradients(): |
| 2363 | if self.gradient_average: |
| 2364 | values.mul_(self.gradient_predivide_factor() / dist.get_world_size(group=dp_group)) |
| 2365 | else: |
| 2366 | values.mul_(1. / dist.get_world_size(group=dp_group)) |
| 2367 | |
| 2368 | indices_device_list = self.sparse_all_gather(indices, dp_group) |
| 2369 | values_device_list = self.sparse_all_gather(values, dp_group) |
| 2370 | |
| 2371 | sparse.indices = torch.cat(indices_device_list).to(torch.long) |
| 2372 | sparse.values = torch.cat(values_device_list).to(original_data_type) |
| 2373 | return sparse |
| 2374 | |
| 2375 | def sparse_all_gather(self, value, dp_group): |
| 2376 | my_size = torch.LongTensor([value.size()[0]]).to(self.device) |
no test coverage detected