From 3000db942b65621a95cd0597dc936777bf684da3 Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Mon, 20 Jul 2026 10:34:51 -0500 Subject: [PATCH 1/4] Add compatibility between training CGs and CP>1 Signed-off-by: Teodor-Dumitru Ene --- .../core/extensions/transformer_engine.py | 5 + megatron/core/packed_seq_params.py | 1 + .../transformer/test_cuda_graphs.py | 107 ++++++++++++++++++ 3 files changed, 113 insertions(+) diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index ac7d5c1da9b..8f98e67007f 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -1754,6 +1754,11 @@ def __init__( self.kept_packed_seq_params.discard("seq_idx") self.kept_packed_seq_params.discard("tokens_per_sample") + # providing `pad_between_seqs` lets callers skip a GPU sync inside TE thd padding. + dpa_forward = te.pytorch.DotProductAttention.forward + if "pad_between_seqs" not in inspect.signature(dpa_forward).parameters: + self.kept_packed_seq_params.discard("pad_between_seqs") + if config.qk_clip or config.log_max_attention_logit: # qk-clip is only supported in TE 2.9.0 and later assert is_te_min_version("2.9.0"), "qk-clip is only supported in TE 2.9.0 and later" diff --git a/megatron/core/packed_seq_params.py b/megatron/core/packed_seq_params.py index bd598bb557a..1c26af56244 100644 --- a/megatron/core/packed_seq_params.py +++ b/megatron/core/packed_seq_params.py @@ -25,6 +25,7 @@ class PackedSeqParams: total_tokens: int = None seq_idx: Tensor = None tokens_per_sample: int = None + pad_between_seqs: bool = None def __post_init__(self): """Pre-compute seq_idx for Mamba mixer CUDA graph compatibility. diff --git a/tests/unit_tests/transformer/test_cuda_graphs.py b/tests/unit_tests/transformer/test_cuda_graphs.py index 52edbf9a264..5290d6982be 100644 --- a/tests/unit_tests/transformer/test_cuda_graphs.py +++ b/tests/unit_tests/transformer/test_cuda_graphs.py @@ -23,6 +23,7 @@ destroy_num_microbatches_calculator, init_num_microbatches_calculator, ) +from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.pipeline_parallel.schedules import set_current_microbatch from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import ( @@ -34,6 +35,7 @@ CudaGraphManager, TECudaGraphHelper, _CudagraphGlobalRecord, + create_cudagraphs, ) from megatron.core.transformer.enums import CudaGraphModule, CudaGraphScope, InferenceCudaGraphScope from megatron.core.transformer.mlp import MLPSubmodules @@ -444,6 +446,111 @@ def test_gpu_cudagraph(self): ) +@pytest.mark.skipif( + not (HAVE_TE and is_te_min_version("1.5.0")), + reason="use_te_rng_tracker requires TransformerEngine version >= 1.5", +) +class TestPackedSeqCudagraphs: + """Training CUDA graphs over thd input with padding between sequences. + + The padded cu_seqlens describe a slot layout that differs from the actual lengths, + and pad_between_seqs is set explicitly so TE does spend a GPU sync inferring it. + cp_size == 2 additionally captures TE's ring-P2P context-parallel attention inside the graphs. + """ + + SEQ_LENGTHS = [7, 5] + SLOT_STARTS = [0, 8, 16] # slot layout aligned to 2 * cp_size for every cp_size tested + BIN_SIZE = 32 + + def teardown_method(self, method): + Utils.destroy_model_parallel() + _CudagraphGlobalRecord.cudagraph_created = False + _CudagraphGlobalRecord.cudagraph_record = [] + CudaGraphManager.global_mempool = None + + def _build_packed_seq_params(self, device): + # Actual boundaries: each sequence's real tokens inside its slot; the trailing bin + # padding [SLOT_STARTS[-1], BIN_SIZE) forms a ghost slot of pad tokens. + boundaries = [0] + for length in self.SEQ_LENGTHS: + boundaries.append(boundaries[-1] + length) + boundaries.append(boundaries[-1] + self.BIN_SIZE - self.SLOT_STARTS[-1]) + cu_seqlens = torch.tensor(boundaries, dtype=torch.int32, device=device) + cu_seqlens_padded = torch.tensor( + self.SLOT_STARTS + [self.BIN_SIZE], dtype=torch.int32, device=device + ) + return PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + cu_seqlens_q_padded=cu_seqlens_padded, + cu_seqlens_kv_padded=cu_seqlens_padded, + max_seqlen_q=self.BIN_SIZE, + max_seqlen_kv=self.BIN_SIZE, + pad_between_seqs=True, + ) + + @pytest.mark.parametrize("cp_size", [1, 2]) + def test_thd_capture_with_pad_between_seqs(self, cp_size): + initialize_rng_tracker(use_te_rng_tracker=True, force_reset=True) + Utils.initialize_model_parallel(context_parallel_size=cp_size) + model_parallel_cuda_manual_seed(123) + + config = TransformerConfig( + num_layers=2, + hidden_size=64, + num_attention_heads=4, + context_parallel_size=cp_size, + bf16=True, + params_dtype=torch.bfloat16, + attention_dropout=0.0, + hidden_dropout=0.0, + cuda_graph_impl="local", + use_cpu_initialization=True, + ) + block = TransformerBlock(config, get_gpt_layer_with_transformer_engine_spec()).cuda() + block.train() + # CUDA-graphed backward assumes DDP-style grad accumulation buffers. + for param in block.parameters(): + param.main_grad = torch.zeros_like(param) + + packed_seq_params = self._build_packed_seq_params(torch.device('cuda')) + # Each CP rank holds its 1/cp_size share of the bin's tokens. + hidden_states = torch.randn( + (self.BIN_SIZE // cp_size, 1, config.hidden_size), + dtype=torch.bfloat16, + device='cuda', + ) + + eager_out = block( + hidden_states=hidden_states, attention_mask=None, packed_seq_params=packed_seq_params + ) + eager_out.sum().backward() + + # This is the primary function under test. + create_cudagraphs() + + for layer in block.layers: + runners = layer.cudagraph_manager.cudagraph_runners + assert len(runners) == 1 + assert runners[0].fwd_graph is not None + + graphed_out = block( + hidden_states=hidden_states, attention_mask=None, packed_seq_params=packed_seq_params + ) + assert torch.allclose(graphed_out.float(), eager_out.float(), rtol=1e-2, atol=1e-2) + graphed_out.sum().backward() + + # Destroy captured graphs deterministically before parallel-state teardown. + for layer in block.layers: + for runner in layer.cudagraph_manager.cudagraph_runners: + if hasattr(runner, "fwd_graph"): + del runner.fwd_graph + if hasattr(runner, "bwd_graph"): + del runner.bwd_graph + torch.cuda.synchronize() + + @pytest.mark.skipif( not (HAVE_TE and is_te_min_version("1.5.0")), reason="use_te_rng_tracker requires TransformerEngine version >= 1.5", From fb17c401fe50ee53a76ef60faabf3dc3dc8ab87e Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Mon, 20 Jul 2026 11:15:23 -0500 Subject: [PATCH 2/4] Address reviewer comments Signed-off-by: Teodor-Dumitru Ene --- megatron/core/extensions/transformer_engine.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index 8f98e67007f..3a8d7f23d09 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -1754,9 +1754,7 @@ def __init__( self.kept_packed_seq_params.discard("seq_idx") self.kept_packed_seq_params.discard("tokens_per_sample") - # providing `pad_between_seqs` lets callers skip a GPU sync inside TE thd padding. - dpa_forward = te.pytorch.DotProductAttention.forward - if "pad_between_seqs" not in inspect.signature(dpa_forward).parameters: + if get_te_version() < PkgVersion("2.2.0"): self.kept_packed_seq_params.discard("pad_between_seqs") if config.qk_clip or config.log_max_attention_logit: From 5b96f1345fa7855bf6a81cb171470dba3783e223 Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Mon, 20 Jul 2026 11:30:54 -0500 Subject: [PATCH 3/4] lint Signed-off-by: Teodor-Dumitru Ene --- tests/unit_tests/transformer/test_cuda_graphs.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/tests/unit_tests/transformer/test_cuda_graphs.py b/tests/unit_tests/transformer/test_cuda_graphs.py index 5290d6982be..fd5b4845188 100644 --- a/tests/unit_tests/transformer/test_cuda_graphs.py +++ b/tests/unit_tests/transformer/test_cuda_graphs.py @@ -517,9 +517,7 @@ def test_thd_capture_with_pad_between_seqs(self, cp_size): packed_seq_params = self._build_packed_seq_params(torch.device('cuda')) # Each CP rank holds its 1/cp_size share of the bin's tokens. hidden_states = torch.randn( - (self.BIN_SIZE // cp_size, 1, config.hidden_size), - dtype=torch.bfloat16, - device='cuda', + (self.BIN_SIZE // cp_size, 1, config.hidden_size), dtype=torch.bfloat16, device='cuda' ) eager_out = block( From f9da2bcdb843b209eeed81325b607468d168f376 Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Mon, 20 Jul 2026 21:46:06 -0500 Subject: [PATCH 4/4] Fix typo Signed-off-by: Teodor-Dumitru Ene --- tests/unit_tests/transformer/test_cuda_graphs.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tests/unit_tests/transformer/test_cuda_graphs.py b/tests/unit_tests/transformer/test_cuda_graphs.py index fd5b4845188..ef556b29f42 100644 --- a/tests/unit_tests/transformer/test_cuda_graphs.py +++ b/tests/unit_tests/transformer/test_cuda_graphs.py @@ -517,7 +517,10 @@ def test_thd_capture_with_pad_between_seqs(self, cp_size): packed_seq_params = self._build_packed_seq_params(torch.device('cuda')) # Each CP rank holds its 1/cp_size share of the bin's tokens. hidden_states = torch.randn( - (self.BIN_SIZE // cp_size, 1, config.hidden_size), dtype=torch.bfloat16, device='cuda' + (self.BIN_SIZE // cp_size, 1, config.hidden_size), + dtype=torch.bfloat16, + device='cuda', + requires_grad=True, ) eager_out = block(