MCPcopy Create free account
hub / github.com/AIS-SNU/Smart-Infinity / sparse_allreduce

Method sparse_allreduce

deepspeed/runtime/engine.py:2350–2373  ·  view source on GitHub ↗
(self, sparse, dp_group)

Source from the content-addressed store, hash-verified

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)

Callers 1

Calls 4

postscale_gradientsMethod · 0.95
sparse_all_gatherMethod · 0.95
get_world_sizeMethod · 0.80

Tested by

no test coverage detected