From c0cfc04ae3356ae4024f7b077cef69cf17b720df Mon Sep 17 00:00:00 2001 From: Tailai Ma Date: Mon, 24 Aug 2026 17:55:15 +0900 Subject: [PATCH] Support logical CP groups with parent P2P transport Signed-off-by: Tailai Ma --- tests/pytorch/test_cuda_graphs.py | 67 +++++ .../dot_product_attention/backends.py | 9 +- .../dot_product_attention/context_parallel.py | 250 ++++++++++++++---- .../dot_product_attention.py | 5 +- .../pytorch/attention/multi_head_attention.py | 3 +- transformer_engine/pytorch/distributed.py | 39 ++- 6 files changed, 314 insertions(+), 59 deletions(-) diff --git a/tests/pytorch/test_cuda_graphs.py b/tests/pytorch/test_cuda_graphs.py index 1b9e11792e..3485c07fae 100644 --- a/tests/pytorch/test_cuda_graphs.py +++ b/tests/pytorch/test_cuda_graphs.py @@ -2,6 +2,8 @@ # # See LICENSE for license information. +import gc +import weakref from typing import Callable, Dict, Iterable, List, Tuple, Union import pytest @@ -22,6 +24,16 @@ is_bf16_available, ) from transformer_engine.pytorch.quantization import FP8GlobalStateManager +from transformer_engine.pytorch.attention.dot_product_attention.context_parallel import ( + _get_cp_p2p_transport_group, + set_cp_p2p_transport_group, +) +from transformer_engine.pytorch.distributed import ( + get_distributed_group_ranks, + get_distributed_rank, + get_distributed_world_size, + is_logical_process_group, +) import transformer_engine.pytorch.ops as te_ops from transformer_engine.common import recipe from utils import ModelConfig, reset_rng_states @@ -39,6 +51,61 @@ } +def test_cp_p2p_transport_group_override(): + class Group: + pass + + logical_group = Group() + transport_group = Group() + + assert _get_cp_p2p_transport_group(logical_group) == (logical_group, False) + set_cp_p2p_transport_group(logical_group, transport_group) + assert _get_cp_p2p_transport_group(logical_group) == (transport_group, True) + set_cp_p2p_transport_group(logical_group, None) + assert _get_cp_p2p_transport_group(logical_group) == (logical_group, False) + + set_cp_p2p_transport_group(logical_group, transport_group) + logical_group_ref = weakref.ref(logical_group) + del logical_group + gc.collect() + assert logical_group_ref() is None + + self_transport_group = Group() + self_transport_group_ref = weakref.ref(self_transport_group) + set_cp_p2p_transport_group(self_transport_group, self_transport_group) + del self_transport_group + gc.collect() + assert self_transport_group_ref() is None + + +def test_logical_cp_group_uses_registered_parent_transport(): + class LogicalGroup: + ranks = (2, 3, 6, 7) + cp_size = 4 + cp_rank = 2 + + class ParentGroup: + pass + + logical_group = LogicalGroup() + parent_group = ParentGroup() + + assert is_logical_process_group(logical_group) + assert get_distributed_world_size(logical_group) == 4 + assert get_distributed_rank(logical_group) == 2 + assert get_distributed_group_ranks(logical_group) == (2, 3, 6, 7) + with pytest.raises(RuntimeError, match="requires a registered parent"): + _get_cp_p2p_transport_group(logical_group) + + set_cp_p2p_transport_group(logical_group, parent_group) + assert _get_cp_p2p_transport_group(logical_group) == (parent_group, True) + + logical_group_ref = weakref.ref(logical_group) + del logical_group + gc.collect() + assert logical_group_ref() is None + + def nvfp4_vanilla(): nvfp4_recipe = recipe.NVFP4BlockScaling() nvfp4_recipe.fp4_quant_fwd_inp = recipe.QParams() diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index a6a8b0b26a..e271b1ad0b 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -48,7 +48,10 @@ META_QKV, ) from transformer_engine.pytorch.quantization import get_fp8_torch_dtype, FP8GlobalStateManager -from transformer_engine.pytorch.distributed import get_distributed_world_size +from transformer_engine.pytorch.distributed import ( + get_distributed_world_size, + is_logical_process_group, +) from transformer_engine.pytorch.jit import no_torch_dynamo from transformer_engine.pytorch.attention.dot_product_attention.context_parallel import ( attn_forward_func_with_cp, @@ -757,7 +760,7 @@ def forward( ), f"FlashAttention does not support qkv_layout = {qkv_layout}!" cp_size = 1 - if isinstance(cp_group, dist_group_type): + if isinstance(cp_group, dist_group_type) or is_logical_process_group(cp_group): cp_size = get_distributed_world_size(cp_group) elif isinstance(cp_group, list): for group in cp_group: @@ -1828,7 +1831,7 @@ def forward( ), f"FusedAttention does not support qkv_layout = {qkv_layout}!" cp_size = 1 - if isinstance(cp_group, dist_group_type): + if isinstance(cp_group, dist_group_type) or is_logical_process_group(cp_group): cp_size = get_distributed_world_size(cp_group) elif isinstance(cp_group, list): for group in cp_group: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 030b1d9cdc..5dbdf41923 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -3,9 +3,11 @@ # See LICENSE for license information. """Context Parallelism.""" -import os + import itertools -from typing import List, Union, Tuple +import os +import weakref +from typing import List, Tuple, Union import torch import transformer_engine_torch as tex @@ -30,9 +32,11 @@ TE_DType, ) from transformer_engine.pytorch.distributed import ( + get_distributed_group_ranks, get_distributed_world_size, get_distributed_rank, gather_along_first_dim, + is_logical_process_group, reduce_scatter_along_first_dim, ) @@ -54,6 +58,28 @@ _seq_chunk_ids_cache_for_reordering_before_attn = {} _seq_chunk_ids_cache_for_reordering_after_attn = {} _softmax_offset_chunk_ids_cache = {} +_cp_p2p_transport_groups = weakref.WeakKeyDictionary() + + +def set_cp_p2p_transport_group(cp_group, transport_group): + """Override only the P2P transport group for a logical CP group.""" + if transport_group is None: + _cp_p2p_transport_groups.pop(cp_group, None) + return + _cp_p2p_transport_groups[cp_group] = weakref.ref(transport_group) + + +def _get_cp_p2p_transport_group(cp_group): + transport_group_ref = _cp_p2p_transport_groups.get(cp_group) + if transport_group_ref is not None: + transport_group = transport_group_ref() + if transport_group is not None: + return transport_group, True + _cp_p2p_transport_groups.pop(cp_group, None) + if is_logical_process_group(cp_group): + raise RuntimeError("A logical CP group requires a registered parent P2P transport group.") + return cp_group, False + # Float8CurrentScaling: fused_attn_bwd takes O in FP8 by default, this flag allows it in F16 _dpa_fp8_cs_o_in_f16 = os.getenv("NVTE_DPA_FP8CS_O_in_F16", "1") == "1" @@ -64,36 +90,39 @@ def flash_attn_p2p_communicate( ): """Point-to-point communications of KV and dKV in Attention with context parallelism""" send_recv_ops = [] + transport_group, transport_overridden = _get_cp_p2p_transport_group(cp_group) + if transport_overridden: + batch_p2p_comm = True if batch_p2p_comm: if rank % 2 == 0: send_op = torch.distributed.P2POp( - torch.distributed.isend, send_tensor, send_dst, cp_group + torch.distributed.isend, send_tensor, send_dst, transport_group ) recv_op = torch.distributed.P2POp( - torch.distributed.irecv, recv_tensor, recv_src, cp_group + torch.distributed.irecv, recv_tensor, recv_src, transport_group ) send_recv_ops.append(send_op) send_recv_ops.append(recv_op) else: recv_op = torch.distributed.P2POp( - torch.distributed.irecv, recv_tensor, recv_src, cp_group + torch.distributed.irecv, recv_tensor, recv_src, transport_group ) send_op = torch.distributed.P2POp( - torch.distributed.isend, send_tensor, send_dst, cp_group + torch.distributed.isend, send_tensor, send_dst, transport_group ) send_recv_ops.append(recv_op) send_recv_ops.append(send_op) send_recv_reqs = torch.distributed.batch_isend_irecv(send_recv_ops) else: if rank % 2 == 0: - send_op = torch.distributed.isend(send_tensor, send_dst, cp_group) - recv_op = torch.distributed.irecv(recv_tensor, recv_src, cp_group) + send_op = torch.distributed.isend(send_tensor, send_dst, transport_group) + recv_op = torch.distributed.irecv(recv_tensor, recv_src, transport_group) send_recv_ops.append(send_op) send_recv_ops.append(recv_op) else: - recv_op = torch.distributed.irecv(recv_tensor, recv_src, cp_group) - send_op = torch.distributed.isend(send_tensor, send_dst, cp_group) + recv_op = torch.distributed.irecv(recv_tensor, recv_src, transport_group) + send_op = torch.distributed.isend(send_tensor, send_dst, transport_group) send_recv_ops.append(recv_op) send_recv_ops.append(send_op) send_recv_reqs = send_recv_ops @@ -101,6 +130,71 @@ def flash_attn_p2p_communicate( return send_recv_reqs +def _logical_cp_all_reduce_max_(tensor, cp_group): + """Apply an in-place max reduction using a logical ring on the parent group.""" + cp_size = get_distributed_world_size(cp_group) + if cp_size == 1: + return tensor + + cp_rank = get_distributed_rank(cp_group) + cp_global_ranks = get_distributed_group_ranks(cp_group) + send_dst = cp_global_ranks[(cp_rank + 1) % cp_size] + recv_src = cp_global_ranks[(cp_rank - 1) % cp_size] + send_buffer = tensor.clone() + for _ in range(cp_size - 1): + recv_buffer = torch.empty_like(send_buffer) + requests = flash_attn_p2p_communicate( + cp_rank, + send_buffer, + send_dst, + recv_buffer, + recv_src, + cp_group, + True, + ) + for request in requests: + request.wait() + torch.maximum(tensor, recv_buffer, out=tensor) + send_buffer = recv_buffer + return tensor + + +def _logical_cp_all_to_all_single(output, input_, cp_group): + """Run equal-split all-to-all as subgroup P2P through the parent group.""" + cp_size = get_distributed_world_size(cp_group) + cp_rank = get_distributed_rank(cp_group) + cp_global_ranks = get_distributed_group_ranks(cp_group) + transport_group, _ = _get_cp_p2p_transport_group(cp_group) + + if input_.shape[0] % cp_size != 0 or output.shape[0] % cp_size != 0: + raise RuntimeError("Logical CP all-to-all requires equal splits along dimension 0.") + + input_chunks = input_.chunk(cp_size, dim=0) + output_chunks = output.chunk(cp_size, dim=0) + output_chunks[cp_rank].copy_(input_chunks[cp_rank]) + operations = [] + for peer_rank, peer_global_rank in enumerate(cp_global_ranks): + if peer_rank == cp_rank: + continue + operations.extend( + [ + torch.distributed.P2POp( + torch.distributed.isend, + input_chunks[peer_rank], + peer_global_rank, + transport_group, + ), + torch.distributed.P2POp( + torch.distributed.irecv, + output_chunks[peer_rank], + peer_global_rank, + transport_group, + ), + ] + ) + return torch.distributed.batch_isend_irecv(operations) if operations else [] + + @jit_fuser def flash_attn_fwd_out_correction_init( out_init_step: torch.Tensor, @@ -540,7 +634,9 @@ def flash_attn_a2a_communicate_softmax_offset( ) torch.cuda.current_stream().wait_stream(cp_stream) output = output.view( - *tensor.shape[:h_dim], cp_size * tensor.shape[h_dim], *tensor.shape[h_dim + 1 :] + *tensor.shape[:h_dim], + cp_size * tensor.shape[h_dim], + *tensor.shape[h_dim + 1 :], ) return output @@ -1382,7 +1478,13 @@ def forward( q, k, v = (q._data, k._data, v._data) chunk_ids_for_a2a = get_seq_chunk_ids_for_reordering_before_attn(cp_size_a2a, q.device) q, k, v = flash_attn_a2a_communicate( - [q, k, v], chunk_ids_for_a2a, seq_dim, cp_size_a2a, cp_group_a2a, cp_stream, True + [q, k, v], + chunk_ids_for_a2a, + seq_dim, + cp_size_a2a, + cp_group_a2a, + cp_stream, + True, ) if fp8 and is_input_fp8: q_fp8, k_fp8, v_fp8 = [ @@ -1480,7 +1582,9 @@ def forward( ), "Sequence length does not meet divisible requirements!" # [b, h, sq, sk] -> [b, h, sq, 2*cp, sk//(2*cp)] attn_bias_ = attn_bias.view( - *attn_bias.shape[:-1], 2 * cp_size, attn_bias.shape[-1] // (2 * cp_size) + *attn_bias.shape[:-1], + 2 * cp_size, + attn_bias.shape[-1] // (2 * cp_size), ) # [b, h, sq, sk] -> [b, h, sq, 2*cp, sk//(2*cp)] @@ -1683,10 +1787,12 @@ def forward( *fused_attn_inputs, *prepare_outputs, section ) else: - out_per_step[i], softmax_lse_per_step[i], rng_states[i] = ( - cp_p2p_fwd_flash_attn( - *flash_attn_inputs, *prepare_outputs, section - ) + ( + out_per_step[i], + softmax_lse_per_step[i], + rng_states[i], + ) = cp_p2p_fwd_flash_attn( + *flash_attn_inputs, *prepare_outputs, section ) elif i <= rank: section = "lower-triangle" @@ -1710,10 +1816,12 @@ def forward( *fused_attn_inputs, *prepare_outputs, section ) else: - out_per_step[i], softmax_lse_per_step[i], rng_states[i] = ( - cp_p2p_fwd_flash_attn( - *flash_attn_inputs, *prepare_outputs, section - ) + ( + out_per_step[i], + softmax_lse_per_step[i], + rng_states[i], + ) = cp_p2p_fwd_flash_attn( + *flash_attn_inputs, *prepare_outputs, section ) else: section = "upper-triangle" @@ -1737,10 +1845,12 @@ def forward( *fused_attn_inputs, *prepare_outputs, section ) else: - out_per_step[i], softmax_lse_per_step[i], rng_states[i] = ( - cp_p2p_fwd_flash_attn( - *flash_attn_inputs, *prepare_outputs, section - ) + ( + out_per_step[i], + softmax_lse_per_step[i], + rng_states[i], + ) = cp_p2p_fwd_flash_attn( + *flash_attn_inputs, *prepare_outputs, section ) else: # all tiles @@ -1831,9 +1941,12 @@ def forward( torch.cuda.current_stream().wait_stream(flash_attn_streams[1]) if return_max_logit: - torch.distributed.all_reduce( - max_logit, op=torch.distributed.ReduceOp.MAX, group=cp_group - ) + if is_logical_process_group(cp_group): + _logical_cp_all_reduce_max_(max_logit, cp_group) + else: + torch.distributed.all_reduce( + max_logit, op=torch.distributed.ReduceOp.MAX, group=cp_group + ) second_half_lse_seqlen = None if causal and rank < (cp_size - 1): @@ -1902,7 +2015,13 @@ def forward( if cp_size_a2a > 1: chunk_ids_for_a2a = get_seq_chunk_ids_for_reordering_after_attn(cp_size_a2a, out.device) out = flash_attn_a2a_communicate( - out, chunk_ids_for_a2a, seq_dim, cp_size_a2a, cp_group_a2a, cp_stream, False + out, + chunk_ids_for_a2a, + seq_dim, + cp_size_a2a, + cp_group_a2a, + cp_stream, + False, ) if use_fused_attention: if qkv_format == "bshd": @@ -2106,12 +2225,17 @@ def backward(ctx, dout, *_args): if attn_biases[0] is not None: # [b, h, sq, 2*cp, sk//(2*cp)] attn_dbias = torch.zeros( - *ctx.attn_bias_shape, dtype=attn_biases[0].dtype, device=attn_biases[0].device + *ctx.attn_bias_shape, + dtype=attn_biases[0].dtype, + device=attn_biases[0].device, ) # [b, h, sq, 2*cp, sk//(2*cp)] -> [b, h, 2, sq//2, 2*cp, sk//(2*cp)] only when sq > 1 (i.e. all supported bias shapes except 111s) if attn_dbias.shape[-3] > 1: attn_dbias_ = attn_dbias.view( - *attn_dbias.shape[:-3], 2, attn_dbias.shape[-3] // 2, *attn_dbias.shape[-2:] + *attn_dbias.shape[:-3], + 2, + attn_dbias.shape[-3] // 2, + *attn_dbias.shape[-2:], ) else: attn_dbias_ = None @@ -2216,7 +2340,10 @@ def backward(ctx, dout, *_args): device=kv.device, ) dkv_recv_buffer = torch.empty_like(dkv_send_buffer) - p2p_comm_buffers = [[kv, dkv_send_buffer], [kv_recv_buffer, dkv_recv_buffer]] + p2p_comm_buffers = [ + [kv, dkv_send_buffer], + [kv_recv_buffer, dkv_recv_buffer], + ] if ctx.fp8_recipe.float8_current_scaling(): dkv_buffer = torch.zeros( kv.shape, @@ -2326,13 +2453,18 @@ def backward(ctx, dout, *_args): batch_p2p_comm, ) else: - dkv_a2a_req = torch.distributed.all_to_all_single( - dkv_send_buffer, - dkv_recv_buffer, - group=ctx.cp_group, - async_op=True, - ) - send_recv_reqs = [dkv_a2a_req] + if is_logical_process_group(ctx.cp_group): + send_recv_reqs = _logical_cp_all_to_all_single( + dkv_send_buffer, dkv_recv_buffer, ctx.cp_group + ) + else: + dkv_a2a_req = torch.distributed.all_to_all_single( + dkv_send_buffer, + dkv_recv_buffer, + group=ctx.cp_group, + async_op=True, + ) + send_recv_reqs = [dkv_a2a_req] else: if i == 0: send_tensor = send_tensor[0] @@ -2341,7 +2473,13 @@ def backward(ctx, dout, *_args): send_tensor = send_tensor[1] recv_tensor = recv_tensor[1] send_recv_reqs = flash_attn_p2p_communicate( - rank, send_tensor, send_dst, recv_tensor, recv_src, ctx.cp_group, batch_p2p_comm + rank, + send_tensor, + send_dst, + recv_tensor, + recv_src, + ctx.cp_group, + batch_p2p_comm, ) kv = p2p_comm_buffers[i % 2][0] @@ -2656,7 +2794,9 @@ def backward(ctx, dout, *_args): dv = dkv_recv_buffer[:, ctx.k_numel :].view(cp_size, *ctx.v_shape) dq, dk, dv = [ ctx.dQKV_quantizer.create_tensor_from_data( - x, fake_dtype=bwd_nominal_dtype, internal=ctx.dQKV_quantizer.internal + x, + fake_dtype=bwd_nominal_dtype, + internal=ctx.dQKV_quantizer.internal, ) for x in [dq, dk, dv] ] @@ -3081,7 +3221,7 @@ def backward(ctx, dout, *_args): rank = get_distributed_rank(ctx.cp_group) (*saved_tensors,) = ctx.saved_tensors - (q, k, v, cu_seqlens_q, cu_seqlens_q_padded) = saved_tensors[:5] + q, k, v, cu_seqlens_q, cu_seqlens_q_padded = saved_tensors[:5] cu_seqlens_kv_per_step = saved_tensors[5:7] out_per_step = saved_tensors[7:9] softmax_lse_per_step = saved_tensors[9:11] @@ -3430,9 +3570,14 @@ def forward( fused_attn_backend = None max_logit = None - QKV_quantizer, O_quantizer, S_quantizer, dQKV_quantizer, dO_quantizer, dP_quantizer = ( - dpa_utils.get_attention_quantizers(fp8, quantizers) - ) + ( + QKV_quantizer, + O_quantizer, + S_quantizer, + dQKV_quantizer, + dO_quantizer, + dP_quantizer, + ) = dpa_utils.get_attention_quantizers(fp8, quantizers) q_fp8, k_fp8, v_fp8 = (None, None, None) if fp8: @@ -4041,9 +4186,14 @@ def attn_forward_func_with_cp( cp_group = cp_group[0] cp_comm_type = "a2a" else: - assert isinstance( - cp_group, dist_group_type - ), f"cp_group must be {dist_group_type} type for {cp_comm_type=}!" + if is_logical_process_group(cp_group): + assert ( + cp_comm_type == "p2p" + ), f"A logical CP group can only be used with cp_comm_type='p2p'; got {cp_comm_type=}." + else: + assert isinstance( + cp_group, dist_group_type + ), f"cp_group must be {dist_group_type} type for {cp_comm_type=}!" assert qkv_format in [ "bshd", @@ -4286,9 +4436,9 @@ def get_batch_on_this_cp_rank( raise ValueError(f"Unsupported qvk_format: {qvk_format}!") if qvk_format == "thd": # Get context parallel size and rank - cp_size = torch.distributed.get_world_size(group=cp_group) + cp_size = get_distributed_world_size(cp_group) if cp_size > 1: - cp_rank = torch.distributed.get_rank(group=cp_group) + cp_rank = get_distributed_rank(cp_group) # Calculate the chunk sizes for each sequence total_slices_of_any_sequence = 2 * cp_size diff --git a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index 2dc42be18a..e6e7271efd 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py @@ -40,6 +40,7 @@ ) from transformer_engine.pytorch.distributed import ( get_distributed_world_size, + is_logical_process_group, checkpoint, set_all_rng_states, CudaRNGStatesTracker, @@ -1236,7 +1237,9 @@ def forward( # adjust max_seqlen and cu_seqlens for CP cp_size = 1 - if isinstance(self.cp_group, dist_group_type): + if isinstance(self.cp_group, dist_group_type) or is_logical_process_group( + self.cp_group + ): cp_size = get_distributed_world_size(self.cp_group) elif isinstance(self.cp_group, list): for group in self.cp_group: diff --git a/transformer_engine/pytorch/attention/multi_head_attention.py b/transformer_engine/pytorch/attention/multi_head_attention.py index d95d327c78..f8672ab381 100644 --- a/transformer_engine/pytorch/attention/multi_head_attention.py +++ b/transformer_engine/pytorch/attention/multi_head_attention.py @@ -26,6 +26,7 @@ from transformer_engine.pytorch.distributed import ( get_distributed_world_size, get_distributed_rank, + is_logical_process_group, ) from transformer_engine.pytorch.attention.dot_product_attention import DotProductAttention @@ -609,7 +610,7 @@ def set_context_parallel_group( across each CP sub-group (e.g., via NVLink), then exchanging KV with p2p between sub-groups (e.g., via IBLink). """ - if isinstance(cp_group, dist_group_type): + if isinstance(cp_group, dist_group_type) or is_logical_process_group(cp_group): self.cp_size = get_distributed_world_size(cp_group) self.cp_rank = get_distributed_rank(cp_group) elif isinstance(cp_group, list): diff --git a/transformer_engine/pytorch/distributed.py b/transformer_engine/pytorch/distributed.py index b80e58fe20..136a59eb1e 100644 --- a/transformer_engine/pytorch/distributed.py +++ b/transformer_engine/pytorch/distributed.py @@ -163,17 +163,34 @@ def set_tensor_model_parallel_attributes( setattr(tensor, "partition_stride", stride) +def is_logical_process_group(group: Any) -> bool: + """Return whether ``group`` is a topology-only CP group descriptor.""" + return ( + group is not None + and hasattr(group, "ranks") + and hasattr(group, "cp_size") + and hasattr(group, "cp_rank") + ) + + @lru_cache -def get_distributed_world_size(group: Optional[dist_group_type] = None) -> int: - """Return world size for the distributed group.""" +def _get_distributed_world_size(group: Optional[dist_group_type] = None) -> int: + """Return world size for a real distributed process group.""" if not torch.distributed.is_initialized(): return 1 return torch.distributed.get_world_size(group=group) +def get_distributed_world_size(group: Optional[dist_group_type] = None) -> int: + """Return world size for a distributed group or logical CP descriptor.""" + if is_logical_process_group(group): + return int(group.cp_size) + return _get_distributed_world_size(group) + + @lru_cache -def get_distributed_rank(group: Optional[dist_group_type] = None) -> int: - """Return my rank for the distributed group.""" +def _get_distributed_rank(group: Optional[dist_group_type] = None) -> int: + """Return my rank for a real distributed process group.""" if not torch.distributed.is_initialized(): raise RuntimeError( "torch.distributed is not initialized. Call torch.distributed.init_process_group() " @@ -182,6 +199,20 @@ def get_distributed_rank(group: Optional[dist_group_type] = None) -> int: return torch.distributed.get_rank(group=group) +def get_distributed_rank(group: Optional[dist_group_type] = None) -> int: + """Return my rank for a distributed group or logical CP descriptor.""" + if is_logical_process_group(group): + return int(group.cp_rank) + return _get_distributed_rank(group) + + +def get_distributed_group_ranks(group) -> Tuple[int, ...]: + """Return global ranks for a ProcessGroup or logical CP descriptor.""" + if is_logical_process_group(group): + return tuple(group.ranks) + return tuple(torch.distributed.get_process_group_ranks(group)) + + def initialize_affine_weight_gpu( weight: torch.Tensor, init_method: Callable,