Skip to content
Draft
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
67 changes: 67 additions & 0 deletions tests/pytorch/test_cuda_graphs.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@
#
# See LICENSE for license information.

import gc
import weakref
from typing import Callable, Dict, Iterable, List, Tuple, Union
import pytest

Expand All @@ -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
Expand All @@ -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()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
Loading
Loading