Add all_to_all_single to GroupCoordinator (#27492)
Co-authored-by: Ming Yang <minos.future@gmail.com>
This commit is contained in:
co-authored by
Ming Yang
parent
5be0b0c8c0
commit
db58e76c33
@@ -196,6 +196,17 @@ def reg_reduce_scatter_tensor(
|
|||||||
group._reduce_scatter_tensor(output, input)
|
group._reduce_scatter_tensor(output, input)
|
||||||
|
|
||||||
|
|
||||||
|
@register_custom_op(mutates_args=["output"])
|
||||||
|
def reg_all_to_all_single(
|
||||||
|
output: torch.Tensor, input: torch.Tensor, group_name: str
|
||||||
|
) -> None:
|
||||||
|
assert group_name in _groups, f"Group {group_name} is not found."
|
||||||
|
group = _groups[group_name]()
|
||||||
|
if group is None:
|
||||||
|
raise ValueError(f"Group {group_name} is destroyed.")
|
||||||
|
group._all_to_all_single(output, input)
|
||||||
|
|
||||||
|
|
||||||
class GroupCoordinator:
|
class GroupCoordinator:
|
||||||
"""
|
"""
|
||||||
PyTorch ProcessGroup wrapper for a group of processes.
|
PyTorch ProcessGroup wrapper for a group of processes.
|
||||||
@@ -776,6 +787,15 @@ class GroupCoordinator:
|
|||||||
else:
|
else:
|
||||||
reg_reduce_scatter_tensor(output, input, group_name=self.unique_name)
|
reg_reduce_scatter_tensor(output, input, group_name=self.unique_name)
|
||||||
|
|
||||||
|
def _all_to_all_single(self, output: torch.Tensor, input: torch.Tensor) -> None:
|
||||||
|
torch.distributed.all_to_all_single(output, input, group=self.device_group)
|
||||||
|
|
||||||
|
def all_to_all_single(self, output: torch.Tensor, input: torch.Tensor):
|
||||||
|
if self.world_size == 1:
|
||||||
|
output.copy_(input)
|
||||||
|
return
|
||||||
|
reg_all_to_all_single(output, input, group_name=self.unique_name)
|
||||||
|
|
||||||
def reduce_scatter(
|
def reduce_scatter(
|
||||||
self,
|
self,
|
||||||
output: torch.Tensor,
|
output: torch.Tensor,
|
||||||
|
|||||||
Reference in New Issue
Block a user