MCPcopy Create free account
hub / github.com/deepspeedai/DeepSpeed / sparse_allreduce

Method sparse_allreduce

deepspeed/runtime/engine.py:3759–3784  ·  view source on GitHub ↗
(self, sparse, dp_group, dp_world_size=None)

Source from the content-addressed store, hash-verified

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)

Callers 1

Calls 5

postscale_gradientsMethod · 0.95
sparse_all_gatherMethod · 0.95
get_world_sizeMethod · 0.80
toMethod · 0.45

Tested by

no test coverage detected