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

Method sparse_all_gather

deepspeed/runtime/engine.py:2375–2400  ·  view source on GitHub ↗
(self, value, dp_group)

Source from the content-addressed store, hash-verified

2373 return sparse
2374
2375 def sparse_all_gather(self, value, dp_group):
2376 my_size = torch.LongTensor([value.size()[0]]).to(self.device)
2377 all_sizes = self.all_gather_scalar(my_size, dp_group)
2378 max_size = torch.cat(all_sizes).max()
2379 fill_size = max_size - my_size
2380
2381 assert value.dim() in [1, 2]
2382 if value.dim() == 1:
2383 if fill_size > 0:
2384 value = torch.cat([value, value.new_empty(fill_size)])
2385 tensor_list = [value.new_empty(max_size) for _ in range(dist.get_world_size(group=dp_group))]
2386 else:
2387 if fill_size > 0:
2388 value = torch.cat([value, value.new_empty(fill_size, value.size()[1])])
2389 tensor_list = [
2390 value.new_empty(max_size,
2391 value.size()[1]) for _ in range(dist.get_world_size(group=dp_group))
2392 ]
2393
2394 dist.all_gather(tensor_list, value, group=dp_group)
2395 tensors = []
2396 for dev_idx, t in enumerate(tensor_list):
2397 size = all_sizes[dev_idx][0]
2398 tensors.append(t.index_select(0, torch.arange(size, dtype=torch.long, device=self.device)))
2399
2400 return tensors
2401
2402 def all_gather_scalar(self, value, dp_group):
2403 tensor_list = [value.new_zeros(value.size()) for _ in range(dist.get_world_size(group=dp_group))]

Callers 1

sparse_allreduceMethod · 0.95

Calls 6

all_gather_scalarMethod · 0.95
get_world_sizeMethod · 0.80
LongTensorMethod · 0.45
sizeMethod · 0.45
all_gatherMethod · 0.45
appendMethod · 0.45

Tested by

no test coverage detected