Skip to content
Merged
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
16 changes: 13 additions & 3 deletions gpt_builders.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
)
from megatron.core.models.gpt.experimental_attention_variant_module_specs import (
get_transformer_block_with_experimental_attention_variant_spec,
get_transformer_layer_with_experimental_attention_variant_spec,
)
from megatron.core.models.gpt.heterogeneous.heterogeneous_layer_specs import (
get_gpt_heterogeneous_layer_spec,
Expand Down Expand Up @@ -66,9 +67,18 @@ def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_
transformer_layer_spec_for_mtp = _get_transformer_layer_spec(use_te, config)
else:
# Define the decoder block spec
decoder_layer_specs = get_gpt_decoder_layer_specs(
config, use_transformer_engine=use_te, normalization=args.normalization, qk_l2_norm=args.qk_l2_norm, vp_stage=vp_stage
)
if args.experimental_attention_variant is not None:
decoder_layer_specs = get_transformer_layer_with_experimental_attention_variant_spec(
config=config
)
else:
decoder_layer_specs = get_gpt_decoder_layer_specs(
config,
use_transformer_engine=use_te,
normalization=args.normalization,
qk_l2_norm=args.qk_l2_norm,
vp_stage=vp_stage,
)
transformer_layer_spec_for_mtp = decoder_layer_specs[-1]
# Use spec of the last layer in decoder block as spec of the transformer layer in MTP
mtp_block_spec = get_gpt_mtp_block_spec(
Expand Down
15 changes: 13 additions & 2 deletions megatron/core/inference/contexts/dynamic_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -1675,6 +1675,7 @@ def apply_rotary_emb_query(
cu_seqlens=cu_seqlens_q,
cp_group=cp_group,
mscale=mscale,
mla_rotary_interleaved=config.multi_latent_attention,
)
return query

Expand Down Expand Up @@ -1709,11 +1710,21 @@ def apply_rotary_emb_key(
f"paused_request_count={self.paused_request_count}"
)
key = apply_rotary_pos_emb(
t=key[:n], freqs=key_emb[:n], config=config, cp_group=cp_group, mscale=mscale
t=key[:n],
freqs=key_emb[:n],
config=config,
cp_group=cp_group,
mscale=mscale,
mla_rotary_interleaved=config.multi_latent_attention,
)
else:
key[:n] = apply_rotary_pos_emb(
t=key[:n], freqs=key_emb[:n], config=config, cp_group=cp_group, mscale=mscale
t=key[:n],
freqs=key_emb[:n],
config=config,
cp_group=cp_group,
mscale=mscale,
mla_rotary_interleaved=config.multi_latent_attention,
)
return key

Expand Down
41 changes: 33 additions & 8 deletions megatron/core/models/common/embeddings/rope_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,8 +93,9 @@ def _apply_rotary_pos_emb_bshd(
t: Tensor,
freqs: Tensor,
rotary_interleaved: bool = False,
multi_latent_attention: bool = False,
mla_rotary_interleaved: bool = False,
mscale: float = 1.0,
multi_latent_attention: Optional[bool] = None,
) -> Tensor:
"""Apply rotary positional embedding to input tensor T.

Expand All @@ -103,16 +104,26 @@ def _apply_rotary_pos_emb_bshd(
Args:
t (Tensor): Input tensor T is of shape [seq_length, ... , dim]
freqs (Tensor): Rotary Positional embedding tensor freq is of shape [seq_length, ..., dim]
rotary_interleaved (bool): Whether to apply interleaving in the rotate half function.
mla_rotary_interleaved (bool): Whether to apply MLA-style interleaving for RoPE.
mscale (float): The scaling factor for the RoPE.

Returns:
Tensor: The input tensor after applying RoPE
"""
if multi_latent_attention is not None:
warnings.warn(
"multi_latent_attention is deprecated. Please use mla_rotary_interleaved instead.",
DeprecationWarning,
)
mla_rotary_interleaved = multi_latent_attention

rot_dim = freqs.shape[-1]

# ideally t_pass is empty so rotary pos embedding is applied to all tensor t
t, t_pass = t[..., :rot_dim], t[..., rot_dim:]

if multi_latent_attention:
if mla_rotary_interleaved:
x1 = t[..., 0::2]
x2 = t[..., 1::2]
t = torch.cat((x1, x2), dim=-1)
Expand Down Expand Up @@ -180,9 +191,10 @@ def _apply_rotary_pos_emb_thd(
cu_seqlens: Tensor,
freqs: Tensor,
rotary_interleaved: bool = False,
multi_latent_attention: bool = False,
mla_rotary_interleaved: bool = False,
mscale: float = 1.0,
cp_group: torch.distributed.ProcessGroup = None,
multi_latent_attention: Optional[bool] = None,
) -> Tensor:
"""A baseline implementation of applying RoPE for `thd` format.

Expand All @@ -196,6 +208,12 @@ def _apply_rotary_pos_emb_thd(
Returns:
Tensor: Shape [t, h, d]. The input tensor after applying RoPE.
"""
if multi_latent_attention is not None:
warnings.warn(
"multi_latent_attention is deprecated. Please use mla_rotary_interleaved instead.",
DeprecationWarning,
)
mla_rotary_interleaved = multi_latent_attention

if cp_group is None:
raise ValueError("cp_group must be provided for THD format RoPE")
Expand Down Expand Up @@ -226,7 +244,7 @@ def _apply_rotary_pos_emb_thd(
t.unsqueeze(1),
freqs_packed,
rotary_interleaved=rotary_interleaved,
multi_latent_attention=multi_latent_attention,
mla_rotary_interleaved=mla_rotary_interleaved,
mscale=mscale,
).squeeze(1)
else:
Expand All @@ -242,7 +260,7 @@ def _apply_rotary_pos_emb_thd(
t.unsqueeze(1),
freqs_packed,
rotary_interleaved=rotary_interleaved,
multi_latent_attention=multi_latent_attention,
mla_rotary_interleaved=mla_rotary_interleaved,
mscale=mscale,
).squeeze(1)

Expand All @@ -254,6 +272,7 @@ def apply_rotary_pos_emb(
cu_seqlens: Optional[Tensor] = None,
mscale: float = 1.0,
cp_group: torch.distributed.ProcessGroup = None,
mla_rotary_interleaved: bool = False,
):
"""
Reroute to the appropriate apply_rotary_pos_emb function depending on
Expand Down Expand Up @@ -282,6 +301,12 @@ def apply_rotary_pos_emb(
"Using unfused implementation."
)
use_unfused = True
if mla_rotary_interleaved:
Comment thread
yuzhongw-nvidia marked this conversation as resolved.
warnings.warn(
"apply_rope_fusion does not support MLA-style interleaving in RoPE."
"Using unfused implementation."
)
use_unfused = True
if not use_unfused:
assert fused_apply_rotary_pos_emb is not None, "apply_rope_fusion is not available."
return fused_apply_rotary_pos_emb(t, freqs, interleaved=config.rotary_interleaved)
Expand All @@ -301,7 +326,7 @@ def apply_rotary_pos_emb(
t,
freqs,
rotary_interleaved=config.rotary_interleaved,
multi_latent_attention=config.multi_latent_attention,
mla_rotary_interleaved=mla_rotary_interleaved,
mscale=mscale,
)
else:
Expand All @@ -310,7 +335,7 @@ def apply_rotary_pos_emb(
cu_seqlens,
freqs,
rotary_interleaved=config.rotary_interleaved,
multi_latent_attention=config.multi_latent_attention,
mla_rotary_interleaved=mla_rotary_interleaved,
mscale=mscale,
cp_group=cp_group,
)
Expand Down Expand Up @@ -339,7 +364,7 @@ def apply_rotary_pos_emb_with_cos_sin(
t,
freqs,
rotary_interleaved=rotary_interleaved,
multi_latent_attention=False,
mla_rotary_interleaved=False,
mscale=1.0,
)
else:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -83,17 +83,6 @@ def get_dsa_module_spec_for_backend(
assert config.multi_latent_attention, "Currently only MLA supports sparse attention."
assert config.qk_l2_norm is False, "qk_l2_norm is not supported with MLA."

linear_q_up_proj = (
backend.column_parallel_layer_norm_linear()
if config.qk_layernorm
else backend.column_parallel_linear()
)
linear_kv_up_proj = (
backend.column_parallel_layer_norm_linear()
if config.qk_layernorm
else backend.column_parallel_linear()
)

# Because TransformerEngine does not support sparse attention yet, we use local
# implementation whether the backend is TransformerEngine or not.
core_attention = ModuleSpec(
Expand All @@ -111,19 +100,27 @@ def get_dsa_module_spec_for_backend(
),
)

# Adjust for RMS norm.
rms_norm = config.normalization == "RMSNorm"
# DSA indexer requires normalized q as input, so here we cannot fuse qk layernorm
# with linear projection and have to use unfused qk layernorm.
qk_norm = (
backend.layer_norm(rms_norm=rms_norm, for_qk=True) if config.qk_layernorm else IdentityOp
)

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_q_up_proj=backend.column_parallel_linear(),
linear_kv_down_proj=backend.linear(),
linear_kv_up_proj=linear_kv_up_proj,
linear_kv_up_proj=backend.column_parallel_linear(),
core_attention=core_attention,
linear_proj=backend.row_parallel_linear(),
q_layernorm=IdentityOp,
kv_layernorm=IdentityOp,
q_layernorm=qk_norm,
kv_layernorm=qk_norm,
),
metainfo={"fuse_input_layernorm": False},
)
Expand Down Expand Up @@ -154,12 +151,12 @@ def get_experimental_attention_variant_module_spec(
##########


def get_transformer_block_with_experimental_attention_variant_spec(
config: TransformerConfig, vp_stage: Optional[int] = None, pp_rank: Optional[int] = None
) -> TransformerBlockSubmodules:
"""Build transformer block spec with experimental attention variants (e.g., linear attention).
def get_transformer_layer_with_experimental_attention_variant_spec(
config: TransformerConfig, backend: BackendSpecProvider = None
) -> List[ModuleSpec]:
"""Build transformer layer specs with experimental attention variants (e.g., linear attention).

This function constructs a heterogeneous transformer block that supports mixing different
This function is for constructing a heterogeneous transformer that supports mixing different
attention mechanisms (experimental vs standard) and MLP types (MoE vs dense) across layers.
**Note that, this API is a experimental API in the short term, and might be deprecated in the
future. In the long run, we will move to a new design that better support hybrid models.**
Expand All @@ -175,22 +172,19 @@ def get_transformer_block_with_experimental_attention_variant_spec(
2. Per-Layer Spec Construction: Iterates through layers, constructing transformer
layer specs based on attention and MLP patterns.

3. Pipeline Slicing: Extracts layer specs for the current pipeline stage.

Args:
config: Transformer configuration containing model hyperparameters and feature flags.
vp_stage: Virtual pipeline stage index for interleaved pipeline parallelism.
pp_rank: Pipeline model parallel rank.

Returns:
TransformerBlockSubmodules containing per-layer specs and final layer norm.
List[ModuleSpec] containing per-layer specs.

Note:
Currently only supports transformer_engine backend. Kitchen backend can be used as a
wrapper with TE fallback for unsupported operations.
"""

backend = _get_backend_spec_provider(config=config)
if backend is None:
backend = _get_backend_spec_provider(config=config)

# Get attention patterns and specs
experimental_attention_pattern = [0] * config.num_layers
Expand Down Expand Up @@ -271,6 +265,42 @@ def get_transformer_block_with_experimental_attention_variant_spec(
)
)

return layer_specs


def get_transformer_block_with_experimental_attention_variant_spec(
config: TransformerConfig, vp_stage: Optional[int] = None, pp_rank: Optional[int] = None
) -> TransformerBlockSubmodules:
"""Build transformer block spec with experimental attention variants (e.g., linear attention).

This function constructs a heterogeneous transformer block that supports mixing different
attention mechanisms (experimental vs standard) and MLP types (MoE vs dense) across layers.
**Note that, this API is a experimental API in the short term, and might be deprecated in the
future. In the long run, we will move to a new design that better support hybrid models.**

Constructing transformer layer specs by
`get_transformer_layer_with_experimental_attention_variant_spec` and then slicing the
layer specs to only include the layers that are built in this pipeline stage.

Args:
config: Transformer configuration containing model hyperparameters and feature flags.
vp_stage: Virtual pipeline stage index for interleaved pipeline parallelism.
pp_rank: Pipeline model parallel rank.

Returns:
TransformerBlockSubmodules containing per-layer specs and final layer norm.

Note:
Currently only supports transformer_engine backend. Kitchen backend can be used as a
wrapper with TE fallback for unsupported operations.
"""

backend = _get_backend_spec_provider(config=config)

layer_specs = get_transformer_layer_with_experimental_attention_variant_spec(
config=config, backend=backend
)

# Slice the layer specs to only include the layers that are built in this pipeline stage.
if config.pipeline_model_parallel_layout is not None:
local_layer_ids = config.pipeline_model_parallel_layout.get_layer_id_list(
Expand All @@ -284,6 +314,7 @@ def get_transformer_block_with_experimental_attention_variant_spec(
layer_specs = [layer_specs[layer_id] for layer_id in local_layer_ids]

# Get GPT decoder block spec
rms_norm = config.normalization == "RMSNorm"

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.

looks like rms_norm var is just added now, does this mean this codepath was broken on main before this change ?

gpt_decoder_block_spec = TransformerBlockSubmodules(
layer_specs=layer_specs, layer_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False)
)
Expand Down
10 changes: 9 additions & 1 deletion megatron/core/models/gpt/gpt_layer_specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -569,6 +569,11 @@ def get_gpt_decoder_layer_specs(
pp_rank: Optional[int] = None,
) -> TransformerBlockSubmodules:
"""GPT block spec."""
assert config.experimental_attention_variant is None, (
"Experimental attention variant is not supported with get_gpt_decoder_layer_specs, "
f"but got {config.experimental_attention_variant=}."
)

if use_transformer_engine:
layer_norm_impl = TENorm
dense_layer_spec = get_gpt_layer_with_transformer_engine_spec(
Expand Down Expand Up @@ -680,13 +685,16 @@ def get_gpt_decoder_block_spec(
layer_specs = get_gpt_decoder_layer_specs(
config, use_transformer_engine, normalization, qk_l2_norm
)

# Slice the layer specs to only include the layers that are built in this pipeline stage.
# Note: MCore layer_number starts at 1
num_layers_to_build = get_num_layers_to_build(config, vp_stage=vp_stage, pp_rank=pp_rank)

if config.pipeline_model_parallel_layout is not None:
layout = config.pipeline_model_parallel_layout
assert isinstance(layout, PipelineParallelLayerLayout)
assert isinstance(
layout, PipelineParallelLayerLayout
), f"Invalid pipeline model parallel layout: {layout}"
local_layer_specs = [
layer_specs[layer_id]
for layer_id in layout.get_layer_id_list(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -605,6 +605,7 @@ def qkv_up_proj_and_rope_apply(q_compressed, kv_compressed, k_pos_emb, rotary_po
cu_seqlens=cu_seqlens_q,
mscale=mscale,
cp_group=self.pg_collection.cp,
mla_rotary_interleaved=True,
)
# k_pos_emb:[num_tokens, 1, qk_pos_emb_head_dim]
k_pos_emb = apply_rotary_pos_emb(
Expand All @@ -614,6 +615,7 @@ def qkv_up_proj_and_rope_apply(q_compressed, kv_compressed, k_pos_emb, rotary_po
cu_seqlens=cu_seqlens_kv,
mscale=mscale,
cp_group=self.pg_collection.cp,
mla_rotary_interleaved=True,
)

# query: [num_tokens, n, (kv_lora_rank + qk_pos_emb_head_dim)]
Expand Down
Loading
Loading