Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 19 additions & 11 deletions examples/multimodal/layer_specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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 (
Expand All @@ -26,6 +26,9 @@

HAVE_TE = True
except ImportError:
TELayerNormColumnParallelLinear = None
TEDotProductAttention = None
TERowParallelLinear = None
HAVE_TE = False

try:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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(
Expand All @@ -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

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Creating a not_none typing utility (or just using the one in PyTorch) might make your lives a bit easier.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yeah that'd be fine! I don't actually want to do it this way (it's clearly sorta ugly), it was just a cheap way to get the type-checker to pass locally while I was drafting. Sorry for the confusion :)

assert TEDotProductAttention is not None
assert TERowParallelLinear is not None
return ModuleSpec(
module=MambaStack,
submodules=MambaStackSubmodules(
Expand Down Expand Up @@ -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,
Expand All @@ -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(
Expand Down
44 changes: 25 additions & 19 deletions examples/multimodal/nvlm/internvit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
Expand All @@ -60,7 +63,7 @@ class InternViTRMSNorm(MegatronModule):

def __init__(
self,
config,
config: TransformerConfig,
hidden_size: int,
eps: float = 1e-6,
sequence_parallel: bool = False,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
)


Expand Down Expand Up @@ -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,
Expand All @@ -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

Expand Down
11 changes: 10 additions & 1 deletion examples/multimodal/radio/radio_g.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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 (
Expand All @@ -24,6 +27,9 @@

HAVE_TE = True
except ImportError:
TELayerNormColumnParallelLinear = None
TEDotProductAttention = None
TERowParallelLinear = None
HAVE_TE = False

try:
Expand Down Expand Up @@ -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(
Expand Down
10 changes: 7 additions & 3 deletions megatron/core/extensions/kitchen.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand All @@ -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

Expand Down
38 changes: 23 additions & 15 deletions megatron/core/extensions/transformer_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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}
Expand Down Expand Up @@ -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
Expand Down
Loading