-
Notifications
You must be signed in to change notification settings - Fork 412
dynamo: route native NCCL collectives at the op's process (sub)group #4380
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
2c6b32d
a502786
b32b151
dde0553
eac638e
df7925d
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||||||||||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
|
|
@@ -83,6 +83,49 @@ def _get_distributed_rank_and_world_size() -> Tuple[int, int]: | |||||||||||||||||||||
| return rank, world_size | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
||||||||||||||||||||||
| def _collective_group_ranks(group_name: Optional[str], world_size: int) -> np.ndarray: | ||||||||||||||||||||||
| """Ranks that participate in this collective, as IDs in the bound communicator. | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| ``addDistCollective`` documents ``groups`` as "a flat array of rank IDs in the | ||||||||||||||||||||||
| communicator", selecting the subset that takes part and defining each one's group-local | ||||||||||||||||||||||
| rank by position. The runtime binds one communicator per engine, so these are that | ||||||||||||||||||||||
| communicator's rank IDs -- global ranks when the world communicator is bound, which is | ||||||||||||||||||||||
| what lets a single engine host collectives on several different subgroups (e.g. context | ||||||||||||||||||||||
| parallel on one mesh axis and tensor parallel on the other). | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| Resolving the op's ``group_name`` is what lets a collective target a subgroup at all; | ||||||||||||||||||||||
| without it every collective is built over the whole world. | ||||||||||||||||||||||
| """ | ||||||||||||||||||||||
| if group_name: | ||||||||||||||||||||||
| try: | ||||||||||||||||||||||
| import torch.distributed as dist | ||||||||||||||||||||||
| from torch.distributed.distributed_c10d import _resolve_process_group | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| # Preserve the group's own ordering: get_process_group_ranks() returns global | ||||||||||||||||||||||
| # ranks indexed by *group* rank, and that mapping is what layout-sensitive | ||||||||||||||||||||||
| # collectives are defined against. all_gather concatenates by group rank, | ||||||||||||||||||||||
| # reduce_scatter sends chunk i to group rank i, all_to_all permutes by it, and a | ||||||||||||||||||||||
| # scatter/gather root names a position in it. Sorting would silently renumber the | ||||||||||||||||||||||
| # group whenever it was not created in ascending order -- e.g. a device mesh | ||||||||||||||||||||||
| # yielding [5, 2] would be read as [2, 5], swapping which rank receives which | ||||||||||||||||||||||
| # slice. Sorting is not needed for agreement either: every member resolves the | ||||||||||||||||||||||
| # same group_name to the same group and therefore already sees the same order. | ||||||||||||||||||||||
| pg = _resolve_process_group(group_name) | ||||||||||||||||||||||
| except Exception as e: | ||||||||||||||||||||||
| # Do not fall back to the world group here. The op named a group, so falling | ||||||||||||||||||||||
| # back would build a collective spanning every rank -- the exact silent | ||||||||||||||||||||||
| # wrong-results failure this function exists to prevent. Fail the build instead. | ||||||||||||||||||||||
| raise RuntimeError( | ||||||||||||||||||||||
| f"Collective names process group '{group_name}', which could not be " | ||||||||||||||||||||||
| f"resolved in this process ({e}). Refusing to fall back to the world " | ||||||||||||||||||||||
| f"group: that would reduce across every rank and silently return wrong " | ||||||||||||||||||||||
| f"results. Ensure the group is created before compiling." | ||||||||||||||||||||||
| ) from e | ||||||||||||||||||||||
| return np.array(dist.get_process_group_ranks(pg), dtype=np.int64) | ||||||||||||||||||||||
| # No group named: the collective is over the world group by construction. | ||||||||||||||||||||||
| return np.arange(world_size, dtype=np.int64) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
||||||||||||||||||||||
| def nccl_all_gather( | ||||||||||||||||||||||
| ctx: ConversionContext, | ||||||||||||||||||||||
| target: Union[Target, str], | ||||||||||||||||||||||
|
|
@@ -224,13 +267,14 @@ def nccl_reduce_scatter( | |||||||||||||||||||||
| return layer.get_output(0) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
||||||||||||||||||||||
| @needs_native_collectives | ||||||||||||||||||||||
| @needs_native_collectives # type: ignore[misc] | ||||||||||||||||||||||
| def nccl_all_gather_native( | ||||||||||||||||||||||
| ctx: ConversionContext, | ||||||||||||||||||||||
| target: Union[Target, str], | ||||||||||||||||||||||
| source_ir: Optional[SourceIR], | ||||||||||||||||||||||
| name: str, | ||||||||||||||||||||||
| plug_inputs: Tuple[Argument, ...], | ||||||||||||||||||||||
| group_name: Optional[str] = None, | ||||||||||||||||||||||
| ) -> trt.ITensor: | ||||||||||||||||||||||
| """ | ||||||||||||||||||||||
| Implement all_gather using native TensorRT DistCollective API. | ||||||||||||||||||||||
|
|
@@ -263,10 +307,8 @@ def nccl_all_gather_native( | |||||||||||||||||||||
| # Use native TensorRT DistCollective API for ALL_GATHER | ||||||||||||||||||||||
| # For ALL_GATHER, the reduce operation and root rank parameters are ignored | ||||||||||||||||||||||
| # The last parameter (group) can be None to include all ranks | ||||||||||||||||||||||
| import numpy as np | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| # Create array of all participating rank IDs [0, 1, 2, ..., world_size-1] | ||||||||||||||||||||||
| groups = np.arange(world_size, dtype=np.int64) | ||||||||||||||||||||||
| groups = _collective_group_ranks(group_name, world_size) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| logger.debug( | ||||||||||||||||||||||
| f"Creating ALL_GATHER layer: groups={groups.tolist()}, groupSize={world_size}" | ||||||||||||||||||||||
|
|
@@ -288,7 +330,10 @@ def nccl_all_gather_native( | |||||||||||||||||||||
| set_layer_name(layer, target, name, source_ir) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| output = layer.get_output(0) | ||||||||||||||||||||||
| layer.num_ranks = world_size | ||||||||||||||||||||||
| # num_ranks is the size of *this collective's* group, not the world size: TensorRT | ||||||||||||||||||||||
| # takes the participants from `groups` and their count from num_ranks, so a subgroup | ||||||||||||||||||||||
| # collective left at world_size describes a group it does not have. | ||||||||||||||||||||||
| layer.num_ranks = len(groups) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| return output | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
@@ -297,14 +342,15 @@ def nccl_all_gather_native( | |||||||||||||||||||||
| raise | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
||||||||||||||||||||||
| @needs_native_collectives | ||||||||||||||||||||||
| @needs_native_collectives # type: ignore[misc] | ||||||||||||||||||||||
| def nccl_reduce_scatter_native( | ||||||||||||||||||||||
| ctx: ConversionContext, | ||||||||||||||||||||||
| target: Union[Target, str], | ||||||||||||||||||||||
| source_ir: Optional[SourceIR], | ||||||||||||||||||||||
| name: str, | ||||||||||||||||||||||
| plug_inputs: Tuple[Argument, ...], | ||||||||||||||||||||||
| reduce_op: str = "sum", | ||||||||||||||||||||||
| group_name: Optional[str] = None, | ||||||||||||||||||||||
| ) -> trt.ITensor: | ||||||||||||||||||||||
| """ | ||||||||||||||||||||||
| Implement reduce_scatter using native TensorRT DistCollective API. | ||||||||||||||||||||||
|
|
@@ -353,7 +399,7 @@ def nccl_reduce_scatter_native( | |||||||||||||||||||||
| trt_reduce_op = reduce_op_map[reduce_op.lower()] | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| try: | ||||||||||||||||||||||
| groups = np.arange(world_size, dtype=np.int64) | ||||||||||||||||||||||
| groups = _collective_group_ranks(group_name, world_size) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| layer = ctx.net.add_dist_collective( | ||||||||||||||||||||||
| input_tensor, | ||||||||||||||||||||||
|
|
@@ -366,7 +412,10 @@ def nccl_reduce_scatter_native( | |||||||||||||||||||||
| set_layer_name(layer, target, name, source_ir) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| output = layer.get_output(0) | ||||||||||||||||||||||
| layer.num_ranks = world_size | ||||||||||||||||||||||
| # num_ranks is the size of *this collective's* group, not the world size: TensorRT | ||||||||||||||||||||||
| # takes the participants from `groups` and their count from num_ranks, so a subgroup | ||||||||||||||||||||||
| # collective left at world_size describes a group it does not have. | ||||||||||||||||||||||
| layer.num_ranks = len(groups) | ||||||||||||||||||||||
| logger.debug( | ||||||||||||||||||||||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. had a doubt in this. The closed PR #4388 listed this as a fix- ayer.num_ranks must be set before layer.get_output(0) so Myelin can infer the collective's output shape (e.g. all_gather dim0 = in_dim0 * num_ranks). But I don't see that in this PR. Its different from the subgroup routing, but does that need to be the case for collectives producing rank dependant outputs like all_gather, reduce_scatter, scatter, gather
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Good question — I dropped it deliberately, and only after testing it, because the reasoning in #4388 doesn't hold up. The shape is stale at the instant of capture: for an all_gather with
I deliberately included the rank-dependent ops you named — What was substantive in that hunk is the value: So: ordering dropped as a no-op, value kept. Happy to restore the ordering as defensive hygiene if you'd prefer it, but I didn't want to carry it as a "fix" when the evidence says it changes nothing. |
||||||||||||||||||||||
| f"Successfully created native REDUCE_SCATTER layer: {name}, reduce_op={reduce_op}, groups={groups.tolist()}" | ||||||||||||||||||||||
| ) | ||||||||||||||||||||||
|
|
@@ -378,14 +427,15 @@ def nccl_reduce_scatter_native( | |||||||||||||||||||||
| raise | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
||||||||||||||||||||||
| @needs_native_collectives | ||||||||||||||||||||||
| @needs_native_collectives # type: ignore[misc] | ||||||||||||||||||||||
| def nccl_all_reduce_native( | ||||||||||||||||||||||
| ctx: ConversionContext, | ||||||||||||||||||||||
| target: Union[Target, str], | ||||||||||||||||||||||
| source_ir: Optional[SourceIR], | ||||||||||||||||||||||
| name: str, | ||||||||||||||||||||||
| plug_inputs: Tuple[Argument, ...], | ||||||||||||||||||||||
| reduce_op: str = "sum", | ||||||||||||||||||||||
| group_name: Optional[str] = None, | ||||||||||||||||||||||
| ) -> trt.ITensor: | ||||||||||||||||||||||
| """ | ||||||||||||||||||||||
| Implement all_reduce using native TensorRT DistCollective API. | ||||||||||||||||||||||
|
|
@@ -438,7 +488,7 @@ def nccl_all_reduce_native( | |||||||||||||||||||||
| # Create array of all participating rank IDs [0, 1, ..., world_size-1] | ||||||||||||||||||||||
| # Passing None for groups can be treated as a no-op by TRT; use an explicit | ||||||||||||||||||||||
| # rank array (same as ALL_GATHER) to ensure the reduction is performed. | ||||||||||||||||||||||
| groups = np.arange(world_size, dtype=np.int64) | ||||||||||||||||||||||
| groups = _collective_group_ranks(group_name, world_size) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| layer = ctx.net.add_dist_collective( | ||||||||||||||||||||||
| input_tensor, | ||||||||||||||||||||||
|
|
@@ -451,7 +501,10 @@ def nccl_all_reduce_native( | |||||||||||||||||||||
| set_layer_name(layer, target, name, source_ir) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| output = layer.get_output(0) | ||||||||||||||||||||||
| layer.num_ranks = world_size | ||||||||||||||||||||||
| # num_ranks is the size of *this collective's* group, not the world size: TensorRT | ||||||||||||||||||||||
| # takes the participants from `groups` and their count from num_ranks, so a subgroup | ||||||||||||||||||||||
| # collective left at world_size describes a group it does not have. | ||||||||||||||||||||||
| layer.num_ranks = len(groups) | ||||||||||||||||||||||
| logger.debug( | ||||||||||||||||||||||
| f"Successfully created native ALL_REDUCE layer: {name}, reduce_op={reduce_op}, groups={groups.tolist()}" | ||||||||||||||||||||||
| ) | ||||||||||||||||||||||
|
|
@@ -463,13 +516,14 @@ def nccl_all_reduce_native( | |||||||||||||||||||||
| raise | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
||||||||||||||||||||||
| @needs_native_collectives | ||||||||||||||||||||||
| @needs_native_collectives # type: ignore[misc] | ||||||||||||||||||||||
| def nccl_all_to_all_native( | ||||||||||||||||||||||
| ctx: ConversionContext, | ||||||||||||||||||||||
| target: Union[Target, str], | ||||||||||||||||||||||
| source_ir: Optional[SourceIR], | ||||||||||||||||||||||
| name: str, | ||||||||||||||||||||||
| plug_inputs: Tuple[Argument, ...], | ||||||||||||||||||||||
| group_name: Optional[str] = None, | ||||||||||||||||||||||
| ) -> trt.ITensor: | ||||||||||||||||||||||
| """ | ||||||||||||||||||||||
| Implement all_to_all using native TensorRT DistCollective API. | ||||||||||||||||||||||
|
|
@@ -503,10 +557,8 @@ def nccl_all_to_all_native( | |||||||||||||||||||||
| # Use native TensorRT DistCollective API for ALL_TO_ALL | ||||||||||||||||||||||
| # For ALL_TO_ALL, the reduce operation and root rank parameters are ignored | ||||||||||||||||||||||
| # The last parameter (group) can be None to include all ranks | ||||||||||||||||||||||
| import numpy as np | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| # Create array of all participating rank IDs [0, 1, 2, ..., world_size-1] | ||||||||||||||||||||||
| groups = np.arange(world_size, dtype=np.int64) | ||||||||||||||||||||||
| groups = _collective_group_ranks(group_name, world_size) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| logger.debug( | ||||||||||||||||||||||
| f"Creating ALL_TO_ALL layer: groups={groups.tolist()}, groupSize={world_size}" | ||||||||||||||||||||||
|
|
@@ -528,7 +580,10 @@ def nccl_all_to_all_native( | |||||||||||||||||||||
| set_layer_name(layer, target, name, source_ir) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| output = layer.get_output(0) | ||||||||||||||||||||||
| layer.num_ranks = world_size | ||||||||||||||||||||||
| # num_ranks is the size of *this collective's* group, not the world size: TensorRT | ||||||||||||||||||||||
| # takes the participants from `groups` and their count from num_ranks, so a subgroup | ||||||||||||||||||||||
| # collective left at world_size describes a group it does not have. | ||||||||||||||||||||||
| layer.num_ranks = len(groups) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| return output | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
@@ -537,14 +592,15 @@ def nccl_all_to_all_native( | |||||||||||||||||||||
| raise | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
||||||||||||||||||||||
| @needs_native_collectives | ||||||||||||||||||||||
| @needs_native_collectives # type: ignore[misc] | ||||||||||||||||||||||
| def nccl_scatter_native( | ||||||||||||||||||||||
| ctx: ConversionContext, | ||||||||||||||||||||||
| target: Union[Target, str], | ||||||||||||||||||||||
| source_ir: Optional[SourceIR], | ||||||||||||||||||||||
| name: str, | ||||||||||||||||||||||
| plug_inputs: Tuple[Argument, ...], | ||||||||||||||||||||||
| root: int = 0, | ||||||||||||||||||||||
| group_name: Optional[str] = None, | ||||||||||||||||||||||
| ) -> trt.ITensor: | ||||||||||||||||||||||
| """ | ||||||||||||||||||||||
| Implement scatter using native TensorRT DistCollective API. | ||||||||||||||||||||||
|
|
@@ -577,10 +633,8 @@ def nccl_scatter_native( | |||||||||||||||||||||
| # Use native TensorRT DistCollective API for SCATTER | ||||||||||||||||||||||
| # For SCATTER, the reduce operation parameter is ignored | ||||||||||||||||||||||
| # The last parameter (group) can be None to include all ranks | ||||||||||||||||||||||
| import numpy as np | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| # Create array of all participating rank IDs [0, 1, 2, ..., world_size-1] | ||||||||||||||||||||||
| groups = np.arange(world_size, dtype=np.int64) | ||||||||||||||||||||||
| groups = _collective_group_ranks(group_name, world_size) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| logger.debug( | ||||||||||||||||||||||
| f"Creating scatter layer: groups={groups.tolist()}, groupSize={world_size}" | ||||||||||||||||||||||
|
|
@@ -602,7 +656,10 @@ def nccl_scatter_native( | |||||||||||||||||||||
| set_layer_name(layer, target, name, source_ir) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| output = layer.get_output(0) | ||||||||||||||||||||||
| layer.num_ranks = world_size | ||||||||||||||||||||||
| # num_ranks is the size of *this collective's* group, not the world size: TensorRT | ||||||||||||||||||||||
| # takes the participants from `groups` and their count from num_ranks, so a subgroup | ||||||||||||||||||||||
| # collective left at world_size describes a group it does not have. | ||||||||||||||||||||||
| layer.num_ranks = len(groups) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| return output | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
@@ -611,14 +668,15 @@ def nccl_scatter_native( | |||||||||||||||||||||
| raise | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
||||||||||||||||||||||
| @needs_native_collectives | ||||||||||||||||||||||
| @needs_native_collectives # type: ignore[misc] | ||||||||||||||||||||||
| def nccl_gather_native( | ||||||||||||||||||||||
| ctx: ConversionContext, | ||||||||||||||||||||||
| target: Union[Target, str], | ||||||||||||||||||||||
| source_ir: Optional[SourceIR], | ||||||||||||||||||||||
| name: str, | ||||||||||||||||||||||
| plug_inputs: Tuple[Argument, ...], | ||||||||||||||||||||||
| root: int = 0, | ||||||||||||||||||||||
| group_name: Optional[str] = None, | ||||||||||||||||||||||
| ) -> trt.ITensor: | ||||||||||||||||||||||
| """ | ||||||||||||||||||||||
| Implement gather using native TensorRT DistCollective API. | ||||||||||||||||||||||
|
|
@@ -651,10 +709,8 @@ def nccl_gather_native( | |||||||||||||||||||||
| # Use native TensorRT DistCollective API for GATHER | ||||||||||||||||||||||
| # For GATHER, the reduce operation parameter is ignored | ||||||||||||||||||||||
| # The last parameter (group) can be None to include all ranks | ||||||||||||||||||||||
| import numpy as np | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| # Create array of all participating rank IDs [0, 1, 2, ..., world_size-1] | ||||||||||||||||||||||
| groups = np.arange(world_size, dtype=np.int64) | ||||||||||||||||||||||
| groups = _collective_group_ranks(group_name, world_size) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| logger.debug( | ||||||||||||||||||||||
| f"Creating gather layer: groups={groups.tolist()}, groupSize={world_size}" | ||||||||||||||||||||||
|
|
@@ -676,7 +732,10 @@ def nccl_gather_native( | |||||||||||||||||||||
| set_layer_name(layer, target, name, source_ir) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| output = layer.get_output(0) | ||||||||||||||||||||||
| layer.num_ranks = world_size | ||||||||||||||||||||||
| # num_ranks is the size of *this collective's* group, not the world size: TensorRT | ||||||||||||||||||||||
| # takes the participants from `groups` and their count from num_ranks, so a subgroup | ||||||||||||||||||||||
| # collective left at world_size describes a group it does not have. | ||||||||||||||||||||||
| layer.num_ranks = len(groups) | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
| return output | ||||||||||||||||||||||
|
|
||||||||||||||||||||||
|
|
||||||||||||||||||||||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Silent world group fallback might not be the right approach here? The way I see it is this can fail to find the ranks in two cases-
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Agreed, and changed in
960c6ca. Falling back to the world group is precisely the failure this function exists to prevent: the op asked for a subgroup, and we'd silently build a collective over every rank and return a wrong answer with no error. That's worse than not compiling.It now distinguishes your two cases:
group_name— the collective is over the world group by construction, soarange(world_size)is correct, not a fallback.group_namegiven but unresolvable — raisesRuntimeErrornaming the group, rather than widening.TestCollectiveGroupRankswas asserting the old fallback behaviour, so I replaced that case with one asserting it raises, plus a separate test that a missinggroup_namestill resolves to the world group.