diff --git a/examples/multimodal/layer_specs.py b/examples/multimodal/layer_specs.py index 4c50ecea10a..13fc63c0bb9 100644 --- a/examples/multimodal/layer_specs.py +++ b/examples/multimodal/layer_specs.py @@ -2,6 +2,10 @@ import torch from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add +from megatron.core.ssm.mamba_block import MambaStack, MambaStackSubmodules +from megatron.core.ssm.mamba_layer import MambaLayer, MambaLayerSubmodules +from megatron.core.ssm.mamba_mixer import MambaMixer, MambaMixerSubmodules +from megatron.core.ssm.mlp_layer import MLPLayer from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear from megatron.core.transformer.attention import SelfAttention, SelfAttentionSubmodules from megatron.core.transformer.dot_product_attention import DotProductAttention @@ -10,10 +14,6 @@ from megatron.core.transformer.mlp import MLP, MLPSubmodules from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.transformer.transformer_layer import TransformerLayer, TransformerLayerSubmodules -from megatron.core.ssm.mamba_block import MambaStack, MambaStackSubmodules -from megatron.core.ssm.mamba_layer import MambaLayer, MambaLayerSubmodules -from megatron.core.ssm.mamba_mixer import MambaMixer, MambaMixerSubmodules -from megatron.core.ssm.mlp_layer import MLPLayer try: from megatron.core.extensions.transformer_engine import ( @@ -26,6 +26,9 @@ HAVE_TE = True except ImportError: + TELayerNormColumnParallelLinear = None + TEDotProductAttention = None + TERowParallelLinear = None HAVE_TE = False try: @@ -54,12 +57,8 @@ def get_layer_spec(is_vit, normalization) -> ModuleSpec: norm = TENorm else: version = torch.__version__.split('.') - version_geq_2_4 = ( - int(TORCH_VERSION[0]) > 2 - or ( - int(TORCH_VERSION[0]) == 2 - and int(TORCH_VERSION[1]) >= 4 - ) + version_geq_2_4 = int(TORCH_VERSION[0]) > 2 or ( + int(TORCH_VERSION[0]) == 2 and int(TORCH_VERSION[1]) >= 4 ) assert version_geq_2_4, "Torch version >= 2.4.0 is required for RMSNorm" if HAVE_APEX: @@ -101,6 +100,9 @@ def get_layer_spec_te(is_vit=False, padding=False) -> ModuleSpec: attn_mask_type = AttnMaskType.padding_causal mlp = get_norm_mlp_module_spec_te() + assert TELayerNormColumnParallelLinear is not None + assert TEDotProductAttention is not None + assert TERowParallelLinear is not None return ModuleSpec( module=TransformerLayer, submodules=TransformerLayerSubmodules( @@ -122,12 +124,16 @@ def get_layer_spec_te(is_vit=False, padding=False) -> ModuleSpec: ), ) + def get_mamba_layer_spec_te(padding=False) -> ModuleSpec: attn_mask_type = AttnMaskType.causal # Padding mask is needed for e.g. Context Parallel. if padding: attn_mask_type = AttnMaskType.padding_causal + assert TELayerNormColumnParallelLinear is not None + assert TEDotProductAttention is not None + assert TERowParallelLinear is not None return ModuleSpec( module=MambaStack, submodules=MambaStackSubmodules( @@ -170,7 +176,8 @@ def get_mamba_layer_spec_te(padding=False) -> ModuleSpec: mlp=ModuleSpec( module=MLP, submodules=MLPSubmodules( - linear_fc1=TELayerNormColumnParallelLinear, linear_fc2=TERowParallelLinear + linear_fc1=TELayerNormColumnParallelLinear, + linear_fc2=TERowParallelLinear, ), ), mlp_bda=get_bias_dropout_add, @@ -179,6 +186,7 @@ def get_mamba_layer_spec_te(padding=False) -> ModuleSpec: ), ) + def get_mlp_module_spec(use_te: bool = True) -> ModuleSpec: # Dense MLP w/ or w/o TE modules. return ModuleSpec( diff --git a/examples/multimodal/nvlm/internvit.py b/examples/multimodal/nvlm/internvit.py index 62f3bdccd85..2cede9d2cd9 100644 --- a/examples/multimodal/nvlm/internvit.py +++ b/examples/multimodal/nvlm/internvit.py @@ -10,11 +10,16 @@ Those code changes are gathered here. """ +from collections.abc import Callable from functools import partial +from typing import cast import torch -from megatron.core.utils import divide +from examples.multimodal.layer_scaling import ( + LayerScalingTransformerLayer, + get_bias_dropout_add_layer_scaling, +) from megatron.core.extensions.transformer_engine import ( TEColumnParallelLinear, TEDotProductAttention, @@ -35,9 +40,7 @@ from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.transformer.transformer_layer import TransformerLayer, TransformerLayerSubmodules from megatron.core.transformer.utils import make_sharded_tensors_for_checkpoint - -from examples.multimodal.layer_scaling import LayerScalingTransformerLayer, get_bias_dropout_add_layer_scaling - +from megatron.core.utils import divide try: import apex @@ -60,7 +63,7 @@ class InternViTRMSNorm(MegatronModule): def __init__( self, - config, + config: TransformerConfig, hidden_size: int, eps: float = 1e-6, sequence_parallel: bool = False, @@ -92,7 +95,7 @@ def _norm(self, x, var): return x * torch.rsqrt(var + self.eps) - def forward(self, x): + def forward(self, x: torch.Tensor) -> torch.Tensor: """Run RMSNorm with an option to compute custom statistic.""" var = None if self._compute_var: @@ -128,10 +131,14 @@ def _gather_var(self, input_, max_dim): if rank < valid_ranks: # Ranks without any dummy attention heads. var = input_.sum(-1, keepdim=True) - elif rank == valid_ranks: # The only rank which may contain 'residual_heads' dummy attention heads. + elif ( + rank == valid_ranks + ): # The only rank which may contain 'residual_heads' dummy attention heads. var = input_[..., :max_dim].sum(-1, keepdim=True) else: - var = input_.sum(-1, keepdim=True) * 0.0 # All heads in these ranks are dummy heads: Zero-out. + var = ( + input_.sum(-1, keepdim=True) * 0.0 + ) # All heads in these ranks are dummy heads: Zero-out. tensor_list = [torch.empty_like(var) for _ in range(world_size)] tensor_list[rank] = var @@ -175,37 +182,35 @@ def __init__( # Need to override linear_qkv, q_layernorm and k_layernorm. qkv_bias = False - self.linear_qkv = build_module( - submodules.linear_qkv, + self.linear_qkv = submodules.linear_qkv( self.config.hidden_size, self.query_projection_size + 2 * self.kv_projection_size, config=self.config, - init_method=self.config.init_method, + init_method=cast(Callable[[torch.Tensor], None], self.config.init_method), gather_output=False, bias=qkv_bias, skip_bias_add=False, is_expert=False, tp_comm_buffer_name='qkv', + tp_group=None, ) qk_layernorm_hidden_size = ( self.hidden_size_per_attention_head * self.num_attention_heads_per_partition ) # 512 for internvit - self.q_layernorm = build_module( - submodules.q_layernorm, + assert submodules.q_layernorm is not None + self.q_layernorm = submodules.q_layernorm( hidden_size=qk_layernorm_hidden_size, config=self.config, eps=self.config.layernorm_epsilon, - compute_var=True, ) - self.k_layernorm = build_module( - submodules.k_layernorm, + assert submodules.k_layernorm is not None + self.k_layernorm = submodules.k_layernorm( hidden_size=qk_layernorm_hidden_size, config=self.config, eps=self.config.layernorm_epsilon, - compute_var=True, ) @@ -245,8 +250,8 @@ def get_internvit_layer_spec(use_te) -> ModuleSpec: linear_qkv=TEColumnParallelLinear if use_te else ColumnParallelLinear, core_attention=TEDotProductAttention if use_te else DotProductAttention, linear_proj=TERowParallelLinear if use_te else RowParallelLinear, - q_layernorm=InternViTRMSNorm, - k_layernorm=InternViTRMSNorm, + q_layernorm=partial(InternViTRMSNorm, compute_var=True), + k_layernorm=partial(InternViTRMSNorm, compute_var=True), ), ), self_attn_bda=get_bias_dropout_add_layer_scaling, @@ -256,6 +261,7 @@ def get_internvit_layer_spec(use_te) -> ModuleSpec: ), ) + def get_internvit300M_layer_spec(use_te) -> ModuleSpec: mlp = get_mlp_module_spec(use_te) # no norm diff --git a/examples/multimodal/radio/radio_g.py b/examples/multimodal/radio/radio_g.py index 3ce793be75d..856eecdf0fe 100644 --- a/examples/multimodal/radio/radio_g.py +++ b/examples/multimodal/radio/radio_g.py @@ -3,6 +3,10 @@ import torch +from examples.multimodal.layer_scaling import ( + LayerScalingTransformerLayer, + get_bias_dropout_add_layer_scaling, +) from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear from megatron.core.transformer.attention import SelfAttention, SelfAttentionSubmodules from megatron.core.transformer.dot_product_attention import DotProductAttention @@ -11,7 +15,6 @@ from megatron.core.transformer.mlp import MLP, MLPSubmodules from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.transformer.transformer_layer import TransformerLayer, TransformerLayerSubmodules -from examples.multimodal.layer_scaling import LayerScalingTransformerLayer, get_bias_dropout_add_layer_scaling try: from megatron.core.extensions.transformer_engine import ( @@ -24,6 +27,9 @@ HAVE_TE = True except ImportError: + TELayerNormColumnParallelLinear = None + TEDotProductAttention = None + TERowParallelLinear = None HAVE_TE = False try: @@ -106,6 +112,9 @@ def get_radio_g_layer_spec_te() -> ModuleSpec: attn_mask_type = AttnMaskType.no_mask mlp = get_norm_mlp_module_spec_te() + assert TELayerNormColumnParallelLinear is not None + assert TEDotProductAttention is not None + assert TERowParallelLinear is not None return ModuleSpec( module=LayerScalingTransformerLayer, submodules=TransformerLayerSubmodules( diff --git a/megatron/core/extensions/kitchen.py b/megatron/core/extensions/kitchen.py index 7f2f1fac9c8..410d404a941 100644 --- a/megatron/core/extensions/kitchen.py +++ b/megatron/core/extensions/kitchen.py @@ -23,7 +23,11 @@ ) from megatron.core.tensor_parallel.utils import divide from megatron.core.transformer.mlp import MLPSubmodules -from megatron.core.transformer.moe.experts import GroupedMLP, SequentialMLP, TEGroupedMLP +from megatron.core.transformer.moe.experts import ( + GroupedMLP, + SequentialMLP, + TEGroupedMLP, +) from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.transformer.utils import make_sharded_tensors_for_checkpoint from megatron.core.utils import get_tensor_model_parallel_group_if_none @@ -1037,7 +1041,7 @@ def column_parallel_linear(self) -> type: """Which column parallel linear module kitchen backend uses""" return KitchenColumnParallelLinear - def row_parallel_linear(self) -> type: + def row_parallel_linear(self) -> type[KitchenRowParallelLinear]: """Which row parallel linear module kitchen backend uses""" return KitchenRowParallelLinear @@ -1052,7 +1056,7 @@ def fuse_layernorm_and_linear(self) -> bool: # explicitly about whether to include a norm. return self.fallback.fuse_layernorm_and_linear() - def column_parallel_layer_norm_linear(self) -> Optional[type]: + def column_parallel_layer_norm_linear(self) -> type[KitchenLayerNormColumnParallelLinear]: """Which module for sequential layernorm and linear""" return KitchenLayerNormColumnParallelLinear diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index e95409e08e9..f8d87269234 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -6,7 +6,7 @@ import os import pickle import warnings -from typing import Any, Callable, List, Optional, Tuple +from typing import TYPE_CHECKING, Any, Callable, List, Optional, Tuple, assert_never import torch import torch.nn.functional as F @@ -55,14 +55,19 @@ ) try: - import transformer_engine as te - HAVE_TE = True + + import transformer_engine as te except ImportError: - from unittest.mock import MagicMock + if TYPE_CHECKING: + import transformer_engine as te - te = MagicMock() - HAVE_TE = False + # Force type checking to treat TE as available + else: + from unittest.mock import MagicMock + + te = MagicMock() + HAVE_TE = False def _get_extra_te_kwargs(config: TransformerConfig): @@ -419,7 +424,7 @@ def __init__( # duplicated across TP ranks setattr(param, "sequence_parallel", self.config.sequence_parallel) - def forward(self, x): + def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: """Forward.""" _is_first_microbatch = ( None if self.disable_parameter_transpose_cache else self.is_first_microbatch @@ -461,7 +466,7 @@ def __init__( output_size: int, *, config: TransformerConfig, - init_method: Callable, + init_method: Callable[[torch.Tensor], None], gather_output: bool, bias: bool, skip_bias_add: bool, @@ -607,7 +612,7 @@ def __init__( self.bias.zero_() setattr(self.bias, "allreduce", True) - def forward(self, x): + def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: """Forward.""" _is_first_microbatch = ( None if self.disable_parameter_transpose_cache else self.is_first_microbatch @@ -849,8 +854,8 @@ def __init__( softmax_scale: Optional[float] = None, k_channels: Optional[int] = None, v_channels: Optional[int] = None, - cp_comm_type: str = "p2p", - pg_collection: ProcessGroupCollection = None, + cp_comm_type: Optional[str] = "p2p", + pg_collection: Optional[ProcessGroupCollection] = None, ): if not HAVE_TE: raise ImportError( @@ -1013,9 +1018,9 @@ def forward( value: Tensor, attention_mask: Tensor, attn_mask_type: AttnMaskType, - attention_bias: Tensor = None, - packed_seq_params: PackedSeqParams = None, - ): + attention_bias: Optional[Tensor] = None, + packed_seq_params: Optional[PackedSeqParams] = None, + ) -> Tensor: """Forward.""" packed_seq_kwargs = ( {key: getattr(packed_seq_params, key) for key in self.kept_packed_seq_params} @@ -1989,7 +1994,10 @@ def fused_apply_rotary_pos_emb_thd( pass try: - from transformer_engine.pytorch import Fp8Padding, Fp8Unpadding # pylint: disable=unused-import + from transformer_engine.pytorch import ( # pylint: disable=unused-import + Fp8Padding, + Fp8Unpadding, + ) except ImportError: Fp8Padding = None diff --git a/megatron/core/extensions/transformer_engine_spec_provider.py b/megatron/core/extensions/transformer_engine_spec_provider.py index a071959bfc9..f366d632cf9 100644 --- a/megatron/core/extensions/transformer_engine_spec_provider.py +++ b/megatron/core/extensions/transformer_engine_spec_provider.py @@ -1,7 +1,7 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import warnings -from typing import Optional, Tuple +from typing import Optional, Tuple, Union from megatron.core.extensions.transformer_engine import ( TEActivationOp, @@ -33,7 +33,7 @@ def column_parallel_linear(self) -> type: """Which column parallel linear module TE backend uses""" return TEColumnParallelLinear - def row_parallel_linear(self) -> type: + def row_parallel_linear(self) -> type[TERowParallelLinear]: """Which row parallel linear module TE backend uses""" return TERowParallelLinear @@ -41,11 +41,13 @@ def fuse_layernorm_and_linear(self) -> bool: """TE backend chooses a single module for layernorm and linear""" return True - def column_parallel_layer_norm_linear(self) -> Optional[type]: + def column_parallel_layer_norm_linear(self) -> type[TELayerNormColumnParallelLinear]: """Which module for sequential layernorm and linear""" return TELayerNormColumnParallelLinear - def layer_norm(self, rms_norm: bool = False, for_qk: bool = False) -> type: + def layer_norm( + self, rms_norm: bool = False, for_qk: bool = False + ) -> Union[type[FusedLayerNorm], type[TENorm]]: """Which module to use for layer norm""" if for_qk and not is_te_min_version("1.9.0"): # TENorm significantly harms convergence when used @@ -54,7 +56,7 @@ def layer_norm(self, rms_norm: bool = False, for_qk: bool = False) -> type: return FusedLayerNorm return TENorm - def core_attention(self) -> type: + def core_attention(self) -> type[TEDotProductAttention]: """Which module to use for attention""" return TEDotProductAttention diff --git a/megatron/core/models/T5/t5_spec.py b/megatron/core/models/T5/t5_spec.py index a3ad84a21d6..7766b07d1ba 100644 --- a/megatron/core/models/T5/t5_spec.py +++ b/megatron/core/models/T5/t5_spec.py @@ -28,6 +28,10 @@ HAVE_TE = True except ImportError: + TELayerNormColumnParallelLinear = None + TEDotProductAttention = None + TEColumnParallelLinear = None + TERowParallelLinear = None HAVE_TE = False try: @@ -49,7 +53,9 @@ def encoder_model_with_transformer_engine_default_spec() -> ModuleSpec: """T5 encoder TE spec (uses Transformer Engine components).""" - + assert TELayerNormColumnParallelLinear is not None + assert TEDotProductAttention is not None + assert TERowParallelLinear is not None return ModuleSpec( module=TransformerLayer, submodules=TransformerLayerSubmodules( @@ -78,7 +84,10 @@ def encoder_model_with_transformer_engine_default_spec() -> ModuleSpec: def decoder_model_with_transformer_engine_default_spec() -> ModuleSpec: """T5 decoder TE spec (uses Transformer Engine components).""" - + assert TELayerNormColumnParallelLinear is not None + assert TEDotProductAttention is not None + assert TEColumnParallelLinear is not None + assert TERowParallelLinear is not None return ModuleSpec( module=TransformerLayer, submodules=TransformerLayerSubmodules( diff --git a/megatron/core/models/backends.py b/megatron/core/models/backends.py index 29169285b3e..02a4ee51bdf 100644 --- a/megatron/core/models/backends.py +++ b/megatron/core/models/backends.py @@ -2,7 +2,7 @@ import warnings from abc import abstractmethod -from typing import Optional, Protocol, Tuple +from typing import Optional, Protocol, Tuple, Union from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear from megatron.core.transformer.dot_product_attention import DotProductAttention @@ -89,7 +89,7 @@ def column_parallel_linear(self) -> type: """Which column parallel linear module the backend uses""" return ColumnParallelLinear - def row_parallel_linear(self) -> type: + def row_parallel_linear(self) -> type[RowParallelLinear]: """Which row parallel linear module the backend uses""" return RowParallelLinear @@ -97,11 +97,13 @@ def fuse_layernorm_and_linear(self) -> bool: """Does the backend choose a single module for layernorm and linear""" return False - def column_parallel_layer_norm_linear(self) -> Optional[type]: + def column_parallel_layer_norm_linear(self) -> None: """Which module for sequential layernorm and linear""" return None - def layer_norm(self, rms_norm: bool = False, for_qk: bool = False) -> type: + def layer_norm( + self, rms_norm: bool = False, for_qk: bool = False + ) -> Union[type['FusedLayerNorm'], type[WrappedTorchNorm]]: """Which module to use for layer norm""" if rms_norm: # Matching get_gpt_layer_local_spec. @@ -110,7 +112,7 @@ def layer_norm(self, rms_norm: bool = False, for_qk: bool = False) -> type: LNImpl = WrappedTorchNorm return LNImpl - def core_attention(self) -> type: + def core_attention(self) -> type[DotProductAttention]: """Which module to use for attention""" return DotProductAttention @@ -145,7 +147,7 @@ def column_parallel_linear(self) -> type: """Which column parallel linear module TE backend uses""" return TEColumnParallelLinear - def row_parallel_linear(self) -> type: + def row_parallel_linear(self) -> type[InferenceRowParallelLinear]: """Which row parallel linear module TE backend uses""" return InferenceRowParallelLinear @@ -153,11 +155,13 @@ def fuse_layernorm_and_linear(self) -> bool: """TE backend chooses a single module for layernorm and linear""" return True - def column_parallel_layer_norm_linear(self) -> Optional[type]: + def column_parallel_layer_norm_linear(self) -> type[InferenceLayerNormColumnParallelLinear]: """Which module for sequential layernorm and linear""" return InferenceLayerNormColumnParallelLinear - def layer_norm(self, rms_norm: bool = False, for_qk: bool = False) -> type: + def layer_norm( + self, rms_norm: bool = False, for_qk: bool = False + ) -> Union[type['FusedLayerNorm'], type[TENorm]]: """Which module to use for layer norm""" if for_qk and not is_te_min_version("1.9.0"): # TENorm significantly harms convergence when used @@ -166,7 +170,7 @@ def layer_norm(self, rms_norm: bool = False, for_qk: bool = False) -> type: return FusedLayerNorm return TENorm - def core_attention(self) -> type: + def core_attention(self) -> type[TEDotProductAttention]: """Which module to use for attention""" return TEDotProductAttention diff --git a/megatron/core/models/bert/bert_layer_specs.py b/megatron/core/models/bert/bert_layer_specs.py index 69cec788b2c..ab51a5a2fd2 100644 --- a/megatron/core/models/bert/bert_layer_specs.py +++ b/megatron/core/models/bert/bert_layer_specs.py @@ -22,6 +22,9 @@ HAVE_TE = True except ImportError: + TELayerNormColumnParallelLinear = None + TEDotProductAttention = None + TERowParallelLinear = None HAVE_TE = False try: @@ -49,7 +52,9 @@ def get_bert_layer_with_transformer_engine_spec(): raise ImportError( "Transformer Engine is not installed. Please use local Bert layer spec instead." ) - + assert TELayerNormColumnParallelLinear is not None + assert TEDotProductAttention is not None + assert TERowParallelLinear is not None return ModuleSpec( module=TransformerLayer, submodules=TransformerLayerSubmodules( diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py index cffeb9234f1..a5ed3d6b207 100755 --- a/megatron/core/models/gpt/gpt_layer_specs.py +++ b/megatron/core/models/gpt/gpt_layer_specs.py @@ -36,6 +36,7 @@ TransformerLayerSubmodules, get_transformer_layer_offset, ) +from megatron.legacy.model.rms_norm import RMSNorm try: import transformer_engine as te # pylint: disable=unused-import @@ -135,6 +136,8 @@ def get_gpt_layer_with_inference_spec( ) else: qk_norm = backend.layer_norm(for_qk=True) + from transformer_engine import pytorch as te_pytorch + return ModuleSpec( module=TransformerLayer, submodules=TransformerLayerSubmodules( diff --git a/megatron/core/models/gpt/heterogeneous/heterogeneous_layer_specs.py b/megatron/core/models/gpt/heterogeneous/heterogeneous_layer_specs.py index b1c2fb79a11..70790723665 100644 --- a/megatron/core/models/gpt/heterogeneous/heterogeneous_layer_specs.py +++ b/megatron/core/models/gpt/heterogeneous/heterogeneous_layer_specs.py @@ -44,6 +44,9 @@ HAVE_TE = True except ImportError: + TELayerNormColumnParallelLinear = None + TEDotProductAttention = None + TERowParallelLinear = None HAVE_TE = False from megatron.core.transformer.torch_norm import WrappedTorchNorm @@ -106,13 +109,27 @@ def _get_heterogenous_attention_spec( ) else: ln = _get_qk_layernorm(use_te, normalization) if qk_layernorm else IdentityOp + if use_te: + assert TELayerNormColumnParallelLinear is not None + assert TEDotProductAttention is not None + assert TERowParallelLinear is not None + linear_qkv = TELayerNormColumnParallelLinear + core_attention = TEDotProductAttention + linear_proj = TERowParallelLinear + + else: + linear_qkv = ColumnParallelLinear + core_attention = DotProductAttention + linear_proj = RowParallelLinear + from transformer_engine import pytorch as te_pytorch + self_attention = ModuleSpec( module=SelfAttention, params={"attn_mask_type": AttnMaskType.causal}, submodules=SelfAttentionSubmodules( - linear_qkv=TELayerNormColumnParallelLinear if use_te else ColumnParallelLinear, - core_attention=TEDotProductAttention if use_te else DotProductAttention, - linear_proj=TERowParallelLinear if use_te else RowParallelLinear, + linear_qkv=linear_qkv, + core_attention=core_attention, + linear_proj=linear_proj, q_layernorm=ln, k_layernorm=ln, ), diff --git a/megatron/core/models/retro/decoder_spec.py b/megatron/core/models/retro/decoder_spec.py index 6539348143f..3c30b9bf314 100644 --- a/megatron/core/models/retro/decoder_spec.py +++ b/megatron/core/models/retro/decoder_spec.py @@ -52,6 +52,9 @@ HAVE_TE = True except ImportError: + TEDotProductAttention = None + TEColumnParallelLinear = None + TERowParallelLinear = None HAVE_TE = False @@ -75,6 +78,9 @@ def get_retro_decoder_layer_te_spec( """ spec = get_gpt_layer_with_transformer_engine_spec() spec.submodules.pre_cross_attn_layernorm = TENorm + assert TEDotProductAttention is not None + assert TEColumnParallelLinear is not None + assert TERowParallelLinear is not None spec.submodules.cross_attention = ModuleSpec( module=RetroDecoderCrossAttention, params={"encoder_block_spec": encoder_block_spec}, diff --git a/megatron/core/models/retro/encoder_spec.py b/megatron/core/models/retro/encoder_spec.py index a7cb76ca19b..2e255073b5e 100644 --- a/megatron/core/models/retro/encoder_spec.py +++ b/megatron/core/models/retro/encoder_spec.py @@ -32,6 +32,9 @@ HAVE_TE = True except ImportError: + TEDotProductAttention = None + TEColumnParallelLinear = None + TERowParallelLinear = None HAVE_TE = False try: @@ -64,6 +67,9 @@ def get_retro_encoder_layer_te_spec() -> ModuleSpec: """ spec = get_gpt_layer_with_transformer_engine_spec() spec.submodules.pre_cross_attn_layernorm = TENorm + assert TEDotProductAttention is not None + assert TEColumnParallelLinear is not None + assert TERowParallelLinear is not None spec.submodules.cross_attention = ModuleSpec( module=RetroEncoderCrossAttention, params={"attn_mask_type": AttnMaskType.padding}, diff --git a/megatron/core/tensor_parallel/inference_layers.py b/megatron/core/tensor_parallel/inference_layers.py index 05f7b88d095..96300fd2f93 100644 --- a/megatron/core/tensor_parallel/inference_layers.py +++ b/megatron/core/tensor_parallel/inference_layers.py @@ -1,7 +1,7 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -from typing import Callable, Optional +from typing import Callable, Optional, Tuple import torch import torch.distributed as dist @@ -86,7 +86,7 @@ def __init__( ), "--transformer-impl=inference_optimized requires --sequence-parallel" @torch.no_grad() - def forward(self, x: torch.Tensor) -> torch.Tensor: + def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: """ Forward pass. """ @@ -108,7 +108,7 @@ def __init__( output_size: int, *, config: ModelParallelConfig, - init_method: Callable, + init_method: Callable[[torch.Tensor], None], bias: bool, input_is_parallel: bool, skip_bias_add: bool, @@ -141,7 +141,7 @@ def __init__( ), "--transformer-impl=inference_optimized requires --sequence-parallel" @torch.no_grad() - def forward(self, x: torch.Tensor) -> torch.Tensor: + def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: """ Forward pass. """ diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index e6e65425b23..6de85f206c5 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -795,7 +795,7 @@ def __init__( output_size, *, config: ModelParallelConfig, - init_method: Callable, + init_method: Callable[[torch.Tensor], None], bias=True, gather_output=False, stride=1, @@ -805,7 +805,7 @@ def __init__( embedding_activation_buffer: Optional[List[torch.Tensor]] = None, grad_output_buffer: Optional[List[torch.Tensor]] = None, is_expert: bool = False, - tp_comm_buffer_name: str = None, # Not used + tp_comm_buffer_name: Optional[str] = None, # Not used disable_grad_reduce: bool = False, tp_group: Optional[torch.distributed.ProcessGroup] = None, ): @@ -948,7 +948,7 @@ def forward( input_: torch.Tensor, weight: Optional[torch.Tensor] = None, runtime_gather_output: Optional[bool] = None, - ): + ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: """Forward of ColumnParallelLinear Args: @@ -1064,6 +1064,9 @@ def __repr__(self): f"out_features={self.output_size}, bias={use_bias}, TP={tp})" ) + def backward_dw(self) -> None: + pass + class RowParallelLinear(torch.nn.Module): """Linear layer with row parallelism. @@ -1222,7 +1225,7 @@ def _forward_impl(self, input, weight, *args, **kwargs): else: return linear_with_grad_accumulation_and_async_allreduce(input, weight, *args, **kwargs) - def forward(self, input_): + def forward(self, input_: torch.Tensor) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: """Forward of RowParallelLinear Args: @@ -1301,3 +1304,6 @@ def __repr__(self): f"{type(self).__name__}(in_features={self.input_size}, " f"out_features={self.output_size}, bias={use_bias}, TP={tp})" ) + + def backward_dw(self) -> None: + pass diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py index 1051799db94..1a69a71d5f5 100644 --- a/megatron/core/transformer/attention.py +++ b/megatron/core/transformer/attention.py @@ -1,13 +1,15 @@ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. from abc import ABC, abstractmethod +from collections.abc import Callable from dataclasses import dataclass -from typing import NoReturn, Optional, Tuple, Union +from typing import NoReturn, Optional, Protocol, Tuple, Union, cast import torch from torch import Tensor +from torch import distributed as dist -from megatron.core import tensor_parallel +from megatron.core import tensor_parallel, typed_torch from megatron.core.inference.contexts import BaseInferenceContext from megatron.core.models.common.embeddings.rope_utils import ( apply_rotary_pos_emb, @@ -25,7 +27,6 @@ from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.identity_op import IdentityOp from megatron.core.transformer.module import MegatronModule -from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.utils import ( deprecate_inference_params, divide, @@ -107,17 +108,143 @@ HAVE_FUSED_QKV_ROPE = False +class Linear(Protocol): + """Protocol for a linear_q or linear_kv layer.""" + + def forward(self, hidden_states: Tensor, /) -> Tuple[Tensor, Optional[Tensor]]: ... + + +class LinearBuilder(Protocol): + """Protocol for building linear_q or linear_kv layers.""" + + def __call__( + self, + input_size: int, + output_size: int, + /, + *, + config: TransformerConfig, + init_method: Callable[[Tensor], None], + gather_output: bool, + bias: bool, + skip_bias_add: bool, + is_expert: bool, + ) -> Linear: ... + + +class LinearQKV(Protocol): + """Protocol for a linear QKV layer.""" + + def forward(self, hidden_states: Tensor, /) -> Tuple[Tensor, Optional[Tensor]]: ... + + def backward_dw(self) -> None: ... + + +class LinearQKVBuilder(Protocol): + """Protocol for building linear QKV layers.""" + + def __call__( + self, + input_size: int, + output_size: int, + /, + *, + config: TransformerConfig, + init_method: Callable[[Tensor], None], + gather_output: bool, + bias: bool, + skip_bias_add: bool, + is_expert: bool, + tp_comm_buffer_name: str, + tp_group: Optional[dist.ProcessGroup], + ) -> LinearQKV: ... + + +class CoreAttention(Protocol): + """Protocol for the core_attention layer.""" + + def forward( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attention_mask: torch.Tensor, + /, + *, + attn_mask_type: AttnMaskType, + attention_bias: Optional[torch.Tensor], + packed_seq_params: Optional[PackedSeqParams], + ) -> torch.Tensor: ... + + +class CoreAttentionBuilder(Protocol): + """Protocol for building core attention layers.""" + + def __call__( + self, + /, + *, + config: TransformerConfig, + layer_number: int, + attn_mask_type: AttnMaskType, + attention_type: str, + cp_comm_type: Optional[str], + softmax_scale: Optional[float], + pg_collection: ProcessGroupCollection, + ) -> CoreAttention: ... + + +class LinearProj(Protocol): + """Protocol for a linear_proj layer.""" + + def forward(self, hidden_states: Tensor, /) -> Tuple[Tensor, Optional[Tensor]]: ... + + def backward_dw(self) -> None: ... + + +class LinearProjBuilder(Protocol): + """Protocol for building linear_proj layers.""" + + def __call__( + self, + input_size: int, + output_size: int, + /, + *, + config: TransformerConfig, + init_method: Callable[[Tensor], None], + bias: bool, + input_is_parallel: bool, + skip_bias_add: bool, + is_expert: bool, + tp_comm_buffer_name: str, + tp_group: Optional[dist.ProcessGroup], + ) -> LinearProj: ... + + +class LayerNorm(Protocol): + """Protocol for a layernorm layer.""" + + def forward(self, input: Tensor, /) -> Tensor: ... + + +class LayerNormBuilder(Protocol): + """Protocol for building layernorm layers.""" + + def __call__(self, *, hidden_size: int, config: TransformerConfig, eps: float) -> LayerNorm: ... + + @dataclass class SelfAttentionSubmodules: """ Configuration class for specifying the submodules of a self-attention. """ - linear_qkv: Union[ModuleSpec, type] = None - core_attention: Union[ModuleSpec, type] = None - linear_proj: Union[ModuleSpec, type] = None - q_layernorm: Union[ModuleSpec, type] = None - k_layernorm: Union[ModuleSpec, type] = None + linear_qkv: LinearQKVBuilder + core_attention: CoreAttentionBuilder + linear_proj: LinearProjBuilder + q_layernorm: Optional[LayerNormBuilder] = None + k_layernorm: Optional[LayerNormBuilder] = None @dataclass @@ -126,10 +253,10 @@ class CrossAttentionSubmodules: Configuration class for specifying the submodules of a cross-attention. """ - linear_q: Union[ModuleSpec, type] = None - linear_kv: Union[ModuleSpec, type] = None - core_attention: Union[ModuleSpec, type] = None - linear_proj: Union[ModuleSpec, type] = None + linear_q: LinearBuilder + linear_kv: LinearBuilder + core_attention: CoreAttentionBuilder + linear_proj: LinearProjBuilder class Attention(MegatronModule, ABC): @@ -146,8 +273,8 @@ def __init__( layer_number: int, attn_mask_type: AttnMaskType, attention_type: str, - cp_comm_type: str = None, - pg_collection: ProcessGroupCollection = None, + cp_comm_type: Optional[str] = None, + pg_collection: Optional[ProcessGroupCollection] = None, ): super().__init__(config=config) @@ -185,8 +312,7 @@ def __init__( self.key_hidden_size = self.hidden_size_per_attention_head self.val_hidden_size = self.hidden_size_per_attention_head - self.core_attention = build_module( - submodules.core_attention, + self.core_attention = submodules.core_attention( config=self.config, layer_number=self.layer_number, attn_mask_type=self.attn_mask_type, @@ -202,12 +328,11 @@ def __init__( ) # Output. - self.linear_proj = build_module( - submodules.linear_proj, + self.linear_proj = submodules.linear_proj( self.query_projection_size, self.config.hidden_size, config=self.config, - init_method=self.config.output_layer_init_method, + init_method=cast(Callable[[torch.Tensor], None], self.config.output_layer_init_method), bias=self.config.add_bias_linear, input_is_parallel=True, skip_bias_add=True, @@ -254,7 +379,7 @@ def custom_forward(*inputs): attention_mask = inputs[3] attn_mask_type = inputs[5] attn_mask_type = AttnMaskType(attn_mask_type.item()) - output_ = self.core_attention( + output_ = typed_torch.apply_module(self.core_attention)( query, key, value, @@ -813,7 +938,7 @@ def forward( ) out = output.transpose(0, 1).contiguous() context_layer = out.view(out.size(0), out.size(1), -1) - output, bias = self.linear_proj(context_layer) + output, bias = typed_torch.apply_module(self.linear_proj)(context_layer) return output, bias if ( @@ -920,7 +1045,7 @@ def forward( else: if inference_context is None or inference_context.is_static_batching(): # Static batching attention kernel. - core_attn_out = self.core_attention( + core_attn_out = typed_torch.apply_module(self.core_attention)( query, key, value, @@ -962,7 +1087,7 @@ def forward( # ================= nvtx_range_push(suffix="linear_proj") - output, bias = self.linear_proj(core_attn_out) + output, bias = typed_torch.apply_module(self.linear_proj)(core_attn_out) nvtx_range_pop(suffix="linear_proj") return output, bias @@ -998,12 +1123,11 @@ def __init__( pg_collection=pg_collection, ) - self.linear_qkv = build_module( - submodules.linear_qkv, + self.linear_qkv = submodules.linear_qkv( self.config.hidden_size, self.query_projection_size + 2 * self.kv_projection_size, config=self.config, - init_method=self.config.init_method, + init_method=cast(Callable[[torch.Tensor], None], self.config.init_method), gather_output=False, bias=self.config.add_bias_linear or self.config.add_qkv_bias, skip_bias_add=False, @@ -1013,8 +1137,7 @@ def __init__( ) if submodules.q_layernorm is not None: - self.q_layernorm = build_module( - submodules.q_layernorm, + self.q_layernorm = submodules.q_layernorm( hidden_size=self.hidden_size_per_attention_head, config=self.config, eps=self.config.layernorm_epsilon, @@ -1023,8 +1146,7 @@ def __init__( self.q_layernorm = None if submodules.k_layernorm is not None: - self.k_layernorm = build_module( - submodules.k_layernorm, + self.k_layernorm = submodules.k_layernorm( hidden_size=self.hidden_size_per_attention_head, config=self.config, eps=self.config.layernorm_epsilon, @@ -1047,6 +1169,17 @@ def run_realtime_tests(self): if not self.config.qk_layernorm: return + assert ( + self.q_layernorm is not None + and hasattr(self.q_layernorm, 'weight') + and hasattr(self.q_layernorm, 'bias') + ) + assert ( + self.k_layernorm is not None + and hasattr(self.k_layernorm, 'weight') + and hasattr(self.k_layernorm, 'bias') + ) + # check that all tensor parallel and data parallel ranks have the same # Q & K layernorm parameters. rank = get_data_parallel_rank() @@ -1109,7 +1242,7 @@ def get_query_key_value_tensors(self, hidden_states, key_value_states=None, spli the unsplit mixed_qkv tensor is returned. """ # Attention heads [sq, b, h] --> [sq, b, ng * (np/ng + 2) * hn)] - mixed_qkv, _ = self.linear_qkv(hidden_states) + mixed_qkv, _ = typed_torch.apply_module(self.linear_qkv)(hidden_states) # [sq, b, hp] --> [sq, b, ng, (np/ng + 2) * hn] new_tensor_shape = mixed_qkv.size()[:-1] + ( @@ -1150,10 +1283,10 @@ def get_query_key_value_tensors(self, hidden_states, key_value_states=None, spli query = query.reshape(query.size(0), query.size(1), -1, self.hidden_size_per_attention_head) if self.q_layernorm is not None: - query = self.q_layernorm(query) + query = typed_torch.apply_module(self.q_layernorm)(query) if self.k_layernorm is not None: - key = self.k_layernorm(key) + key = typed_torch.apply_module(self.k_layernorm)(key) if self.config.test_mode: self.run_realtime_tests() @@ -1193,8 +1326,8 @@ def __init__( submodules: CrossAttentionSubmodules, layer_number: int, attn_mask_type=AttnMaskType.padding, - cp_comm_type: str = None, - pg_collection: ProcessGroupCollection = None, + cp_comm_type: Optional[str] = None, + pg_collection: Optional[ProcessGroupCollection] = None, ): super().__init__( config=config, @@ -1210,24 +1343,22 @@ def __init__( raise ValueError("Group query attention is not currently supported in cross attention.") assert self.query_projection_size == self.kv_projection_size - self.linear_q = build_module( - submodules.linear_q, + self.linear_q = submodules.linear_q( self.config.hidden_size, self.query_projection_size, config=self.config, - init_method=self.config.init_method, + init_method=cast(Callable[[torch.Tensor], None], self.config.init_method), gather_output=False, bias=self.config.add_bias_linear, skip_bias_add=False, is_expert=False, ) - self.linear_kv = build_module( - submodules.linear_kv, + self.linear_kv = submodules.linear_kv( self.config.hidden_size, 2 * self.kv_projection_size, config=self.config, - init_method=self.config.init_method, + init_method=cast(Callable[[torch.Tensor], None], self.config.init_method), gather_output=False, bias=self.config.add_bias_linear, skip_bias_add=False, @@ -1241,7 +1372,7 @@ def get_query_key_value_tensors(self, hidden_states, key_value_states, split_qkv """ assert split_qkv, "split_qkv must be True for CrossAttention" # Attention heads [sk, b, h] --> [sk, b, (np * 2 * hn)] - mixed_kv, _ = self.linear_kv(key_value_states) + mixed_kv, _ = typed_torch.apply_module(self.linear_kv)(key_value_states) # [sk, b, (np * 2 * hn)] --> [sk, b, np, 2 * hn] new_tensor_shape = mixed_kv.size()[:-1] + ( @@ -1254,7 +1385,7 @@ def get_query_key_value_tensors(self, hidden_states, key_value_states, split_qkv (key, value) = tensor_parallel.split_tensor_along_last_dim(mixed_kv, 2) # Attention head [sq, b, h] --> [sq, b, hp] - query, _ = self.linear_q(hidden_states) + query, _ = typed_torch.apply_module(self.linear_q)(hidden_states) # [sq, b, hp] --> [sq, b, np, hn] new_tensor_shape = query.size()[:-1] + ( diff --git a/megatron/core/transformer/dot_product_attention.py b/megatron/core/transformer/dot_product_attention.py index f3711c86ebd..eed6301a007 100644 --- a/megatron/core/transformer/dot_product_attention.py +++ b/megatron/core/transformer/dot_product_attention.py @@ -45,10 +45,10 @@ def __init__( layer_number: int, attn_mask_type: AttnMaskType, attention_type: str, - attention_dropout: float = None, - softmax_scale: float = None, - cp_comm_type: str = None, - pg_collection: ProcessGroupCollection = None, + attention_dropout: Optional[float] = None, + softmax_scale: Optional[float] = None, + cp_comm_type: Optional[str] = None, + pg_collection: Optional[ProcessGroupCollection] = None, ): super().__init__(config=config) @@ -143,10 +143,10 @@ def forward( key: Tensor, value: Tensor, attention_mask: Tensor, - attn_mask_type: AttnMaskType = None, - attention_bias: Tensor = None, + attn_mask_type: Optional[AttnMaskType] = None, + attention_bias: Optional[Tensor] = None, packed_seq_params: Optional[PackedSeqParams] = None, - ): + ) -> Tensor: """Forward.""" assert packed_seq_params is None, ( "Packed sequence is not supported by DotProductAttention." diff --git a/megatron/core/transformer/identity_op.py b/megatron/core/transformer/identity_op.py index 5d9388ffcc6..970cafb592f 100644 --- a/megatron/core/transformer/identity_op.py +++ b/megatron/core/transformer/identity_op.py @@ -1,4 +1,6 @@ # Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. +from typing import Any + import torch @@ -7,10 +9,10 @@ class IdentityOp(torch.nn.Module): This is a placeholder for IdentityOp(x) -> x """ - def __init__(self, *args, **kwargs): + def __init__(self, *args: Any, **kwargs: Any): super().__init__() - def forward(self, x, *args, **kwargs): + def forward(self, x: torch.Tensor, *args: Any, **kwargs: Any) -> torch.Tensor: return x @@ -21,8 +23,8 @@ class IdentityFuncOp(IdentityOp): return a function at runtime based on passed arguments """ - def __init__(self, *args, **kwargs): + def __init__(self, *args: Any, **kwargs: Any): super().__init__() - def forward(self, *args, **kwargs): + def forward(self, *args: Any, **kwargs: Any): return super().forward diff --git a/megatron/core/transformer/spec_utils.py b/megatron/core/transformer/spec_utils.py index b3de8541734..73f46a8df54 100644 --- a/megatron/core/transformer/spec_utils.py +++ b/megatron/core/transformer/spec_utils.py @@ -2,7 +2,7 @@ import types from dataclasses import dataclass, field -from typing import Tuple, Union +from typing import Any, Tuple, Union @dataclass @@ -24,7 +24,16 @@ class ModuleSpec: module: Union[Tuple, type] params: dict = field(default_factory=lambda: {}) - submodules: type = None + submodules: object = None + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + """Builds an instance of the module from the spec. + + Args: + *args: Positional arguments to be passed to the module init. + **kwargs: Keyword arguments to be passed to the module init. + """ + return build_module(self, *args, **kwargs) def import_module(module_path: Tuple[str]): diff --git a/megatron/core/typed_torch.py b/megatron/core/typed_torch.py new file mode 100644 index 00000000000..8ca163315a8 --- /dev/null +++ b/megatron/core/typed_torch.py @@ -0,0 +1,31 @@ +"""Utilities for improved type hinting with torch interfaces.""" + +from collections.abc import Callable +from typing import Generic, ParamSpec, Protocol, TypeVar + +import torch + +P = ParamSpec('P') +R_co = TypeVar('R_co', covariant=True) + + +class _Module(Generic[P, R_co], Protocol): + """Protocol allowing us to unwrap `forward`.""" + + def forward(self, *args: P.args, **kwargs: P.kwargs) -> R_co: ... + + +def apply_module(m: _Module[P, R_co], *, check_subclass: bool = True) -> Callable[P, R_co]: + """Returns the provided module unchanged, but with correct type hints. + + Args: + m: An instance of a subclass of `torch.nn.Module`. + check_subclass: If `True`, checks that `m` is a subclass of + `torch.nn.Module` and raises a `TypeError` if not. + + Returns: + That module unchanged, but with correct type hints. + """ + if check_subclass and not issubclass(type(m), torch.nn.Module): + raise TypeError(f'{type(m)} is not a subclass of torch.nn.Module') + return m # type: ignore