Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 18 additions & 2 deletions py/torch_tensorrt/dynamo/conversion/custom_ops_converters.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,7 @@ def fused_nccl_all_gather(
SourceIR.ATEN,
name,
[args[0]],
group_name=args[2] if len(args) > 2 else None,
)

@dynamo_tensorrt_converter(
Expand All @@ -138,6 +139,7 @@ def fused_nccl_reduce_scatter(
name,
[args[0]],
reduce_op=reduce_op,
group_name=args[3] if len(args) > 3 else None,
)

@dynamo_tensorrt_converter(
Expand All @@ -161,6 +163,7 @@ def fused_nccl_all_reduce(
name,
[args[0]],
reduce_op=reduce_op,
group_name=args[2] if len(args) > 2 else None,
)

@dynamo_tensorrt_converter(
Expand All @@ -182,6 +185,7 @@ def fused_nccl_all_to_all(
SourceIR.ATEN,
name,
[args[0]],
group_name=args[3] if len(args) > 3 else None,
)

@dynamo_tensorrt_converter(
Expand All @@ -199,7 +203,13 @@ def fused_nccl_scatter(
"""Scatter using native TensorRT DistCollective API."""
root = args[1] if len(args) > 1 else 0
return impl.nccl_ops.nccl_scatter_native(
ctx, target, SourceIR.ATEN, name, [args[0]], root=root
ctx,
target,
SourceIR.ATEN,
name,
[args[0]],
root=root,
group_name=args[2] if len(args) > 2 else None,
)

@dynamo_tensorrt_converter(
Expand All @@ -217,7 +227,13 @@ def fused_nccl_gather(
"""Gather using native TensorRT DistCollective API."""
root = args[1] if len(args) > 1 else 0
return impl.nccl_ops.nccl_gather_native(
ctx, target, SourceIR.ATEN, name, [args[0]], root=root
ctx,
target,
SourceIR.ATEN,
name,
[args[0]],
root=root,
group_name=args[2] if len(args) > 2 else None,
)


Expand Down
111 changes: 85 additions & 26 deletions py/torch_tensorrt/dynamo/conversion/impl/nccl_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Copy link
Copy Markdown
Collaborator

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-

  1. torch.distributed is not initialized
  2. distributed is initialized, group_name was provided but torch could not resolve it. I dont think we should silently fallback to world here

Copy link
Copy Markdown
Author

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:

  • No group_name — the collective is over the world group by construction, so arange(world_size) is correct, not a fallback.
  • group_name given but unresolvable — raises RuntimeError naming the group, rather than widening.
Collective names process group '<name>', which could not be resolved in this process (...).
Refusing to fall back to the world group: that would reduce across every rank and silently
return wrong results. Ensure the group is created before compiling.

TestCollectiveGroupRanks was asserting the old fallback behaviour, so I replaced that case with one asserting it raises, plus a separate test that a missing group_name still resolves to the world group.


def nccl_all_gather(
ctx: ConversionContext,
target: Union[Target, str],
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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}"
Expand All @@ -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

Expand All @@ -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.
Expand Down Expand Up @@ -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,
Expand All @@ -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(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The 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

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The 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 num_ranks=4 and input (2,64), reading get_output(0).shape before setting num_ranks gives (2,64) instead of (8,64). So the concern was real. But TensorRT re-evaluates that shape lazily, on main the two lines are adjacent with nothing reading the shape in between, and the built engine is byte-identical either way. On 8× A100 / TRT 11.2.0.86, building the same network both orders and hashing the serialized engine:

op num_ranks first get_output first identical
ALL_GATHER (16,64) md5 0ebe6c0b97d4 (16,64) md5 0ebe6c0b97d4 yes
ALL_REDUCE (4,64) md5 18e50672a18a (4,64) md5 18e50672a18a yes
ALL_TO_ALL (4,64) md5 a0ad8806638d (4,64) md5 a0ad8806638d yes
REDUCE_SCATTER (1,64) md5 547445add902 (1,64) md5 547445add902 yes

I deliberately included the rank-dependent ops you named — all_gather and reduce_scatter are the two whose output shape actually scales with num_ranks, and they come out identical too.

What was substantive in that hunk is the value: num_ranks has to be the size of the collective's own group, not world_size. That part is kept here (layer.num_ranks = len(groups)), and it's inseparable from subgroup routing — with groups a subgroup, world_size describes a group the collective doesn't have. TestNativeCollectiveNumRanks covers it for all six converters and fails if you put world_size back.

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()}"
)
Expand All @@ -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.
Expand Down Expand Up @@ -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,
Expand All @@ -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()}"
)
Expand All @@ -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.
Expand Down Expand Up @@ -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}"
Expand All @@ -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

Expand All @@ -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.
Expand Down Expand Up @@ -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}"
Expand All @@ -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

Expand All @@ -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.
Expand Down Expand Up @@ -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}"
Expand All @@ -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

Expand Down
Loading
Loading