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
54 changes: 31 additions & 23 deletions megatron/core/models/bert/bert_layer_specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
)


Expand Down
10 changes: 8 additions & 2 deletions megatron/core/models/bert/bert_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,17 +15,18 @@
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,
)
from megatron.core.transformer.enums import AttnBackend, AttnMaskType, ModelType
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():
Expand Down Expand Up @@ -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
Expand Down
Loading