diff --git a/megatron/core/models/bert/bert_layer_specs.py b/megatron/core/models/bert/bert_layer_specs.py index 8415ef02cc5..832d55a35d5 100644 --- a/megatron/core/models/bert/bert_layer_specs.py +++ b/megatron/core/models/bert/bert_layer_specs.py @@ -45,40 +45,48 @@ HAVE_APEX = False -def get_bert_layer_with_transformer_engine_spec(): - """Use this spec to use lower-level Transformer Engine modules (required for fp8 training). +def get_bert_layer_with_transformer_engine_submodules() -> TransformerLayerSubmodules: + """Use these submodules to use lower-level Transformer Engine modules (required for fp8 training). Returns: - ModuleSpec: Module specification with TE modules + TransformerLayerSubmodules: Submodules with TE modules. """ if not HAVE_TE: raise ImportError( "Transformer Engine is not installed. Please use local Bert layer spec instead." ) - return ModuleSpec( - module=TransformerLayer, - submodules=TransformerLayerSubmodules( - self_attention=ModuleSpec( - module=SelfAttention, - params={"attn_mask_type": AttnMaskType.padding}, - submodules=SelfAttentionSubmodules( - linear_qkv=not_none(TELayerNormColumnParallelLinear), - core_attention=not_none(TEDotProductAttention), - linear_proj=TERowParallelLinear, - q_layernorm=IdentityOp, - k_layernorm=IdentityOp, - ), + return TransformerLayerSubmodules( + self_attention=ModuleSpec( + module=SelfAttention, + params={"attn_mask_type": AttnMaskType.padding}, + submodules=SelfAttentionSubmodules( + linear_qkv=not_none(TELayerNormColumnParallelLinear), + core_attention=not_none(TEDotProductAttention), + linear_proj=not_none(TERowParallelLinear), + q_layernorm=IdentityOp, + k_layernorm=IdentityOp, ), - self_attn_bda=get_bias_dropout_add, - mlp=ModuleSpec( - module=MLP, - submodules=MLPSubmodules( - linear_fc1=TELayerNormColumnParallelLinear, linear_fc2=TERowParallelLinear - ), + ), + self_attn_bda=get_bias_dropout_add, + mlp=ModuleSpec( + module=MLP, + submodules=MLPSubmodules( + linear_fc1=TELayerNormColumnParallelLinear, linear_fc2=TERowParallelLinear ), - mlp_bda=get_bias_dropout_add, ), + mlp_bda=get_bias_dropout_add, + ) + + +def get_bert_layer_with_transformer_engine_spec(): + """Use this spec to use lower-level Transformer Engine modules (required for fp8 training). + + Returns: + ModuleSpec: Module specification with TE modules + """ + return ModuleSpec( + module=TransformerLayer, submodules=get_bert_layer_with_transformer_engine_submodules() ) diff --git a/megatron/core/models/bert/bert_model.py b/megatron/core/models/bert/bert_model.py index abe9bc1c9b7..fd6aa4bc03d 100644 --- a/megatron/core/models/bert/bert_model.py +++ b/megatron/core/models/bert/bert_model.py @@ -15,6 +15,7 @@ from megatron.core.models.common.embeddings.rotary_pos_embedding import RotaryEmbedding from megatron.core.models.common.language_module.language_module import LanguageModule from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.transformer.attention import SelfAttentionSubmodules from megatron.core.transformer.dot_product_attention import ( DotProductAttention as MCoreDotProductAttention, ) @@ -22,10 +23,10 @@ from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.transformer.transformer_block import TransformerBlock from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.transformer.transformer_layer import TransformerLayerSubmodules from megatron.core.transformer.utils import get_linear_layer -from megatron.core.utils import deprecate_inference_params +from megatron.core.utils import deprecate_inference_params, is_te_min_version from megatron.core.utils import get_te_version as _get_te_version -from megatron.core.utils import is_te_min_version def get_te_version(): @@ -187,6 +188,11 @@ def _sanity_check_attention_and_get_attn_mask_dimension(self) -> str: """ attention_backend = self.config.attention_backend attn_mask_dimensions = None + assert isinstance(self.transformer_layer_spec.submodules, TransformerLayerSubmodules) + assert isinstance( + self.transformer_layer_spec.submodules.self_attention.submodules, + SelfAttentionSubmodules, + ) # For local layer spec we just use b1ss if ( self.transformer_layer_spec.submodules.self_attention.submodules.core_attention diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py index 49501ee54eb..ef088c41a1d 100755 --- a/megatron/core/models/gpt/gpt_layer_specs.py +++ b/megatron/core/models/gpt/gpt_layer_specs.py @@ -1,5 +1,5 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - +import functools import warnings from typing import Optional, Union @@ -37,6 +37,7 @@ TransformerLayerSubmodules, get_transformer_layer_offset, ) +from megatron.core.typed_torch import copy_signature from megatron.core.utils import is_te_min_version try: @@ -72,12 +73,12 @@ HAVE_APEX = False -def get_gpt_layer_with_inference_spec( +def get_gpt_layer_with_inference_submodules( qk_layernorm: Optional[bool] = False, multi_latent_attention: Optional[bool] = False, qk_l2_norm: Optional[bool] = False, -) -> ModuleSpec: - """Use this spec to use inference optimized linear layers. +) -> TransformerLayerSubmodules: + """Returns the TransformerLayerSubmodules to use for inference optimized linear layers. Args: qk_layernorm (bool, optional): To use layernorm for queries/keys. Defaults to False. multi_latent_attention (bool, optional): To use MLA. Defaults to False. @@ -107,68 +108,74 @@ def get_gpt_layer_with_inference_spec( if qk_layernorm else backend.column_parallel_linear() ) - return ModuleSpec( - module=TransformerLayer, - submodules=TransformerLayerSubmodules( - input_layernorm=backend.layer_norm(), - self_attention=ModuleSpec( - module=MLASelfAttention, - params={"attn_mask_type": AttnMaskType.causal}, - submodules=MLASelfAttentionSubmodules( - linear_q_proj=backend.column_parallel_linear(), - linear_q_down_proj=backend.linear(), - linear_q_up_proj=linear_q_up_proj, - linear_kv_down_proj=backend.linear(), - linear_kv_up_proj=linear_kv_up_proj, - core_attention=backend.core_attention(), - linear_proj=backend.row_parallel_linear(), - q_layernorm=IdentityOp, - kv_layernorm=IdentityOp, - ), + return TransformerLayerSubmodules( + input_layernorm=backend.layer_norm(), + self_attention=ModuleSpec( + module=MLASelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=MLASelfAttentionSubmodules( + linear_q_proj=backend.column_parallel_linear(), + linear_q_down_proj=backend.linear(), + linear_q_up_proj=linear_q_up_proj, + linear_kv_down_proj=backend.linear(), + linear_kv_up_proj=linear_kv_up_proj, + core_attention=backend.core_attention(), + linear_proj=backend.row_parallel_linear(), + q_layernorm=IdentityOp, + kv_layernorm=IdentityOp, ), - self_attn_bda=get_bias_dropout_add, - pre_mlp_layernorm=IdentityOp, - mlp=mlp, - mlp_bda=get_bias_dropout_add, ), + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=IdentityOp, + mlp=mlp, + mlp_bda=get_bias_dropout_add, ) else: qk_norm = backend.layer_norm(for_qk=True) - return ModuleSpec( - module=TransformerLayer, - submodules=TransformerLayerSubmodules( - self_attention=ModuleSpec( - module=SelfAttention, - params={"attn_mask_type": AttnMaskType.causal}, - submodules=SelfAttentionSubmodules( - linear_qkv=backend.column_parallel_layer_norm_linear(), - core_attention=backend.core_attention(), - linear_proj=backend.row_parallel_linear(), - q_layernorm=( - L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) - ), - k_layernorm=( - L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) - ), + return TransformerLayerSubmodules( + self_attention=ModuleSpec( + module=SelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=SelfAttentionSubmodules( + linear_qkv=backend.column_parallel_layer_norm_linear(), + core_attention=backend.core_attention(), + linear_proj=backend.row_parallel_linear(), + q_layernorm=( + L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) + ), + k_layernorm=( + L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) ), ), - self_attn_bda=get_bias_dropout_add, - pre_mlp_layernorm=IdentityOp, - mlp=mlp, - mlp_bda=get_bias_dropout_add, - sharded_state_dict_keys_map={ - "mlp.0.weight": "mlp.linear_fc1.layer_norm_weight", - "mlp.0.bias": "mlp.linear_fc1.layer_norm_bias", - "mlp.1.basic_ops.0.weight": "mlp.linear_fc1.weight", - "mlp.1.basic_ops.1.bias": "mlp.linear_fc1.bias", - "mlp.3.basic_ops.0.weight": "mlp.linear_fc2.weight", - "mlp.3.basic_ops.1.bias": "mlp.linear_fc2.bias", - }, ), + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=IdentityOp, + mlp=mlp, + mlp_bda=get_bias_dropout_add, + sharded_state_dict_keys_map={ + "mlp.0.weight": "mlp.linear_fc1.layer_norm_weight", + "mlp.0.bias": "mlp.linear_fc1.layer_norm_bias", + "mlp.1.basic_ops.0.weight": "mlp.linear_fc1.weight", + "mlp.1.basic_ops.1.bias": "mlp.linear_fc1.bias", + "mlp.3.basic_ops.0.weight": "mlp.linear_fc2.weight", + "mlp.3.basic_ops.1.bias": "mlp.linear_fc2.bias", + }, ) -def get_gpt_layer_with_transformer_engine_spec( +@functools.wraps(get_gpt_layer_with_inference_submodules) +@copy_signature(get_gpt_layer_with_inference_submodules) +def get_gpt_layer_with_inference_spec(*args, **kwargs) -> ModuleSpec: + """Use this spec to use inference optimized linear layers. + + See get_gpt_layer_with_inference_submodules for arguments. + """ + return ModuleSpec( + module=TransformerLayer, submodules=get_gpt_layer_with_inference_submodules(*args, **kwargs) + ) + + +def get_gpt_layer_with_transformer_engine_submodules( num_experts: Optional[int] = None, moe_grouped_gemm: Optional[bool] = False, qk_layernorm: Optional[bool] = False, @@ -181,7 +188,7 @@ def get_gpt_layer_with_transformer_engine_spec( use_te_activation_func: bool = False, use_kitchen_attention: bool = False, kitchen_attention_backend: str = "sdpa", -) -> ModuleSpec: +) -> TransformerLayerSubmodules: """Use this spec to use lower-level Transformer Engine modules (required for fp8 training). @@ -242,68 +249,75 @@ def get_gpt_layer_with_transformer_engine_spec( if qk_layernorm else backend.column_parallel_linear() ) - return ModuleSpec( - module=TransformerLayer, - submodules=TransformerLayerSubmodules( - input_layernorm=backend.layer_norm(), - self_attention=ModuleSpec( - module=MLASelfAttention, - params={"attn_mask_type": AttnMaskType.causal}, - submodules=MLASelfAttentionSubmodules( - linear_q_proj=backend.column_parallel_linear(), - linear_q_down_proj=backend.linear(), - linear_q_up_proj=linear_q_up_proj, - linear_kv_down_proj=backend.linear(), - linear_kv_up_proj=linear_kv_up_proj, - core_attention=backend.core_attention(), - linear_proj=backend.row_parallel_linear(), - q_layernorm=IdentityOp, - kv_layernorm=IdentityOp, - ), + return TransformerLayerSubmodules( + input_layernorm=backend.layer_norm(), + self_attention=ModuleSpec( + module=MLASelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=MLASelfAttentionSubmodules( + linear_q_proj=backend.column_parallel_linear(), + linear_q_down_proj=backend.linear(), + linear_q_up_proj=linear_q_up_proj, + linear_kv_down_proj=backend.linear(), + linear_kv_up_proj=linear_kv_up_proj, + core_attention=backend.core_attention(), + linear_proj=backend.row_parallel_linear(), + q_layernorm=IdentityOp, + kv_layernorm=IdentityOp, ), - self_attn_bda=get_bias_dropout_add, - pre_mlp_layernorm=backend.layer_norm() if num_experts else IdentityOp, - mlp=mlp, - mlp_bda=get_bias_dropout_add, ), + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=backend.layer_norm() if num_experts else IdentityOp, + mlp=mlp, + mlp_bda=get_bias_dropout_add, ) else: qk_norm = backend.layer_norm(for_qk=True) - return ModuleSpec( - module=TransformerLayer, - submodules=TransformerLayerSubmodules( - self_attention=ModuleSpec( - module=SelfAttention, - params={"attn_mask_type": AttnMaskType.causal}, - submodules=SelfAttentionSubmodules( - linear_qkv=backend.column_parallel_layer_norm_linear(), - core_attention=backend.core_attention(), - linear_proj=backend.row_parallel_linear(), - q_layernorm=( - L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) - ), - k_layernorm=( - L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) - ), + return TransformerLayerSubmodules( + self_attention=ModuleSpec( + module=SelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=SelfAttentionSubmodules( + linear_qkv=backend.column_parallel_layer_norm_linear(), + core_attention=backend.core_attention(), + linear_proj=backend.row_parallel_linear(), + q_layernorm=( + L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) + ), + k_layernorm=( + L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) ), ), - self_attn_bda=get_bias_dropout_add, - pre_mlp_layernorm=backend.layer_norm() if num_experts else IdentityOp, - mlp=mlp, - mlp_bda=get_bias_dropout_add, - sharded_state_dict_keys_map={ - "mlp.0.weight": "mlp.linear_fc1.layer_norm_weight", - "mlp.0.bias": "mlp.linear_fc1.layer_norm_bias", - "mlp.1.basic_ops.0.weight": "mlp.linear_fc1.weight", - "mlp.1.basic_ops.1.bias": "mlp.linear_fc1.bias", - "mlp.3.basic_ops.0.weight": "mlp.linear_fc2.weight", - "mlp.3.basic_ops.1.bias": "mlp.linear_fc2.bias", - }, ), + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=backend.layer_norm() if num_experts else IdentityOp, + mlp=mlp, + mlp_bda=get_bias_dropout_add, + sharded_state_dict_keys_map={ + "mlp.0.weight": "mlp.linear_fc1.layer_norm_weight", + "mlp.0.bias": "mlp.linear_fc1.layer_norm_bias", + "mlp.1.basic_ops.0.weight": "mlp.linear_fc1.weight", + "mlp.1.basic_ops.1.bias": "mlp.linear_fc1.bias", + "mlp.3.basic_ops.0.weight": "mlp.linear_fc2.weight", + "mlp.3.basic_ops.1.bias": "mlp.linear_fc2.bias", + }, ) -def get_gpt_layer_local_spec( +@functools.wraps(get_gpt_layer_with_transformer_engine_submodules) +@copy_signature(get_gpt_layer_with_transformer_engine_submodules) +def get_gpt_layer_with_transformer_engine_spec(*args, **kwargs) -> ModuleSpec: + """Use this spec to use lower-level Transformer Engine modules (required for fp8 training). + + See get_gpt_layer_with_transformer_engine_submodules for arguments. + """ + return ModuleSpec( + module=TransformerLayer, + submodules=get_gpt_layer_with_transformer_engine_submodules(*args, **kwargs), + ) + + +def get_gpt_layer_local_submodules( num_experts: Optional[int] = None, moe_grouped_gemm: Optional[bool] = False, qk_layernorm: Optional[bool] = False, @@ -315,7 +329,7 @@ def get_gpt_layer_local_spec( use_kitchen: bool = False, use_kitchen_attention: bool = False, kitchen_attention_backend: str = "sdpa", -) -> ModuleSpec: +) -> TransformerLayerSubmodules: """Use this spec for an implementation using only modules in Megatron-Core. @@ -365,63 +379,69 @@ def get_gpt_layer_local_spec( if multi_latent_attention: assert qk_l2_norm is False, "qk_l2_norm is not supported with MLA." - return ModuleSpec( - module=TransformerLayer, - submodules=TransformerLayerSubmodules( - input_layernorm=layer_norm, - self_attention=ModuleSpec( - module=MLASelfAttention, - params={"attn_mask_type": AttnMaskType.causal}, - submodules=MLASelfAttentionSubmodules( - linear_q_proj=backend.column_parallel_linear(), - linear_q_down_proj=backend.column_parallel_linear(), - linear_q_up_proj=backend.column_parallel_linear(), - linear_kv_down_proj=backend.column_parallel_linear(), - linear_kv_up_proj=backend.column_parallel_linear(), - core_attention=backend.core_attention(), - linear_proj=backend.row_parallel_linear(), - q_layernorm=qk_norm if qk_layernorm else IdentityOp, - kv_layernorm=qk_norm if qk_layernorm else IdentityOp, - ), + return TransformerLayerSubmodules( + input_layernorm=layer_norm, + self_attention=ModuleSpec( + module=MLASelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=MLASelfAttentionSubmodules( + linear_q_proj=backend.column_parallel_linear(), + linear_q_down_proj=backend.column_parallel_linear(), + linear_q_up_proj=backend.column_parallel_linear(), + linear_kv_down_proj=backend.column_parallel_linear(), + linear_kv_up_proj=backend.column_parallel_linear(), + core_attention=backend.core_attention(), + linear_proj=backend.row_parallel_linear(), + q_layernorm=qk_norm if qk_layernorm else IdentityOp, + kv_layernorm=qk_norm if qk_layernorm else IdentityOp, ), - self_attn_bda=get_bias_dropout_add, - pre_mlp_layernorm=layer_norm, - mlp=mlp, - mlp_bda=get_bias_dropout_add, ), + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=layer_norm, + mlp=mlp, + mlp_bda=get_bias_dropout_add, ) else: - return ModuleSpec( - module=TransformerLayer, - submodules=TransformerLayerSubmodules( - input_layernorm=layer_norm, - self_attention=ModuleSpec( - module=SelfAttention, - params={"attn_mask_type": AttnMaskType.causal}, - submodules=SelfAttentionSubmodules( - linear_qkv=backend.column_parallel_linear(), - core_attention=backend.core_attention(), - linear_proj=backend.row_parallel_linear(), - q_layernorm=( - L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) - ), - k_layernorm=( - L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) - ), + return TransformerLayerSubmodules( + input_layernorm=layer_norm, + self_attention=ModuleSpec( + module=SelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=SelfAttentionSubmodules( + linear_qkv=backend.column_parallel_linear(), + core_attention=backend.core_attention(), + linear_proj=backend.row_parallel_linear(), + q_layernorm=( + L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) + ), + k_layernorm=( + L2Norm if qk_l2_norm else (qk_norm if qk_layernorm else IdentityOp) ), ), - self_attn_bda=get_bias_dropout_add, - pre_mlp_layernorm=layer_norm, - mlp=mlp, - mlp_bda=get_bias_dropout_add, - sharded_state_dict_keys_map={ - "input_layernorm.": "self_attention.linear_qkv.layer_norm_", - "pre_mlp_layernorm.": "mlp.linear_fc1.layer_norm_", - }, ), + self_attn_bda=get_bias_dropout_add, + pre_mlp_layernorm=layer_norm, + mlp=mlp, + mlp_bda=get_bias_dropout_add, + sharded_state_dict_keys_map={ + "input_layernorm.": "self_attention.linear_qkv.layer_norm_", + "pre_mlp_layernorm.": "mlp.linear_fc1.layer_norm_", + }, ) +@functools.wraps(get_gpt_layer_local_submodules) +@copy_signature(get_gpt_layer_local_submodules) +def get_gpt_layer_local_spec(*args, **kwargs) -> ModuleSpec: + """Use this spec for an implementation using only modules in Megatron-Core. + + See get_gpt_layer_local_submodules for arguments. + """ + return ModuleSpec( + module=TransformerLayer, submodules=get_gpt_layer_local_submodules(*args, **kwargs) + ) + + def _get_mlp_module_spec( use_te: Optional[bool] = True, num_experts: Optional[int] = None, diff --git a/megatron/core/models/multimodal/llava_model.py b/megatron/core/models/multimodal/llava_model.py index af0bcf6e9fd..d6fab90b00a 100644 --- a/megatron/core/models/multimodal/llava_model.py +++ b/megatron/core/models/multimodal/llava_model.py @@ -17,8 +17,10 @@ from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer import MegatronModule +from megatron.core.transformer.attention import SelfAttentionSubmodules from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.transformer.transformer_layer import TransformerLayerSubmodules from megatron.core.utils import deprecate_inference_params, log_single_rank try: @@ -158,9 +160,17 @@ def __init__( self.context_parallel_lm = language_transformer_config.context_parallel_size if self.sequence_parallel_lm or self.context_parallel_lm > 1: if not language_model_type.startswith('nemotron5-hybrid'): - attn_module = language_transformer_layer_spec.submodules.self_attention + assert isinstance( + language_transformer_layer_spec.submodules, TransformerLayerSubmodules + ) + assert isinstance( + language_transformer_layer_spec.submodules.self_attention.submodules, + SelfAttentionSubmodules, + ) assert ( - attn_module.submodules.core_attention == TEDotProductAttention and HAVE_TE + language_transformer_layer_spec.submodules.self_attention.submodules.core_attention + == TEDotProductAttention + and HAVE_TE ), "Sequence/Context Parallelism is supported only with TE DotProductAttention." if self.context_parallel_lm > 1: self.cp_group = self.pg_collection.cp diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index 2edb652bfc6..d8dfb84d0fe 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -3,7 +3,7 @@ import warnings from contextlib import nullcontext from dataclasses import dataclass -from typing import Callable, List, Optional, Union +from typing import Callable, List, Optional, Union, cast import torch from torch import Tensor @@ -24,7 +24,10 @@ from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.transformer_block import TransformerBlockSubmodules from megatron.core.transformer.transformer_config import TransformerConfig -from megatron.core.transformer.transformer_layer import get_transformer_layer_offset +from megatron.core.transformer.transformer_layer import ( + TransformerLayerSubmodules, + get_transformer_layer_offset, +) from megatron.core.utils import ( get_pg_rank, is_torch_min_version, @@ -621,6 +624,8 @@ def __init__( self.vp_stage = vp_stage self.cp_group = pg_collection.cp + assert isinstance(self.submodules.transformer_layer.submodules, TransformerLayerSubmodules) + self_attention_spec = self.submodules.transformer_layer.submodules.self_attention attn_mask_type = self_attention_spec.params.get('attn_mask_type', '') assert attn_mask_type in SUPPORTED_ATTN_MASK, ( diff --git a/megatron/core/typed_torch.py b/megatron/core/typed_torch.py index bcbf388facc..5661e2e9b98 100644 --- a/megatron/core/typed_torch.py +++ b/megatron/core/typed_torch.py @@ -2,8 +2,9 @@ """Utilities for improved type hinting with torch interfaces.""" from __future__ import annotations +import inspect from collections.abc import Callable -from typing import Generic, ParamSpec, Protocol, TypeVar +from typing import Any, Concatenate, Generic, Literal, ParamSpec, Protocol, TypeVar, overload import torch @@ -48,3 +49,154 @@ def not_none(value: T | None) -> T: if value is None: raise ValueError('Expected value to be not None') return value + + +R_src = TypeVar('R_src') +R_dst = TypeVar('R_dst') +P_src = ParamSpec('P_src') +P_dst = ParamSpec('P_dst') +First_dst = TypeVar('First_dst') + + +@overload +def copy_signature( + source: Callable[P_src, Any], + /, + *, + handle_return_type: Literal['preserve'] = 'preserve', + handle_first_src_param: Literal['copy'] = 'copy', + handle_first_dst_param: Literal['drop'] = 'drop', +) -> Callable[[Callable[..., R_dst]], Callable[P_src, R_dst]]: ... + + +@overload +def copy_signature( + source: Callable[P_src, R_src], + /, + *, + handle_return_type: Literal['overwrite'], + handle_first_src_param: Literal['copy'] = 'copy', + handle_first_dst_param: Literal['drop'] = 'drop', +) -> Callable[[Callable[..., Any]], Callable[P_src, R_src]]: ... + + +@overload +def copy_signature( + source: Callable[Concatenate[Any, P_src], Any], + /, + *, + handle_return_type: Literal['preserve'] = 'preserve', + handle_first_src_param: Literal['skip'], + handle_first_dst_param: Literal['drop'] = 'drop', +) -> Callable[[Callable[..., R_dst]], Callable[P_src, R_dst]]: ... + + +@overload +def copy_signature( + source: Callable[Concatenate[Any, P_src], R_src], + /, + *, + handle_return_type: Literal['overwrite'], + handle_first_src_param: Literal['skip'], + handle_first_dst_param: Literal['drop'] = 'drop', +) -> Callable[[Callable[..., Any]], Callable[P_src, R_src]]: ... + + +@overload +def copy_signature( + source: Callable[P_src, Any], + /, + *, + handle_return_type: Literal['preserve'] = 'preserve', + handle_first_src_param: Literal['copy'] = 'copy', + handle_first_dst_param: Literal['preserve'], +) -> Callable[ + [Callable[Concatenate[First_dst, ...], R_dst]], Callable[Concatenate[First_dst, P_src], R_dst] +]: ... + + +@overload +def copy_signature( + source: Callable[P_src, R_src], + /, + *, + handle_return_type: Literal['overwrite'], + handle_first_src_param: Literal['copy'] = 'copy', + handle_first_dst_param: Literal['preserve'], +) -> Callable[ + [Callable[Concatenate[First_dst, ...], Any]], Callable[Concatenate[First_dst, P_src], R_src] +]: ... + + +@overload +def copy_signature( + source: Callable[Concatenate[Any, P_src], Any], + /, + *, + handle_return_type: Literal['preserve'] = 'preserve', + handle_first_src_param: Literal['skip'], + handle_first_dst_param: Literal['preserve'], +) -> Callable[ + [Callable[Concatenate[First_dst, ...], R_dst]], Callable[Concatenate[First_dst, P_src], R_dst] +]: ... + + +@overload +def copy_signature( + source: Callable[Concatenate[Any, P_src], R_src], + /, + *, + handle_return_type: Literal['overwrite'], + handle_first_src_param: Literal['skip'], + handle_first_dst_param: Literal['preserve'], +) -> Callable[ + [Callable[Concatenate[First_dst, ...], Any]], Callable[Concatenate[First_dst, P_src], R_src] +]: ... + + +def copy_signature( + source: Callable[..., Any], + /, + *, + handle_return_type: Literal['preserve', 'overwrite'] = 'preserve', + handle_first_src_param: Literal['copy', 'skip'] = 'copy', + handle_first_dst_param: Literal['preserve', 'drop'] = 'drop', +): + """Decorator to copy the signature from one function to another. + + Args: + source: The function or callable from which to copy the signature. + handle_return_type: How to handle the return type annotation. + 'preserve' to keep the decorated function's return type, + 'overwrite' to use the source function's return type. + handle_first_src_param: How to handle the first parameter of the source function. + 'copy' to include it in the decorated function's signature, + 'skip' to exclude it. Useful for removing 'self' or 'cls'. + handle_first_dst_param: How to handle the first parameter of the decorated function. + 'preserve' to keep it in the decorated function's signature, + 'drop' to exclude it. Useful for preserving 'self' or 'cls'. + + Returns: + A decorator that copies the signature from `func` to the decorated function. + """ + source_signature = inspect.signature(source) + + def decorator(decorated: Callable[..., Any], /) -> Callable[..., Any]: + dest_signature = inspect.signature(decorated) + new_params = [] + if handle_first_dst_param == 'preserve': + new_params.append(next(iter(dest_signature.parameters.values()))) + src_params_iter = iter(source_signature.parameters.values()) + if handle_first_src_param == 'skip': + next(src_params_iter) + new_params.extend(src_params_iter) + new_signature = dest_signature.replace(parameters=new_params) + if handle_return_type == 'overwrite': + new_signature = new_signature.replace( + return_annotation=source_signature.return_annotation + ) + + decorated.__signature__ = new_signature # type: ignore + return decorated + + return decorator diff --git a/tests/functional_tests/test_cases/common/ckpt_converter/__main__.py b/tests/functional_tests/test_cases/common/ckpt_converter/__main__.py index 62d86a29358..fd4d58c2747 100644 --- a/tests/functional_tests/test_cases/common/ckpt_converter/__main__.py +++ b/tests/functional_tests/test_cases/common/ckpt_converter/__main__.py @@ -19,7 +19,6 @@ from megatron.core import parallel_state from megatron.core.datasets.gpt_dataset import _get_ltor_masks_and_position_ids from megatron.core.enums import ModelType -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec from megatron.core.models.multimodal.llava_model import DEFAULT_IMAGE_TOKEN_INDEX, LLaVAModel from megatron.core.models.vision.vit_layer_specs import get_vit_layer_with_transformer_engine_spec from megatron.core.pipeline_parallel import get_forward_backward_func diff --git a/tests/functional_tests/test_cases/common/moe_perf/__main__.py b/tests/functional_tests/test_cases/common/moe_perf/__main__.py index ace44c7ca4f..3fabf6e0236 100644 --- a/tests/functional_tests/test_cases/common/moe_perf/__main__.py +++ b/tests/functional_tests/test_cases/common/moe_perf/__main__.py @@ -18,8 +18,7 @@ from megatron.core.config import set_experimental_flag from megatron.core.fp8_utils import get_fp8_context from megatron.core.models.gpt.gpt_layer_specs import ( - get_gpt_layer_local_spec, - get_gpt_layer_with_transformer_engine_spec, + get_gpt_layer_with_transformer_engine_submodules, ) from megatron.core.transformer.moe.fused_a2a import HAVE_DEEP_EP, HAVE_HYBRIDEP from megatron.core.transformer.moe.moe_layer import MoELayer @@ -89,10 +88,9 @@ def _build_transformer_config(case: MoEPerformanceCase) -> TransformerConfig: # NOTE: Only TE backend is covered in this test. def _resolve_moe_submodules(case: MoEPerformanceCase): - layer_spec = get_gpt_layer_with_transformer_engine_spec( + return get_gpt_layer_with_transformer_engine_submodules( num_experts=case.model.num_experts, moe_grouped_gemm=True - ) - return layer_spec.submodules.mlp.submodules + ).mlp.submodules def _load_baselines() -> Dict[str, Dict[str, float]]: diff --git a/tests/unit_tests/dist_checkpointing/models/test_mlp_glu.py b/tests/unit_tests/dist_checkpointing/models/test_mlp_glu.py index 0970e2adc8a..037e368ea2f 100644 --- a/tests/unit_tests/dist_checkpointing/models/test_mlp_glu.py +++ b/tests/unit_tests/dist_checkpointing/models/test_mlp_glu.py @@ -13,7 +13,9 @@ get_param_id_to_sharded_param_map, optim_state_to_sharding_state, ) -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, +) from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.mlp import MLP, apply_swiglu_sharded_factory from megatron.core.transformer.transformer_config import TransformerConfig @@ -32,7 +34,7 @@ def initialize_mlp(glu=True): gated_linear_unit=glu, ) return MLP( - transformer_config, get_gpt_layer_with_transformer_engine_spec().submodules.mlp.submodules + transformer_config, get_gpt_layer_with_transformer_engine_submodules().mlp.submodules ) diff --git a/tests/unit_tests/dist_checkpointing/models/test_moe_experts.py b/tests/unit_tests/dist_checkpointing/models/test_moe_experts.py index ca546d746af..3b04f73ee2e 100644 --- a/tests/unit_tests/dist_checkpointing/models/test_moe_experts.py +++ b/tests/unit_tests/dist_checkpointing/models/test_moe_experts.py @@ -1,6 +1,4 @@ # Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. -import os - import pytest import torch from transformer_engine.pytorch.fp8 import check_fp8_support, fp8_autocast @@ -17,11 +15,13 @@ FullyParallelSaveStrategyWrapper, ) from megatron.core.models.gpt.gpt_layer_specs import ( - get_gpt_layer_local_spec, - get_gpt_layer_with_transformer_engine_spec, + get_gpt_layer_local_submodules, + get_gpt_layer_with_transformer_engine_submodules, ) from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.mlp import MLPSubmodules from megatron.core.transformer.moe.experts import GroupedMLP, SequentialMLP, TEGroupedMLP +from megatron.core.transformer.moe.moe_layer import MoESubmodules from megatron.core.transformer.moe.moe_utils import get_default_pg_collection from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import is_te_min_version @@ -54,33 +54,39 @@ def initialize_expert_layer(seed, glu=True, expert_type='sequential', fp8=False, if expert_type == 'grouped': model = GroupedMLP(num_local_experts, transformer_config, pg_collection) elif expert_type == 'te_grouped': - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + layer_submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=num_moe_experts, moe_grouped_gemm=True ) + assert isinstance(layer_submodules.mlp.submodules, MoESubmodules) + assert isinstance(layer_submodules.mlp.submodules.experts.submodules, MLPSubmodules) model = TEGroupedMLP( num_local_experts, transformer_config, - transformer_layer_spec.submodules.mlp.submodules.experts.submodules, + layer_submodules.mlp.submodules.experts.submodules, pg_collection, ) elif expert_type == 'sequential': - transformer_layer_spec = get_gpt_layer_local_spec( + layer_submodules = get_gpt_layer_local_submodules( num_experts=num_moe_experts, moe_grouped_gemm=False ) + assert isinstance(layer_submodules.mlp.submodules, MoESubmodules) + assert isinstance(layer_submodules.mlp.submodules.experts.submodules, MLPSubmodules) model = SequentialMLP( num_local_experts, transformer_config, - transformer_layer_spec.submodules.mlp.submodules.experts.submodules, + layer_submodules.mlp.submodules.experts.submodules, pg_collection, ) elif expert_type == 'te_sequential': - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + layer_submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=num_moe_experts, moe_grouped_gemm=False ) + assert isinstance(layer_submodules.mlp.submodules, MoESubmodules) + assert isinstance(layer_submodules.mlp.submodules.experts.submodules, MLPSubmodules) model = SequentialMLP( num_local_experts, transformer_config, - transformer_layer_spec.submodules.mlp.submodules.experts.submodules, + layer_submodules.mlp.submodules.experts.submodules, pg_collection, ) else: diff --git a/tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.py b/tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.py index e83f7142284..c110b9263e2 100644 --- a/tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.py +++ b/tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.py @@ -9,7 +9,9 @@ from megatron.core import parallel_state from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig from megatron.core.distributed.param_and_grad_buffer import partition_buckets -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, +) from megatron.core.transformer import TransformerConfig from megatron.core.transformer.moe.moe_layer import MoELayer from tests.unit_tests.test_utilities import TestModel, Utils @@ -41,15 +43,13 @@ def __init__( params_dtype=torch.bfloat16, add_bias_linear=False, ) - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=num_moe_experts, moe_grouped_gemm=moe_grouped_gemm ) super().__init__() self.layers = torch.nn.ModuleList( [ - MoELayer( - transformer_config, transformer_layer_spec.submodules.mlp.submodules - ).cuda() + MoELayer(transformer_config, submodules.mlp.submodules).cuda() for _ in range(num_layers) ] ) diff --git a/tests/unit_tests/inference/model_inference_wrappers/gpt/test_gpt_inference_wrapper.py b/tests/unit_tests/inference/model_inference_wrappers/gpt/test_gpt_inference_wrapper.py index 07afebe1067..8d460f1bc24 100644 --- a/tests/unit_tests/inference/model_inference_wrappers/gpt/test_gpt_inference_wrapper.py +++ b/tests/unit_tests/inference/model_inference_wrappers/gpt/test_gpt_inference_wrapper.py @@ -13,10 +13,7 @@ from megatron.core.inference.model_inference_wrappers.inference_wrapper_config import ( InferenceWrapperConfig, ) -from megatron.core.models.gpt.gpt_layer_specs import ( - get_gpt_layer_local_spec, - get_gpt_layer_with_transformer_engine_spec, -) +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec from megatron.core.models.gpt.gpt_model import GPTModel from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.transformer_config import TransformerConfig diff --git a/tests/unit_tests/inference/text_generation_controllers/test_vlm_text_generation_controller.py b/tests/unit_tests/inference/text_generation_controllers/test_vlm_text_generation_controller.py index 31bf415ba56..653baf76b8a 100644 --- a/tests/unit_tests/inference/text_generation_controllers/test_vlm_text_generation_controller.py +++ b/tests/unit_tests/inference/text_generation_controllers/test_vlm_text_generation_controller.py @@ -23,12 +23,14 @@ from megatron.core.inference.text_generation_controllers.vlm_text_generation_controller import ( VLMTextGenerationController, ) -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules from megatron.core.models.multimodal.llava_model import LLaVAModel from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.enums import AttnBackend from megatron.core.transformer.module import Float16Module +from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.transformer.transformer_layer import TransformerLayer from tests.unit_tests.test_utilities import Utils @@ -69,15 +71,19 @@ def setup_method(self, method): bf16=True, ) - language_layer_spec = get_gpt_layer_local_spec() - vision_layer_spec = copy.deepcopy(language_layer_spec) - vision_projection_spec = copy.deepcopy(language_layer_spec.submodules.mlp.submodules) + language_layer_submodules = get_gpt_layer_local_submodules() + vision_layer_spec = ModuleSpec( + module=TransformerLayer, submodules=copy.deepcopy(language_layer_submodules) + ) + vision_projection_spec = copy.deepcopy(language_layer_submodules.mlp.submodules) language_config.language_model_type = "dummy" vision_config.vision_model_type = "clip" self.model = LLaVAModel( language_transformer_config=language_config, - language_transformer_layer_spec=language_layer_spec, + language_transformer_layer_spec=ModuleSpec( + module=TransformerLayer, submodules=language_layer_submodules + ), language_vocab_size=self.language_vocab_size, language_max_sequence_length=self.language_max_sequence_length, vision_transformer_config=vision_config, diff --git a/tests/unit_tests/models/test_bert_model.py b/tests/unit_tests/models/test_bert_model.py index c878fd4d3b5..e4841c5ec45 100644 --- a/tests/unit_tests/models/test_bert_model.py +++ b/tests/unit_tests/models/test_bert_model.py @@ -10,12 +10,15 @@ from megatron.core.models.bert.bert_layer_specs import ( bert_layer_local_spec, - bert_layer_with_transformer_engine_spec, + get_bert_layer_with_transformer_engine_spec, + get_bert_layer_with_transformer_engine_submodules, ) from megatron.core.models.bert.bert_model import BertModel from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.enums import AttnBackend, AttnMaskType +from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.transformer.transformer_layer import TransformerLayer from tests.unit_tests.test_utilities import Utils @@ -40,7 +43,7 @@ def setup_method(self, method): self.bert_model = BertModel( config=transformer_config, num_tokentypes=0, - transformer_layer_spec=bert_layer_with_transformer_engine_spec, + transformer_layer_spec=get_bert_layer_with_transformer_engine_spec(), vocab_size=100, max_sequence_length=4, ) @@ -112,7 +115,7 @@ def setup_method(self, method): self.bert_model = BertModel( config=self.transformer_config, num_tokentypes=0, - transformer_layer_spec=bert_layer_with_transformer_engine_spec, + transformer_layer_spec=get_bert_layer_with_transformer_engine_spec(), vocab_size=100, max_sequence_length=4, ) @@ -139,16 +142,15 @@ def test_local_spec_exception(self, mocker): @pytest.mark.internal def test_transformer_engine_version_1_10(self, mocker): - bert_layer_with_transformer_engine_spec.submodules.self_attention.params[ - 'attn_mask_type' - ] == AttnMaskType.arbitrary + submodules = get_bert_layer_with_transformer_engine_submodules() + submodules.self_attention.params['attn_mask_type'] = AttnMaskType.arbitrary mocker.patch("megatron.core.utils.get_te_version", return_value=PkgVersion("1.10")) - self.bert_model.transformer_layer_spec = bert_layer_with_transformer_engine_spec + self.bert_model.transformer_layer_spec = ModuleSpec( + module=TransformerLayer, submodules=submodules + ) attn_mask_dimensions = self.bert_model._sanity_check_attention_and_get_attn_mask_dimension() - attn_mask_type = self.bert_model.transformer_layer_spec.submodules.self_attention.params[ - 'attn_mask_type' - ] + attn_mask_type = submodules.self_attention.params['attn_mask_type'] assert ( attn_mask_type == AttnMaskType.padding ), f"Exepcted attn mask type to be padding, but got {attn_mask_type}" @@ -160,7 +162,7 @@ def test_transformer_engine_version_1_10(self, mocker): def test_transformer_engine_version_1_7_to_1_10_flash_attn(self, mocker): self.bert_model.config.attention_backend = AttnBackend.flash mocker.patch("megatron.core.utils.get_te_version", return_value=PkgVersion("1.8")) - self.bert_model.transformer_layer_spec = bert_layer_with_transformer_engine_spec + self.bert_model.transformer_layer_spec = get_bert_layer_with_transformer_engine_spec() attn_mask_dimensions = self.bert_model._sanity_check_attention_and_get_attn_mask_dimension() assert ( attn_mask_dimensions == "b11s" @@ -170,15 +172,14 @@ def test_transformer_engine_version_1_7_to_1_10_flash_attn(self, mocker): @pytest.mark.flaky @pytest.mark.flaky_in_dev def test_transformer_engine_version_1_7_to_1_10_rng_error(self, mocker): - bert_layer_with_transformer_engine_spec.submodules.self_attention.params[ - 'attn_mask_type' - ] == AttnMaskType.padding + submodules = get_bert_layer_with_transformer_engine_submodules() + submodules.self_attention.params['attn_mask_type'] = AttnMaskType.padding mocker.patch("megatron.core.utils.get_te_version", return_value=PkgVersion("1.8")) with pytest.raises(Exception) as exc_info: self.bert_model = BertModel( config=self.transformer_config, num_tokentypes=0, - transformer_layer_spec=bert_layer_with_transformer_engine_spec, + transformer_layer_spec=ModuleSpec(module=TransformerLayer, submodules=submodules), vocab_size=100, max_sequence_length=4, ) @@ -191,15 +192,14 @@ def test_transformer_engine_version_1_7_to_1_10_rng_error(self, mocker): @pytest.mark.internal def test_transformer_engine_version_1_7_to_1_10_unfused_attention(self, mocker): self.bert_model.config.attention_backend = AttnBackend.unfused - bert_layer_with_transformer_engine_spec.submodules.self_attention.params[ - 'attn_mask_type' - ] == AttnMaskType.padding + submodules = get_bert_layer_with_transformer_engine_submodules() + submodules.self_attention.params['attn_mask_type'] = AttnMaskType.padding mocker.patch("megatron.core.utils.get_te_version", return_value=PkgVersion("1.8")) - self.bert_model.transformer_layer_spec = bert_layer_with_transformer_engine_spec + self.bert_model.transformer_layer_spec = ModuleSpec( + module=TransformerLayer, submodules=submodules + ) attn_mask_dimensions = self.bert_model._sanity_check_attention_and_get_attn_mask_dimension() - attn_mask_type = self.bert_model.transformer_layer_spec.submodules.self_attention.params[ - 'attn_mask_type' - ] + attn_mask_type = submodules.self_attention.params['attn_mask_type'] assert ( attn_mask_type == AttnMaskType.arbitrary ), f"Exepcted attn mask type to be arbitrary, but got {attn_mask_type}" @@ -218,7 +218,7 @@ def test_transformer_engine_version_less_than_1_7(self, mocker): self.bert_model = BertModel( config=self.transformer_config, num_tokentypes=0, - transformer_layer_spec=bert_layer_with_transformer_engine_spec, + transformer_layer_spec=get_bert_layer_with_transformer_engine_spec(), vocab_size=100, max_sequence_length=4, ) diff --git a/tests/unit_tests/models/test_llava_model.py b/tests/unit_tests/models/test_llava_model.py index cee6d2b0b27..685e20a58dd 100644 --- a/tests/unit_tests/models/test_llava_model.py +++ b/tests/unit_tests/models/test_llava_model.py @@ -8,14 +8,18 @@ from megatron.core import parallel_state as ps from megatron.core.inference.contexts import StaticInferenceContext -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, +) from megatron.core.models.multimodal import context_parallel from megatron.core.models.multimodal.llava_model import LLaVAModel from megatron.core.models.vision.vit_layer_specs import get_vit_layer_with_transformer_engine_spec from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.enums import AttnMaskType +from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.transformer.transformer_layer import TransformerLayer from megatron.core.utils import is_te_min_version from megatron.training.global_vars import set_args from tests.unit_tests.test_utilities import Utils @@ -47,15 +51,19 @@ def setup_method(self, method): use_cpu_initialization=False, ) - language_layer_spec = get_gpt_layer_with_transformer_engine_spec() - vision_layer_spec = deepcopy(language_layer_spec) - vision_projection_spec = deepcopy(language_layer_spec.submodules.mlp.submodules) + language_layer_submodules = get_gpt_layer_with_transformer_engine_submodules() + vision_layer_spec = ModuleSpec( + module=TransformerLayer, submodules=deepcopy(language_layer_submodules) + ) + vision_projection_spec = deepcopy(language_layer_submodules.mlp.submodules) language_config.language_model_type = "dummy" vision_config.vision_model_type = "clip" self.model = LLaVAModel( language_transformer_config=language_config, - language_transformer_layer_spec=language_layer_spec, + language_transformer_layer_spec=ModuleSpec( + module=TransformerLayer, submodules=language_layer_submodules + ), language_vocab_size=8192, language_max_sequence_length=4096, vision_transformer_config=vision_config, @@ -481,16 +489,20 @@ def setup_and_teardown_llava_model(request): use_cpu_initialization=False, ) - language_layer_spec = get_gpt_layer_with_transformer_engine_spec() - vision_layer_spec = deepcopy(language_layer_spec) - vision_projection_spec = deepcopy(language_layer_spec.submodules.mlp.submodules) + language_layer_submodules = get_gpt_layer_with_transformer_engine_submodules() + vision_layer_spec = ModuleSpec( + module=TransformerLayer, submodules=deepcopy(language_layer_submodules) + ) + vision_projection_spec = deepcopy(language_layer_submodules.mlp.submodules) language_config.language_model_type = "dummy" vision_model_type = request.param vision_config.vision_model_type = vision_model_type model = LLaVAModel( language_transformer_config=language_config, - language_transformer_layer_spec=language_layer_spec, + language_transformer_layer_spec=ModuleSpec( + module=TransformerLayer, submodules=language_layer_submodules + ), language_vocab_size=2048, language_max_sequence_length=4096, vision_transformer_config=vision_config, @@ -573,31 +585,33 @@ def _init_llava_model(self, cp_size, tp_size, sequence_parallel): context_parallel_size=1, ) - language_layer_spec = get_gpt_layer_with_transformer_engine_spec() + language_layer_submodules = get_gpt_layer_with_transformer_engine_submodules() # SP/CP either requires user to ensure token lengths do not require padding OR change mask type to padding if ( - language_layer_spec.submodules.self_attention.params.get('attn_mask_type', '') + language_layer_submodules.self_attention.params.get('attn_mask_type', '') == AttnMaskType.causal ): - language_layer_spec.submodules.self_attention.params['attn_mask_type'] = ( + language_layer_submodules.self_attention.params['attn_mask_type'] = ( AttnMaskType.padding_causal ) elif ( - language_layer_spec.submodules.self_attention.params.get('attn_mask_type', '') + language_layer_submodules.self_attention.params.get('attn_mask_type', '') == AttnMaskType.no_mask ): - language_layer_spec.submodules.self_attention.params['attn_mask_type'] = ( - AttnMaskType.padding - ) + language_layer_submodules.self_attention.params['attn_mask_type'] = AttnMaskType.padding - vision_layer_spec = deepcopy(language_layer_spec) - vision_projection_spec = deepcopy(language_layer_spec.submodules.mlp.submodules) + vision_layer_spec = ModuleSpec( + module=TransformerLayer, submodules=deepcopy(language_layer_submodules) + ) + vision_projection_spec = deepcopy(language_layer_submodules.mlp.submodules) language_config.language_model_type = "dummy" vision_config.vision_model_type = "clip" model = LLaVAModel( language_transformer_config=language_config, - language_transformer_layer_spec=language_layer_spec, + language_transformer_layer_spec=ModuleSpec( + module=TransformerLayer, submodules=language_layer_submodules + ), language_vocab_size=8192, language_max_sequence_length=4096, vision_transformer_config=vision_config, diff --git a/tests/unit_tests/ssm/test_mamba_layer.py b/tests/unit_tests/ssm/test_mamba_layer.py index 25d24443d38..26c7ba4d92f 100644 --- a/tests/unit_tests/ssm/test_mamba_layer.py +++ b/tests/unit_tests/ssm/test_mamba_layer.py @@ -5,7 +5,8 @@ from megatron.core.models.mamba.mamba_layer_specs import mamba_stack_spec from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.ssm.mamba_layer import MambaLayer +from megatron.core.ssm.mamba_block import MambaStackSubmodules +from megatron.core.ssm.mamba_layer import MambaLayer, MambaLayerSubmodules from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import TransformerConfig from tests.unit_tests.test_utilities import Utils @@ -25,9 +26,14 @@ def setup_method(self, method): num_attention_heads=1, use_cpu_initialization=True, ) - modules = mamba_stack_spec.submodules.mamba_layer.submodules + assert isinstance(mamba_stack_spec.submodules, MambaStackSubmodules) + assert isinstance(mamba_stack_spec.submodules.mamba_layer.submodules, MambaLayerSubmodules) pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp']) - self.layer = MambaLayer(transformer_config, modules, pg_collection=pg_collection) + self.layer = MambaLayer( + transformer_config, + mamba_stack_spec.submodules.mamba_layer.submodules, + pg_collection=pg_collection, + ) def teardown_method(self, method): Utils.destroy_model_parallel() diff --git a/tests/unit_tests/ssm/test_mamba_mixer.py b/tests/unit_tests/ssm/test_mamba_mixer.py index 587cfc66967..655526fe794 100644 --- a/tests/unit_tests/ssm/test_mamba_mixer.py +++ b/tests/unit_tests/ssm/test_mamba_mixer.py @@ -6,7 +6,9 @@ from megatron.core.inference.contexts.static_context import StaticInferenceContext from megatron.core.models.mamba.mamba_layer_specs import mamba_stack_spec from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.ssm.mamba_mixer import MambaMixer +from megatron.core.ssm.mamba_block import MambaStackSubmodules +from megatron.core.ssm.mamba_layer import MambaLayerSubmodules +from megatron.core.ssm.mamba_mixer import MambaMixer, MambaMixerSubmodules from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import TransformerConfig from tests.unit_tests.test_utilities import Utils @@ -36,11 +38,16 @@ def get_mixer(self, tp_size=1, cp_size=1, use_mem_eff_path=True): num_attention_heads=1, use_cpu_initialization=True, ) - modules = mamba_stack_spec.submodules.mamba_layer.submodules.mixer.submodules + assert isinstance(mamba_stack_spec.submodules, MambaStackSubmodules) + assert isinstance(mamba_stack_spec.submodules.mamba_layer.submodules, MambaLayerSubmodules) + assert isinstance( + mamba_stack_spec.submodules.mamba_layer.submodules.mixer.submodules, + MambaMixerSubmodules, + ) pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp']) mixer = MambaMixer( transformer_config, - modules, + mamba_stack_spec.submodules.mamba_layer.submodules.mixer.submodules, transformer_config.hidden_size, layer_number=1, use_mem_eff_path=use_mem_eff_path, @@ -119,12 +126,17 @@ def test_error_check(self, hidden_size, ngroups, tp_size, expected_error_message use_cpu_initialization=True, mamba_num_groups=ngroups, ) - submodules = mamba_stack_spec.submodules.mamba_layer.submodules.mixer.submodules + assert isinstance(mamba_stack_spec.submodules, MambaStackSubmodules) + assert isinstance(mamba_stack_spec.submodules.mamba_layer.submodules, MambaLayerSubmodules) + assert isinstance( + mamba_stack_spec.submodules.mamba_layer.submodules.mixer.submodules, + MambaMixerSubmodules, + ) pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp']) with pytest.raises(AssertionError, match=expected_error_message): MambaMixer( transformer_config, - submodules, + mamba_stack_spec.submodules.mamba_layer.submodules.mixer.submodules, transformer_config.hidden_size, pg_collection=pg_collection, ) diff --git a/tests/unit_tests/test_typed_torch.py b/tests/unit_tests/test_typed_torch.py new file mode 100644 index 00000000000..724749dd79a --- /dev/null +++ b/tests/unit_tests/test_typed_torch.py @@ -0,0 +1,183 @@ +import inspect +from typing import Any + +import pytest + +from megatron.core.typed_torch import copy_signature, not_none + + +def source_func(a: int, *, b: str) -> str: + """Sample function to copy the signature from.""" + return str(a) + b + + +class SourceClass: + """Sample class with a method to copy the signature from.""" + + def method(self, a: int, *, b: str) -> str: + """Sample method to copy the signature from.""" + return str(a) + b + + +@copy_signature(source_func) +def dest_func_from_func(*args: Any, **kwargs: Any) -> list[str]: + """Function with copied signature from source_func.""" + return [source_func(*args, **kwargs)] + + +@copy_signature(source_func, handle_return_type='overwrite') +def dest_func_from_func_overwrite(*args: Any, **kwargs: Any) -> object: + """Function with copied signature from source_func, but overwritten return type.""" + return source_func(*args, **kwargs) + + +@copy_signature(SourceClass.method, handle_first_src_param='skip') +def dest_func_from_method(*args: Any, **kwargs: Any) -> int: + """Function with copied signature from SourceClass.method.""" + return len(SourceClass().method(*args, **kwargs)) + + +@copy_signature(SourceClass.method, handle_return_type='overwrite', handle_first_src_param='skip') +def dest_func_from_method_overwrite(*args: Any, **kwargs: Any) -> object: + """Function with copied signature from SourceClass.method, but overwritten return type.""" + return SourceClass().method(*args, **kwargs) + + +class DestClass: + """Class with methods that have copied signatures.""" + + @copy_signature(source_func, handle_first_dst_param='preserve') + def dest_method_from_func(self, *args: Any, **kwargs: Any) -> list[str]: + """Method with copied signature from source_func.""" + return [source_func(*args, **kwargs)] + + @copy_signature(source_func, handle_return_type='overwrite', handle_first_dst_param='preserve') + def dest_method_from_func_overwrite(self, *args: Any, **kwargs: Any) -> object: + """Method with copied signature from source_func, but overwritten return type.""" + return source_func(*args, **kwargs) + + @classmethod + @copy_signature( + SourceClass.method, handle_first_src_param='skip', handle_first_dst_param='preserve' + ) + def dest_method_from_method(cls, *args: Any, **kwargs: Any) -> int: + """Class method with copied signature from SourceClass.method.""" + return len(SourceClass().method(*args, **kwargs)) + + @copy_signature( + SourceClass.method, + handle_return_type='overwrite', + handle_first_src_param='skip', + handle_first_dst_param='preserve', + ) + def dest_method_from_method_overwrite(self, *args: Any, **kwargs: Any) -> object: + """Method with copied signature from SourceClass.method, but overwritten return type.""" + return SourceClass().method(*args, **kwargs) + + +class TestCopySignature: + def test_original_return_type(self): + """Test that the original return types are preserved.""" + f2f: list[str] = dest_func_from_func(1, b='a') + assert f2f == ['1a'] + assert inspect.signature(dest_func_from_func) == inspect.Signature( + [ + inspect.Parameter('a', inspect.Parameter.POSITIONAL_OR_KEYWORD, annotation=int), + inspect.Parameter('b', inspect.Parameter.KEYWORD_ONLY, annotation=str), + ], + return_annotation=list[str], + ) + + m2f: int = dest_func_from_method(1, b='a') + assert m2f == 2 + assert inspect.signature(dest_func_from_method) == inspect.Signature( + [ + inspect.Parameter('a', inspect.Parameter.POSITIONAL_OR_KEYWORD, annotation=int), + inspect.Parameter('b', inspect.Parameter.KEYWORD_ONLY, annotation=str), + ], + return_annotation=int, + ) + + f2m: list[str] = DestClass().dest_method_from_func( + 1, b='a' + ) + DestClass.dest_method_from_func(DestClass(), 1, b='a') + assert f2m == ['1a', '1a'] + assert inspect.signature(DestClass.dest_method_from_func) == inspect.Signature( + [ + inspect.Parameter('self', inspect.Parameter.POSITIONAL_OR_KEYWORD), + inspect.Parameter('a', inspect.Parameter.POSITIONAL_OR_KEYWORD, annotation=int), + inspect.Parameter('b', inspect.Parameter.KEYWORD_ONLY, annotation=str), + ], + return_annotation=list[str], + ) + + m2m: int = DestClass.dest_method_from_method(1, b='a') + assert m2m == 2 + assert inspect.signature(DestClass.dest_method_from_method) == inspect.Signature( + [ + inspect.Parameter('a', inspect.Parameter.POSITIONAL_OR_KEYWORD, annotation=int), + inspect.Parameter('b', inspect.Parameter.KEYWORD_ONLY, annotation=str), + ], + return_annotation=int, + ) + + def test_overwritten_return_type(self): + """Test that the return types are overwritten correctly.""" + f2f: str = dest_func_from_func_overwrite(1, b='a') + assert f2f == '1a' + assert inspect.signature(dest_func_from_func_overwrite) == inspect.Signature( + [ + inspect.Parameter('a', inspect.Parameter.POSITIONAL_OR_KEYWORD, annotation=int), + inspect.Parameter('b', inspect.Parameter.KEYWORD_ONLY, annotation=str), + ], + return_annotation=str, + ) + + m2f: str = dest_func_from_method_overwrite(1, b='a') + assert m2f == '1a' + assert inspect.signature(dest_func_from_method_overwrite) == inspect.Signature( + [ + inspect.Parameter('a', inspect.Parameter.POSITIONAL_OR_KEYWORD, annotation=int), + inspect.Parameter('b', inspect.Parameter.KEYWORD_ONLY, annotation=str), + ], + return_annotation=str, + ) + + f2m: str = DestClass().dest_method_from_func_overwrite( + 1, b='a' + ) + DestClass.dest_method_from_func_overwrite(DestClass(), 1, b='a') + assert f2m == '1a1a' + assert inspect.signature(DestClass.dest_method_from_func_overwrite) == inspect.Signature( + [ + inspect.Parameter('self', inspect.Parameter.POSITIONAL_OR_KEYWORD), + inspect.Parameter('a', inspect.Parameter.POSITIONAL_OR_KEYWORD, annotation=int), + inspect.Parameter('b', inspect.Parameter.KEYWORD_ONLY, annotation=str), + ], + return_annotation=str, + ) + + m2m: str = DestClass().dest_method_from_method_overwrite(1, b='a') + assert m2m == '1a' + assert inspect.signature(DestClass.dest_method_from_method_overwrite) == inspect.Signature( + [ + inspect.Parameter('self', inspect.Parameter.POSITIONAL_OR_KEYWORD), + inspect.Parameter('a', inspect.Parameter.POSITIONAL_OR_KEYWORD, annotation=int), + inspect.Parameter('b', inspect.Parameter.KEYWORD_ONLY, annotation=str), + ], + return_annotation=str, + ) + + +class TestNotNone: + """Tests not_none.""" + + def test_none(self): + """Test that passing None raises a ValueError.""" + with pytest.raises(ValueError, match=r'Expected value to be not None'): + not_none(None) + + def test_not_none(self): + """Test that passing a non-None value returns the value.""" + value = 42 + result = not_none(value) + assert result == value diff --git a/tests/unit_tests/test_utils.py b/tests/unit_tests/test_utils.py index fddde8dd36e..69b901ceb2a 100644 --- a/tests/unit_tests/test_utils.py +++ b/tests/unit_tests/test_utils.py @@ -14,8 +14,13 @@ import megatron.core.utils as util import megatron.training.utils as training_util from megatron.core import config -from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.distributed import ( + DistributedDataParallel, + DistributedDataParallelConfig, +) +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, +) from megatron.core.optimizer import OptimizerConfig, get_megatron_optimizer from megatron.core.transformer import TransformerConfig from megatron.core.transformer.moe.moe_layer import MoELayer @@ -313,12 +318,12 @@ def test_param_norm_moe(use_distributed_optimizer: bool): add_bias_linear=False, bf16=True, ) - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( - num_experts=2, moe_grouped_gemm=True - ) - model = MoELayer(transformer_config, transformer_layer_spec.submodules.mlp.submodules).to( - device='cuda' - ) + model = MoELayer( + transformer_config, + get_gpt_layer_with_transformer_engine_submodules( + num_experts=2, moe_grouped_gemm=True + ).mlp.submodules, + ).to(device='cuda') model.requires_grad_(True) # Initialize the model with all 1.0 for weights. for param in model.parameters(): diff --git a/tests/unit_tests/transformer/moe/test_grouped_mlp.py b/tests/unit_tests/transformer/moe/test_grouped_mlp.py index f215f9008b2..67543168480 100644 --- a/tests/unit_tests/transformer/moe/test_grouped_mlp.py +++ b/tests/unit_tests/transformer/moe/test_grouped_mlp.py @@ -5,8 +5,8 @@ import torch.nn.functional as F from megatron.core.models.gpt.gpt_layer_specs import ( - get_gpt_layer_local_spec, - get_gpt_layer_with_transformer_engine_spec, + get_gpt_layer_local_submodules, + get_gpt_layer_with_transformer_engine_submodules, ) from megatron.core.transformer.module import Float16Module from megatron.core.transformer.moe import grouped_gemm_util as gg @@ -69,8 +69,8 @@ def setup_method(self, method, use_cpu_initialization=False, swiglu=True): ## Vanilla sequential GEMM # Set random seed for reproducability _set_random_seed(seed_=123, data_parallel_random_init=False) - transformer_layer_spec = get_gpt_layer_local_spec(self.num_experts, moe_grouped_gemm=False) - self.sequential_mlp = MoELayer(tf_config, transformer_layer_spec.submodules.mlp.submodules) + submodules = get_gpt_layer_local_submodules(self.num_experts, moe_grouped_gemm=False) + self.sequential_mlp = MoELayer(tf_config, submodules.mlp.submodules) self.args = parse_args(ignore_unknown_args=True) self.args.bf16 = True @@ -83,10 +83,12 @@ def setup_method(self, method, use_cpu_initialization=False, swiglu=True): ## Grouped GEMM _set_random_seed(seed_=123, data_parallel_random_init=False) tf_config.moe_grouped_gemm = True - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( - self.num_experts, moe_grouped_gemm=True + self.grouped_mlp = MoELayer( + tf_config, + get_gpt_layer_with_transformer_engine_submodules( + self.num_experts, moe_grouped_gemm=True + ).mlp.submodules, ) - self.grouped_mlp = MoELayer(tf_config, transformer_layer_spec.submodules.mlp.submodules) self.grouped_mlp = Float16Module(self.grouped_mlp.config, self.grouped_mlp).module print("done intializing for grouped gemm") @@ -259,8 +261,8 @@ def setup_method(self, method, use_cpu_initialization=False, swiglu=True): ## Vanilla sequential GEMM # Set random seed for reproducability _set_random_seed(seed_=123, data_parallel_random_init=False) - transformer_layer_spec = get_gpt_layer_local_spec(self.num_experts, moe_grouped_gemm=False) - self.sequential_mlp = MoELayer(tf_config, transformer_layer_spec.submodules.mlp.submodules) + submodules = get_gpt_layer_local_submodules(self.num_experts, moe_grouped_gemm=False) + self.sequential_mlp = MoELayer(tf_config, submodules.mlp.submodules) self.args = parse_args(ignore_unknown_args=True) self.args.bf16 = True @@ -271,11 +273,13 @@ def setup_method(self, method, use_cpu_initialization=False, swiglu=True): ## Grouped GEMM _set_random_seed(seed_=123, data_parallel_random_init=False) - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( - self.num_experts, moe_grouped_gemm=True - ) tf_config.moe_grouped_gemm = True - self.grouped_mlp = MoELayer(tf_config, transformer_layer_spec.submodules.mlp.submodules) + self.grouped_mlp = MoELayer( + tf_config, + get_gpt_layer_with_transformer_engine_submodules( + self.num_experts, moe_grouped_gemm=True + ).mlp.submodules, + ) assert isinstance(self.grouped_mlp.experts, TEGroupedMLP) self.grouped_mlp = Float16Module(self.grouped_mlp.config, self.grouped_mlp).module diff --git a/tests/unit_tests/transformer/moe/test_moe_layer.py b/tests/unit_tests/transformer/moe/test_moe_layer.py index 11bd09f8449..f65c1e01caf 100644 --- a/tests/unit_tests/transformer/moe/test_moe_layer.py +++ b/tests/unit_tests/transformer/moe/test_moe_layer.py @@ -5,8 +5,8 @@ from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_decoder_block_spec, - get_gpt_layer_local_spec, - get_gpt_layer_with_transformer_engine_spec, + get_gpt_layer_local_submodules, + get_gpt_layer_with_transformer_engine_submodules, ) from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.moe.moe_layer import MoELayer @@ -45,12 +45,10 @@ def test_te_moe_layer(self, num_moe_experts, moe_token_dispatcher_type, grouped_ moe_ffn_hidden_size=128, add_bias_linear=False, ) - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=num_moe_experts, moe_grouped_gemm=grouped_gemm ) - moe_layer = MoELayer( - self.transformer_config, transformer_layer_spec.submodules.mlp.submodules - ) + moe_layer = MoELayer(self.transformer_config, submodules.mlp.submodules) Utils.destroy_model_parallel() @pytest.mark.parametrize("moe_token_dispatcher_type", ["allgather", "alltoall"]) @@ -73,12 +71,10 @@ def test_legacy_moe_layer(self, num_moe_experts, moe_token_dispatcher_type, grou moe_grouped_gemm=grouped_gemm, add_bias_linear=False, ) - transformer_layer_spec = get_gpt_layer_local_spec( + transformer_layer_submodules = get_gpt_layer_local_submodules( num_experts=num_moe_experts, moe_grouped_gemm=grouped_gemm ) - moe_layer = MoELayer( - self.transformer_config, transformer_layer_spec.submodules.mlp.submodules - ) + moe_layer = MoELayer(self.transformer_config, transformer_layer_submodules.mlp.submodules) Utils.destroy_model_parallel() @pytest.mark.skip( @@ -110,7 +106,7 @@ def test_moe_with_late_initialize( bf16=True, params_dtype=torch.bfloat16, ) - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=num_moe_experts, moe_grouped_gemm=grouped_gemm ) @@ -118,9 +114,7 @@ def test_moe_with_late_initialize( Utils.fake_initialize_model_parallel( tensor_model_parallel_size=tp_size, expert_model_parallel_size=ep_size ) - moe_layer = MoELayer( - transformer_config, transformer_layer_spec.submodules.mlp.submodules - ).cuda() + moe_layer = MoELayer(transformer_config, submodules.mlp.submodules).cuda() Utils.initialize_model_parallel( tensor_model_parallel_size=tp_size, expert_model_parallel_size=ep_size @@ -236,13 +230,11 @@ def test_moe_layer_fp16_forward_backward( params_dtype=torch.float16, ) - transformer_layer_spec = get_gpt_layer_local_spec( + submodules = get_gpt_layer_local_submodules( num_experts=num_moe_experts, moe_grouped_gemm=False ) - moe_layer = MoELayer( - transformer_config, transformer_layer_spec.submodules.mlp.submodules - ).cuda() + moe_layer = MoELayer(transformer_config, submodules.mlp.submodules).cuda() hidden_states = torch.randn( sequence_length, diff --git a/tests/unit_tests/transformer/moe/test_moe_layer_discrepancy.py b/tests/unit_tests/transformer/moe/test_moe_layer_discrepancy.py index 4386a844dd2..d60fbe336b8 100644 --- a/tests/unit_tests/transformer/moe/test_moe_layer_discrepancy.py +++ b/tests/unit_tests/transformer/moe/test_moe_layer_discrepancy.py @@ -7,10 +7,8 @@ from megatron.core import parallel_state from megatron.core.models.gpt.gpt_layer_specs import ( - get_gpt_layer_local_spec, - get_gpt_layer_with_transformer_engine_spec, + get_gpt_layer_with_transformer_engine_submodules, ) -from megatron.core.transformer.moe.moe_layer import MoELayer from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.transformer.transformer_layer import TransformerLayer from megatron.training.initialize import _set_random_seed @@ -49,7 +47,7 @@ def test_moe_layer_dispatcher_discrepancy( expert_model_parallel_size=ep_size, sequence_parallel=True if (tp_size > 1) else False, ) - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=num_moe_experts, moe_grouped_gemm=grouped_gemm ) # Init input and layer @@ -58,21 +56,13 @@ def test_moe_layer_dispatcher_discrepancy( # Init allgather moe layer _set_random_seed(seed_=123, data_parallel_random_init=False) - layer = ( - TransformerLayer(self.transformer_config, transformer_layer_spec.submodules) - .cuda() - .float() - ) + layer = TransformerLayer(self.transformer_config, submodules).cuda().float() ag_moe_layer = layer.mlp ag_moe_layer.eval() # Init a2a moe layer self.transformer_config.moe_token_dispatcher_type = "alltoall" _set_random_seed(seed_=123, data_parallel_random_init=False) - layer = ( - TransformerLayer(self.transformer_config, transformer_layer_spec.submodules) - .cuda() - .float() - ) + layer = TransformerLayer(self.transformer_config, submodules).cuda().float() a2a_moe_layer = layer.mlp a2a_moe_layer.eval() @@ -134,7 +124,7 @@ def test_moe_layer_ag_dispatcher_discrepancy( sequence_parallel=True if (tp_size > 1 and ep_size > 1) else False, bf16=True, ) - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=num_moe_experts, moe_grouped_gemm=grouped_gemm ) # Init input and layer @@ -143,11 +133,7 @@ def test_moe_layer_ag_dispatcher_discrepancy( # Init allgather moe layer _set_random_seed(seed_=123, data_parallel_random_init=False) - layer = ( - TransformerLayer(self.transformer_config, transformer_layer_spec.submodules) - .cuda() - .bfloat16() - ) + layer = TransformerLayer(self.transformer_config, submodules).cuda().bfloat16() ag_moe_layer = layer.mlp ag_moe_layer.eval() @@ -205,7 +191,7 @@ def test_moe_layer_a2a_dispatcher_discrepancy( sequence_parallel=True if (tp_size > 1 and ep_size > 1) else False, bf16=True, ) - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=num_moe_experts, moe_grouped_gemm=grouped_gemm ) # Init input and layer @@ -213,11 +199,7 @@ def test_moe_layer_a2a_dispatcher_discrepancy( input = torch.randn(1, 4096, 4096).cuda().bfloat16() # Init a2a moe layer - layer = ( - TransformerLayer(self.transformer_config, transformer_layer_spec.submodules) - .cuda() - .bfloat16() - ) + layer = TransformerLayer(self.transformer_config, submodules).cuda().bfloat16() a2a_moe_layer = layer.mlp a2a_moe_layer.eval() diff --git a/tests/unit_tests/transformer/moe/test_routers.py b/tests/unit_tests/transformer/moe/test_routers.py index 4d6b5ee2c3e..f2ae2e3fcc0 100644 --- a/tests/unit_tests/transformer/moe/test_routers.py +++ b/tests/unit_tests/transformer/moe/test_routers.py @@ -1,11 +1,10 @@ # Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. -from typing import cast import pytest import torch -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules from megatron.core.transformer.moe.moe_layer import MoELayer from megatron.core.transformer.moe.moe_utils import get_updated_expert_bias, router_gating_linear from megatron.core.transformer.moe.router import Router @@ -43,12 +42,10 @@ def setup_method(self, method): params_dtype=torch.bfloat16, add_bias_linear=False, ) - transformer_layer_spec = get_gpt_layer_local_spec( + submodules = get_gpt_layer_local_submodules( num_experts=num_moe_experts, moe_grouped_gemm=False ) - self.sequential_mlp = MoELayer( - self.transformer_config, transformer_layer_spec.submodules.mlp.submodules - ) + self.sequential_mlp = MoELayer(self.transformer_config, submodules.mlp.submodules) self.router = cast(Router, self.sequential_mlp.router) def teardown_method(self, method): @@ -314,12 +311,10 @@ def setup_method(self, method): ) # init MoE layer - transformer_layer_spec = get_gpt_layer_local_spec( + submodules = get_gpt_layer_local_submodules( num_experts=num_moe_experts, moe_grouped_gemm=False ) - self.moe_layer = MoELayer( - self.transformer_config, transformer_layer_spec.submodules.mlp.submodules - ).cuda() + self.moe_layer = MoELayer(self.transformer_config, submodules.mlp.submodules).cuda() self.router = cast(Router, self.moe_layer.router) def teardown_method(self, method): @@ -421,12 +416,10 @@ def setup_method(self, method): params_dtype=torch.bfloat16, add_bias_linear=False, ) - transformer_layer_spec = get_gpt_layer_local_spec( + submodules = get_gpt_layer_local_submodules( num_experts=num_moe_experts, moe_grouped_gemm=False ) - self.moe_layer = MoELayer( - self.transformer_config, transformer_layer_spec.submodules.mlp.submodules - ) + self.moe_layer = MoELayer(self.transformer_config, submodules.mlp.submodules) self.router = cast(Router, self.moe_layer.router) assert self.router.expert_bias is not None assert self.router.local_tokens_per_expert is not None @@ -469,15 +462,11 @@ def test_router_forward_aux_free(self): def test_router_forward_fusion_equivalence(self, score_function): with torch.no_grad(): # Build two fresh routers to avoid bias update interference - transformer_layer_spec = get_gpt_layer_local_spec( + submodules = get_gpt_layer_local_submodules( num_experts=self.transformer_config.num_moe_experts, moe_grouped_gemm=False ) - moe_layer_ref = MoELayer( - self.transformer_config, transformer_layer_spec.submodules.mlp.submodules - ) - moe_layer_fused = MoELayer( - self.transformer_config, transformer_layer_spec.submodules.mlp.submodules - ) + moe_layer_ref = MoELayer(self.transformer_config, submodules.mlp.submodules) + moe_layer_fused = MoELayer(self.transformer_config, submodules.mlp.submodules) router_ref = moe_layer_ref.router.cuda() router_fused = moe_layer_fused.router.cuda() diff --git a/tests/unit_tests/transformer/moe/test_sequential_mlp.py b/tests/unit_tests/transformer/moe/test_sequential_mlp.py index a80e2a2ff1b..e618f6a8318 100644 --- a/tests/unit_tests/transformer/moe/test_sequential_mlp.py +++ b/tests/unit_tests/transformer/moe/test_sequential_mlp.py @@ -5,7 +5,7 @@ import torch from megatron.core.extensions.transformer_engine import TEColumnParallelLinear, TERowParallelLinear -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.mlp import MLPSubmodules @@ -37,12 +37,10 @@ def setup_method(self, method): moe_router_topk=1, add_bias_linear=False, ) - transformer_layer_spec = get_gpt_layer_local_spec( + submodules = get_gpt_layer_local_submodules( num_experts=num_moe_experts, moe_grouped_gemm=False ) - self.sequential_mlp = MoELayer( - transformer_config, transformer_layer_spec.submodules.mlp.submodules - ) + self.sequential_mlp = MoELayer(transformer_config, submodules.mlp.submodules) def teardown_method(self, method): Utils.destroy_model_parallel() diff --git a/tests/unit_tests/transformer/moe/test_shared_experts.py b/tests/unit_tests/transformer/moe/test_shared_experts.py index 6df4d2fd369..d99dc1a0b05 100644 --- a/tests/unit_tests/transformer/moe/test_shared_experts.py +++ b/tests/unit_tests/transformer/moe/test_shared_experts.py @@ -3,7 +3,7 @@ import pytest import torch -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.moe.moe_layer import MoELayer from megatron.core.transformer.transformer_config import TransformerConfig @@ -41,12 +41,10 @@ def test_gpu_forward(self, shared_expert_gate): add_bias_linear=False, moe_shared_expert_gate=shared_expert_gate, ) - transformer_layer_spec = get_gpt_layer_local_spec( + submodules = get_gpt_layer_local_submodules( num_experts=num_moe_experts, moe_grouped_gemm=False ) - self.moe_layer = MoELayer( - transformer_config, transformer_layer_spec.submodules.mlp.submodules - ) + self.moe_layer = MoELayer(transformer_config, submodules.mlp.submodules) assert isinstance(self.moe_layer, MoELayer) @@ -103,12 +101,10 @@ def test_gpu_forward(self): moe_router_topk=1, add_bias_linear=False, ) - transformer_layer_spec = get_gpt_layer_local_spec( + submodules = get_gpt_layer_local_submodules( num_experts=num_moe_experts, moe_grouped_gemm=False ) - self.moe_layer = MoELayer( - transformer_config, transformer_layer_spec.submodules.mlp.submodules - ) + self.moe_layer = MoELayer(transformer_config, submodules.mlp.submodules) assert isinstance(self.moe_layer, MoELayer) diff --git a/tests/unit_tests/transformer/moe/test_token_dispatcher.py b/tests/unit_tests/transformer/moe/test_token_dispatcher.py index 24617952b94..a340857a0ee 100644 --- a/tests/unit_tests/transformer/moe/test_token_dispatcher.py +++ b/tests/unit_tests/transformer/moe/test_token_dispatcher.py @@ -7,7 +7,7 @@ import torch from megatron.core import config, parallel_state -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules from megatron.core.transformer.moe.moe_layer import MoELayer from megatron.core.transformer.moe.moe_utils import get_capacity from megatron.core.transformer.transformer_config import TransformerConfig @@ -99,15 +99,11 @@ def __init__( self.moe_layer = self.new_moe_layer() def new_moe_layer(self, **kargs): - transformer_layer_spec = get_gpt_layer_local_spec( + submodules = get_gpt_layer_local_submodules( num_experts=self.config.num_moe_experts, moe_grouped_gemm=self.config.moe_grouped_gemm ) new_config = dataclasses.replace(self.config, **kargs) - moe_layer = ( - MoELayer(new_config, transformer_layer_spec.submodules.mlp.submodules) - .cuda() - .to(dtype=self.test_dtype) - ) + moe_layer = MoELayer(new_config, submodules.mlp.submodules).cuda().to(dtype=self.test_dtype) moe_layer.set_layer_number(0) return moe_layer diff --git a/tests/unit_tests/transformer/test_attention.py b/tests/unit_tests/transformer/test_attention.py index d7771d0920d..704cc78d004 100644 --- a/tests/unit_tests/transformer/test_attention.py +++ b/tests/unit_tests/transformer/test_attention.py @@ -8,7 +8,9 @@ import megatron.core.parallel_state as parallel_state from megatron.core.hyper_comm_grid import HyperCommGrid -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, +) from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import TransformerConfig @@ -43,7 +45,7 @@ def setup_method(self, output_gate): ) self.parallel_attention = SelfAttention( self.transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, ) @@ -136,7 +138,7 @@ def test_checkpointed_gpu_forward(self): transformer_config.recompute_granularity = 'selective' checkpointed_parallel_attention = SelfAttention( transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, ) config = checkpointed_parallel_attention.config @@ -186,7 +188,7 @@ def test_clip_qk_disabled_raises_error(self): ) attention = SelfAttention( transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, ) @@ -206,7 +208,7 @@ def test_clip_qk_none_logits_raises_error(self): ) attention = SelfAttention( transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, ) @@ -226,7 +228,7 @@ def test_clip_qk_below_threshold_no_update(self): ) attention = SelfAttention( transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, ) attention.cuda() @@ -260,7 +262,7 @@ def test_clip_qk_above_threshold_updates_weights(self): ) attention = SelfAttention( transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, ) attention.cuda() @@ -295,7 +297,7 @@ def test_clip_qk_gqa_configuration(self): ) attention = SelfAttention( transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, ) attention.cuda() @@ -329,7 +331,7 @@ def test_clip_qk_mixed_logits(self): ) attention = SelfAttention( transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, ) attention.cuda() @@ -374,7 +376,7 @@ def run_self_attention(self, pg_collection): ) self.self_attention = SelfAttention( self.transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, attn_mask_type=AttnMaskType.causal, pg_collection=pg_collection, diff --git a/tests/unit_tests/transformer/test_attention_no_rope.py b/tests/unit_tests/transformer/test_attention_no_rope.py index 30e11609e57..2c79d7a3011 100644 --- a/tests/unit_tests/transformer/test_attention_no_rope.py +++ b/tests/unit_tests/transformer/test_attention_no_rope.py @@ -3,7 +3,9 @@ import pytest import torch -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, +) from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.attention import SelfAttention from megatron.core.transformer.enums import AttnMaskType @@ -30,7 +32,7 @@ def setup_method(self, method): ) self.parallel_attention = SelfAttention( self.transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, attn_mask_type=AttnMaskType.causal, ) @@ -176,7 +178,7 @@ def test_checkpointed_gpu_forward(self): checkpointed_parallel_attention = SelfAttention( transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, attn_mask_type=AttnMaskType.causal, ) diff --git a/tests/unit_tests/transformer/test_attention_packed_seq.py b/tests/unit_tests/transformer/test_attention_packed_seq.py index e6e2c847395..9c8f5bd2ce6 100644 --- a/tests/unit_tests/transformer/test_attention_packed_seq.py +++ b/tests/unit_tests/transformer/test_attention_packed_seq.py @@ -3,7 +3,9 @@ import pytest import torch -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, +) from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.attention import SelfAttention @@ -62,7 +64,7 @@ def setup_method(self, method): ) self.parallel_attention = SelfAttention( self.transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, attn_mask_type=AttnMaskType.causal, ) @@ -138,7 +140,7 @@ def test_checkpointed_gpu_forward(self): transformer_config.recompute_granularity = 'selective' checkpointed_parallel_attention = SelfAttention( transformer_config, - get_gpt_layer_with_transformer_engine_spec().submodules.self_attention.submodules, + get_gpt_layer_with_transformer_engine_submodules().self_attention.submodules, layer_number=1, attn_mask_type=AttnMaskType.causal, ) diff --git a/tests/unit_tests/transformer/test_cuda_graphs.py b/tests/unit_tests/transformer/test_cuda_graphs.py index 4696a3ed439..5a8da60f2ee 100644 --- a/tests/unit_tests/transformer/test_cuda_graphs.py +++ b/tests/unit_tests/transformer/test_cuda_graphs.py @@ -37,6 +37,7 @@ from megatron.core.transformer.moe.fused_a2a import reset_hybrid_ep_buffer from megatron.core.transformer.transformer_block import TransformerBlock from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.transformer.transformer_layer import TransformerLayerSubmodules from megatron.core.utils import is_fa_min_version, is_te_min_version from megatron.training.arguments import core_transformer_config_from_args, parse_args, validate_args from megatron.training.global_vars import ( @@ -327,6 +328,7 @@ def setup_method(self, method): # Get layer specs language_layer_spec = get_gpt_layer_with_transformer_engine_spec() vision_layer_spec = get_vit_layer_with_transformer_engine_spec() + assert isinstance(language_layer_spec.submodules, TransformerLayerSubmodules) vision_projection_spec = deepcopy(language_layer_spec.submodules.mlp.submodules) # Set vision model type diff --git a/tests/unit_tests/transformer/test_mlp.py b/tests/unit_tests/transformer/test_mlp.py index d2c25e0cc53..45e0df9ad0a 100644 --- a/tests/unit_tests/transformer/test_mlp.py +++ b/tests/unit_tests/transformer/test_mlp.py @@ -1,9 +1,10 @@ # Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. + import pytest import torch -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.mlp import MLP from megatron.core.transformer.transformer_config import TransformerConfig @@ -18,7 +19,7 @@ def setup_method(self, method): transformer_config = TransformerConfig( num_layers=2, hidden_size=12, num_attention_heads=4, use_cpu_initialization=True ) - self.mlp = MLP(transformer_config, get_gpt_layer_local_spec().submodules.mlp.submodules) + self.mlp = MLP(transformer_config, get_gpt_layer_local_submodules().mlp.submodules) def teardown_method(self, method): Utils.destroy_model_parallel() diff --git a/tests/unit_tests/transformer/test_multi_latent_attention.py b/tests/unit_tests/transformer/test_multi_latent_attention.py index bc8514ee561..2cea21ce3e5 100644 --- a/tests/unit_tests/transformer/test_multi_latent_attention.py +++ b/tests/unit_tests/transformer/test_multi_latent_attention.py @@ -15,13 +15,20 @@ from megatron.core.models.common.embeddings.rope_utils import ( get_pos_emb_on_this_cp_rank as get_tensor_on_this_cp_rank, ) -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_spec, + get_gpt_layer_with_transformer_engine_submodules, +) from megatron.core.models.gpt.gpt_model import GPTModel from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.attention import Attention from megatron.core.transformer.enums import AttnMaskType -from megatron.core.transformer.multi_latent_attention import MLASelfAttention, MultiLatentAttention +from megatron.core.transformer.multi_latent_attention import ( + MLASelfAttention, + MLASelfAttentionSubmodules, + MultiLatentAttention, +) from megatron.core.transformer.transformer_config import MLATransformerConfig from megatron.core.utils import is_te_min_version, is_torch_min_version from megatron.training.arguments import parse_args @@ -92,9 +99,10 @@ def make_test_packed_seq_params_with_padding( def get_mla_self_attn_submodules(linear_qkv_down_proj=None): - submodules = get_gpt_layer_with_transformer_engine_spec( + submodules = get_gpt_layer_with_transformer_engine_submodules( multi_latent_attention=True - ).submodules.self_attention.submodules + ).self_attention.submodules + assert isinstance(submodules, MLASelfAttentionSubmodules) if linear_qkv_down_proj is not None: submodules.linear_q_down_proj = linear_qkv_down_proj submodules.linear_kv_down_proj = linear_qkv_down_proj diff --git a/tests/unit_tests/transformer/test_spec_customization.py b/tests/unit_tests/transformer/test_spec_customization.py index d6b2a2eb623..c2ab094d0cf 100755 --- a/tests/unit_tests/transformer/test_spec_customization.py +++ b/tests/unit_tests/transformer/test_spec_customization.py @@ -14,7 +14,7 @@ TERowParallelLinear, ) from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules from megatron.core.parallel_state import get_context_parallel_group, get_tensor_model_parallel_group from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed @@ -203,20 +203,18 @@ def test_transformer_block_custom(self): transformer_config = TransformerConfig( num_layers=2, hidden_size=12, num_attention_heads=4, use_cpu_initialization=True ) - layer_local_spec = get_gpt_layer_local_spec() + submodules = get_gpt_layer_local_submodules() # The following way can be used to pass a different `TransformerLayer` # and internally the `TransformerBlock` would fan out the single # `ModuleSpec` layer spec provided to all the layers of the block. - layer_spec1 = ModuleSpec(module=TransformerLayer, submodules=layer_local_spec.submodules) + layer_spec1 = ModuleSpec(module=TransformerLayer, submodules=submodules) model_parallel_cuda_manual_seed(123) torch.manual_seed(0) parallel_transformer_block1 = TransformerBlock(transformer_config, layer_spec1) layer_spec2 = TransformerBlockSubmodules( - layer_specs=[ - ModuleSpec(module=TransformerLayer, submodules=layer_local_spec.submodules) - ] + layer_specs=[ModuleSpec(module=TransformerLayer, submodules=submodules)] * transformer_config.num_layers, layer_norm=TENorm, ) @@ -252,12 +250,10 @@ def test_transformer_block_custom(self): def test_l2_qk_norm(self): """Test L2 normalization for QK vectors using local spec.""" - layer_spec = get_gpt_layer_local_spec(qk_l2_norm=True) + submodules = get_gpt_layer_local_submodules(qk_l2_norm=True) # Build the self-attention module from the spec - self_attention = build_module( - layer_spec.submodules.self_attention, config=self.config, layer_number=1 - ) + self_attention = build_module(submodules.self_attention, config=self.config, layer_number=1) assert isinstance(self_attention, SelfAttention) # Verify that q_layernorm and k_layernorm are L2Norm instances diff --git a/tests/unit_tests/transformer/test_submodule_callables.py b/tests/unit_tests/transformer/test_submodule_callables.py index 73059495c06..03e2d751a52 100644 --- a/tests/unit_tests/transformer/test_submodule_callables.py +++ b/tests/unit_tests/transformer/test_submodule_callables.py @@ -3,7 +3,9 @@ import torch from megatron.core.models.gpt.fine_grained_callables import build_layer_callables -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, +) from megatron.core.transformer.transformer_layer import TransformerLayer from megatron.core.utils import is_te_min_version from tests.unit_tests.a2a_overlap.utils import ( @@ -140,13 +142,13 @@ def test_1f1b_overlap(self, dispatcher_type, grouped_gemm, permute_fusion): config = get_test_config(extra_kwargs=extra_kwargs, moe_grouped_gemm=grouped_gemm) microbatches = 4 with deterministic_mode(): - transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + transformer_layer_submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=8, moe_grouped_gemm=grouped_gemm, qk_layernorm=True, multi_latent_attention=True, ) - model = TransformerLayer(config, transformer_layer_spec.submodules) + model = TransformerLayer(config, transformer_layer_submodules) params = reset_model(model) input_tensors = [build_data() for _ in range(microbatches)] diff --git a/tests/unit_tests/transformer/test_transformer_block_custom_pgs.py b/tests/unit_tests/transformer/test_transformer_block_custom_pgs.py index bb64efe7449..c55babe35ca 100644 --- a/tests/unit_tests/transformer/test_transformer_block_custom_pgs.py +++ b/tests/unit_tests/transformer/test_transformer_block_custom_pgs.py @@ -68,7 +68,7 @@ def __init__( # Temporarily replace attention and MLP with IdentityOp, # This is a temporary workaround for the test until we have a better interface # will rebuild them with custom process groups after super init - def _modify_submodules(submodules): + def _modify_submodules(submodules: TransformerLayerSubmodules): submodules.self_attention = IdentityOp submodules.mlp = IdentityOp return submodules diff --git a/tests/unit_tests/transformer/test_transformer_layer.py b/tests/unit_tests/transformer/test_transformer_layer.py index 7db9aa30fec..da1f9ce5860 100644 --- a/tests/unit_tests/transformer/test_transformer_layer.py +++ b/tests/unit_tests/transformer/test_transformer_layer.py @@ -7,7 +7,9 @@ from megatron.core import parallel_state from megatron.core.dist_checkpointing.mapping import ShardedObject, ShardedTensor from megatron.core.inference.contexts import StaticInferenceContext -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_submodules, +) from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.transformer.transformer_layer import ( @@ -26,7 +28,7 @@ def setup_method(self, method): num_layers=2, hidden_size=12, num_attention_heads=4, use_cpu_initialization=True ) self.parallel_transformer_layer = TransformerLayer( - transformer_config, get_gpt_layer_with_transformer_engine_spec().submodules + transformer_config, get_gpt_layer_with_transformer_engine_submodules() ) def teardown_method(self, method): @@ -81,7 +83,7 @@ def test( use_cpu_initialization=True, ) parallel_transformer_layer = TransformerLayer( - transformer_config, get_gpt_layer_with_transformer_engine_spec().submodules + transformer_config, get_gpt_layer_with_transformer_engine_submodules() ) parallel_transformer_layer.cuda() @@ -265,7 +267,7 @@ def test_sharded_state_dict(self, tp_pp, order): num_layers=2, hidden_size=128, num_attention_heads=8, use_cpu_initialization=True ) parallel_transformer_layer = TransformerLayer( - transformer_config, get_gpt_layer_with_transformer_engine_spec().submodules + transformer_config, get_gpt_layer_with_transformer_engine_submodules() ) sharded_state_dict = parallel_transformer_layer.sharded_state_dict()