diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py index 3fa91068769..8c564a71487 100644 --- a/megatron/core/transformer/transformer_layer.py +++ b/megatron/core/transformer/transformer_layer.py @@ -324,6 +324,10 @@ def __init__( name (str | None): module instance name passed top-down from its paranet module """ self.submodules_config = submodules + # GraphableMegatronModule may create a local CUDA graph manager during super().__init__(). + # Dense layers need a default before that hook runs, while MoETransformerLayer sets this + # to True before entering this constructor. + self.is_moe_layer = getattr(self, "is_moe_layer", False) super().__init__(config=config, vp_stage=vp_stage) if pg_collection is None: @@ -548,7 +552,7 @@ def create_mcore_cudagraph_manager(self, config): CudaGraphModule.mlp in self.config.cuda_graph_modules and self.submodules_config.mlp != IdentityOp ): - # Cudagraphing MoE layers are supposed handled by MoeTransforerLayer + # Cudagraphing MoE layers is handled by MoETransformerLayer. assert not self.is_moe_layer self.cudagraph_manager = CudaGraphManager(config) diff --git a/tests/unit_tests/transformer/test_cuda_graphs.py b/tests/unit_tests/transformer/test_cuda_graphs.py index f387e58953e..b54ba1d3c01 100644 --- a/tests/unit_tests/transformer/test_cuda_graphs.py +++ b/tests/unit_tests/transformer/test_cuda_graphs.py @@ -450,6 +450,26 @@ def test_gpu_cudagraph(self): .fwd_graph ) + @pytest.mark.skipif( + not (HAVE_TE and is_te_min_version("1.5.0")), + reason="use_te_rng_tracker requires TransformerEngine version >= 1.5", + ) + def test_dense_mlp_scope_constructs(self): + config = TransformerConfig( + num_layers=8, + hidden_size=64, + num_attention_heads=4, + use_cpu_initialization=True, + cuda_graph_impl="local", + cuda_graph_modules=[CudaGraphModule.mlp], + ) + + block = TransformerBlock(config, get_gpt_layer_with_transformer_engine_spec()) + + assert block.layers + assert all(not layer.is_moe_layer for layer in block.layers) + assert all(isinstance(layer.cudagraph_manager, CudaGraphManager) for layer in block.layers) + @pytest.mark.skipif( not (HAVE_TE and is_te_min_version("1.5.0")),