From 2ce8b3332c4d3dbb127e34539b9b90aba7f34149 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Wed, 21 Jan 2026 14:19:46 +0800 Subject: [PATCH 01/18] Fix several bugs of experimental attention variant --- gpt_builders.py | 18 ++++-- ...rimental_attention_variant_module_specs.py | 56 +++++++++++++++---- megatron/core/models/gpt/gpt_layer_specs.py | 12 +++- 3 files changed, 68 insertions(+), 18 deletions(-) diff --git a/gpt_builders.py b/gpt_builders.py index f7a34e7203a..aa47e8458fa 100644 --- a/gpt_builders.py +++ b/gpt_builders.py @@ -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, @@ -66,10 +67,19 @@ 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 - ) - transformer_layer_spec_for_mtp = decoder_layer_specs[-1] + 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, + ) + mtp_transformer_layer_spec = 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( config, diff --git a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py index 8f6b1a1a3f8..a59a222e252 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -154,12 +154,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.** @@ -175,22 +175,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 @@ -271,6 +268,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( @@ -284,6 +317,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" gpt_decoder_block_spec = TransformerBlockSubmodules( layer_specs=layer_specs, layer_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False) ) diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py index 5aa12747f3c..28107556b21 100755 --- a/megatron/core/models/gpt/gpt_layer_specs.py +++ b/megatron/core/models/gpt/gpt_layer_specs.py @@ -565,10 +565,13 @@ def get_gpt_decoder_layer_specs( use_transformer_engine: bool, normalization: Optional[str] = None, qk_l2_norm: Optional[bool] = False, - vp_stage: Optional[int] = None, - 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( @@ -680,13 +683,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( From e01fdd3dcecf5073724e4c813fbfbedb491f0b16 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Fri, 13 Feb 2026 15:50:25 +0800 Subject: [PATCH 02/18] Update experimental_attention_variant_module_specs.py --- ...rimental_attention_variant_module_specs.py | 27 +++++++++---------- 1 file changed, 12 insertions(+), 15 deletions(-) diff --git a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py index a59a222e252..26e35ce0621 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -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( @@ -111,19 +100,27 @@ def get_dsa_module_spec_for_backend( ), ) + # Adjust for RMS norm. + rms_norm = config.normalization == "RMSNorm" + 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}, ) From 6e13bf8fd8c3e8d32b849e6b70e66feb72917f7e Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Wed, 25 Feb 2026 01:58:36 -0800 Subject: [PATCH 03/18] add comment about why using unfused qk layernorm --- .../models/gpt/experimental_attention_variant_module_specs.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py index 26e35ce0621..78710317f12 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -102,6 +102,8 @@ 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 From 61965852f284acb5e943826152bd48ea0ba9a7fa Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Wed, 25 Feb 2026 23:03:24 -0800 Subject: [PATCH 04/18] fix --- gpt_builders.py | 2 +- .../models/gpt/experimental_attention_variant_module_specs.py | 4 +--- 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/gpt_builders.py b/gpt_builders.py index aa47e8458fa..57b5179b1a0 100644 --- a/gpt_builders.py +++ b/gpt_builders.py @@ -79,7 +79,7 @@ def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None, pg_ qk_l2_norm=args.qk_l2_norm, vp_stage=vp_stage, ) - mtp_transformer_layer_spec = decoder_layer_specs[-1] + 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( config, diff --git a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py index 78710317f12..8231a2a3764 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -105,9 +105,7 @@ def get_dsa_module_spec_for_backend( # 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 + backend.layer_norm(rms_norm=rms_norm, for_qk=True) if config.qk_layernorm else IdentityOp ) attention = ModuleSpec( From 2b6349309b19a0250ae3b9a9880d44c802bda4f6 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Tue, 3 Mar 2026 03:39:20 -0800 Subject: [PATCH 05/18] fix rope order --- .../transformer/experimental_attention_variant/dsa.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/megatron/core/transformer/experimental_attention_variant/dsa.py b/megatron/core/transformer/experimental_attention_variant/dsa.py index 3734db7043f..743029df3a4 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -778,10 +778,12 @@ def __init__( def _apply_rope(self, x: torch.Tensor, rotary_pos_emb: torch.Tensor, mscale: float): """Apply RoPE to the input tensor.""" - # x_nope [seqlen, batch, *, index_head_dim - qk_pos_emb_head_dim] # x_pe [seqlen, batch, *, qk_pos_emb_head_dim] - x_nope, x_pe = torch.split( - x, [self.index_head_dim - self.qk_pos_emb_head_dim, self.qk_pos_emb_head_dim], dim=-1 + # x_nope [seqlen, batch, *, index_head_dim - qk_pos_emb_head_dim] + # To align with DeepSeek's implementation, + # x_pe is placed at the front, and x_nope is placed at the back. + x_pe, x_nope = torch.split( + x, [self.qk_pos_emb_head_dim, self.index_head_dim - self.qk_pos_emb_head_dim], dim=-1 ) x_pe = apply_rotary_pos_emb( x_pe, @@ -792,7 +794,7 @@ def _apply_rope(self, x: torch.Tensor, rotary_pos_emb: torch.Tensor, mscale: flo cp_group=self.pg_collection.cp, ) # [seqlen, batch, *, index_head_dim] - x = torch.cat([x_nope, x_pe], dim=-1) + x = torch.cat([x_pe, x_nope], dim=-1) return x def forward_before_topk( From 865801c37e1e3afc51003429ca4d2cbde632131c Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Mon, 2 Mar 2026 00:41:46 -0800 Subject: [PATCH 06/18] refactor apply_rotary_pos_emb and disable rope interleaving for DSA indexer Update attention.py Update transformer.py Update test_rope.py rename multi_latent_attention -> mla_rotary_interleaved --- .../inference/contexts/dynamic_context.py | 15 ++++++- .../models/common/embeddings/rope_utils.py | 41 +++++++++++++++---- .../experimental_attention_variant/dsa.py | 3 ++ .../transformer/multi_latent_attention.py | 2 + .../fusions/test_mla_yarn_rope_apply.py | 9 +++- 5 files changed, 59 insertions(+), 11 deletions(-) diff --git a/megatron/core/inference/contexts/dynamic_context.py b/megatron/core/inference/contexts/dynamic_context.py index 4a0d0cba518..cbbbfbd47f9 100644 --- a/megatron/core/inference/contexts/dynamic_context.py +++ b/megatron/core/inference/contexts/dynamic_context.py @@ -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 @@ -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 diff --git a/megatron/core/models/common/embeddings/rope_utils.py b/megatron/core/models/common/embeddings/rope_utils.py index 2fd19194813..b990615da29 100644 --- a/megatron/core/models/common/embeddings/rope_utils.py +++ b/megatron/core/models/common/embeddings/rope_utils.py @@ -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. @@ -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) @@ -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. @@ -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") @@ -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: @@ -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) @@ -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 @@ -282,6 +301,12 @@ def apply_rotary_pos_emb( "Using unfused implementation." ) use_unfused = True + if mla_rotary_interleaved: + 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) @@ -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: @@ -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, ) @@ -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: diff --git a/megatron/core/transformer/experimental_attention_variant/dsa.py b/megatron/core/transformer/experimental_attention_variant/dsa.py index 743029df3a4..5c5f77363dc 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -792,6 +792,9 @@ def _apply_rope(self, x: torch.Tensor, rotary_pos_emb: torch.Tensor, mscale: flo cu_seqlens=None, mscale=mscale, cp_group=self.pg_collection.cp, + # This flag is for the MLA-style interleaving in RoPE. + # Set it to False, as indexer does not apply interleaved RoPE. + mla_rotary_interleaved=False, ) # [seqlen, batch, *, index_head_dim] x = torch.cat([x_pe, x_nope], dim=-1) diff --git a/megatron/core/transformer/multi_latent_attention.py b/megatron/core/transformer/multi_latent_attention.py index fcafccff246..8023f53056e 100644 --- a/megatron/core/transformer/multi_latent_attention.py +++ b/megatron/core/transformer/multi_latent_attention.py @@ -930,6 +930,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( @@ -939,6 +940,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, (qk_head_dim + v_head_dim)] diff --git a/tests/unit_tests/fusions/test_mla_yarn_rope_apply.py b/tests/unit_tests/fusions/test_mla_yarn_rope_apply.py index d644a8ccf2a..abd4af1f3a8 100644 --- a/tests/unit_tests/fusions/test_mla_yarn_rope_apply.py +++ b/tests/unit_tests/fusions/test_mla_yarn_rope_apply.py @@ -91,7 +91,13 @@ def _test_fused_apply_mla_rope_for_q(input_format): no_pe, pe = torch.split(pytorch_fwd_input, [q_dim, emb_dim], dim=-1) pe_output = apply_rotary_pos_emb( - pe, freqs, transformer_config, cu_seqlens=cu_seqlens, mscale=mscale, cp_group=FakeCPGroup() + pe, + freqs, + transformer_config, + cu_seqlens=cu_seqlens, + mscale=mscale, + cp_group=FakeCPGroup(), + mla_rotary_interleaved=True, ) pytorch_output = torch.concat([no_pe, pe_output], dim=-1) pytorch_output.backward(pytorch_bwd_input, retain_graph=True) @@ -190,6 +196,7 @@ def _test_fused_apply_mla_rope_for_kv(input_format): cu_seqlens=cu_seqlens, mscale=mscale, cp_group=FakeCPGroup(), + mla_rotary_interleaved=True, ) if input_format == "sbhd": pe_output = pe_output.expand(-1, -1, num_heads, -1) From 835724d2b1b3a75187205c7a712bca0f6256cfd4 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Tue, 3 Mar 2026 23:08:17 -0800 Subject: [PATCH 07/18] add a UT for apply rope refactor --- .../fusions/test_mla_yarn_rope_apply.py | 62 +++++++++++++++++++ 1 file changed, 62 insertions(+) diff --git a/tests/unit_tests/fusions/test_mla_yarn_rope_apply.py b/tests/unit_tests/fusions/test_mla_yarn_rope_apply.py index abd4af1f3a8..9059d0157aa 100644 --- a/tests/unit_tests/fusions/test_mla_yarn_rope_apply.py +++ b/tests/unit_tests/fusions/test_mla_yarn_rope_apply.py @@ -1,12 +1,18 @@ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +import warnings +from unittest.mock import MagicMock, patch + import pytest import torch from megatron.core.models.common.embeddings import apply_rotary_pos_emb +from megatron.core.models.common.embeddings import rope_utils as rope_utils_module from megatron.core.models.common.embeddings.yarn_rotary_pos_embedding import YarnRotaryEmbedding +from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import is_torch_min_version +from tests.unit_tests.test_utilities import Utils try: from megatron.core.fusions.fused_mla_yarn_rope_apply import ( @@ -261,3 +267,59 @@ def test_forward_backward_for_q(self, input_format): def test_forward_backward_for_kv(self, input_format): _test_fused_apply_mla_rope_for_kv(input_format) + + +class TestApplyRotaryPosEmbMlaFusionConflict: + """Test apply_rotary_pos_emb: mla_rotary_interleaved vs apply_rope_fusion conflict.""" + + def setup_method(self): + Utils.initialize_model_parallel(1, 1) + model_parallel_cuda_manual_seed(123) + self.seq_len = 16 + self.num_heads = 2 + self.kv_channels = 32 + self.rot_dim = self.kv_channels + + def teardown_method(self): + Utils.destroy_model_parallel() + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + def test_mla_rotary_interleaved_with_apply_rope_fusion_emits_warning_and_uses_unfused(self): + """When apply_rope_fusion=True and mla_rotary_interleaved=True, expect warning and unfused path.""" + config = TransformerConfig( + num_attention_heads=self.num_heads, + num_layers=1, + apply_rope_fusion=True, + rotary_interleaved=False, + ) + t = torch.randn( + self.seq_len, 1, self.num_heads, self.kv_channels, device="cuda", dtype=torch.float32 + ) + freqs = torch.randn(self.seq_len, 1, 1, self.rot_dim, device="cuda", dtype=torch.float32) + + fused_mock = MagicMock(return_value=t.clone()) + with ( + patch.object(rope_utils_module, "fused_apply_rotary_pos_emb", fused_mock), + patch.object( + rope_utils_module, + "_apply_rotary_pos_emb_bshd", + wraps=rope_utils_module._apply_rotary_pos_emb_bshd, + ) as unfused_spy, + ): + with warnings.catch_warnings(record=True) as w: + warnings.simplefilter("always") + out = apply_rotary_pos_emb(t, freqs, config, mla_rotary_interleaved=True) + # Should have warned about MLA + fusion conflict + mla_fusion_warnings = [ + x for x in w if "apply_rope_fusion does not support MLA-style" in str(x.message) + ] + assert ( + len(mla_fusion_warnings) >= 1 + ), "Expected warning when mla_rotary_interleaved and apply_rope_fusion both enabled" + # Fused kernel must not be used + fused_mock.assert_not_called() + # Unfused path must have been used + unfused_spy.assert_called_once() + call_kw = unfused_spy.call_args[1] + assert call_kw["mla_rotary_interleaved"] is True + assert out.shape == t.shape From 19e2b9b08d83b48c3133219fe53e01f0f794066e Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Tue, 3 Mar 2026 23:07:49 -0800 Subject: [PATCH 08/18] add a UT for exp spec --- ...rimental_attention_variant_module_specs.py | 628 ++++++++++++++++++ 1 file changed, 628 insertions(+) create mode 100644 tests/unit_tests/models/test_experimental_attention_variant_module_specs.py diff --git a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py new file mode 100644 index 00000000000..8dcb88a3b31 --- /dev/null +++ b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py @@ -0,0 +1,628 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +from unittest.mock import MagicMock, patch + +import pytest + +from megatron.core.transformer.enums import AttnMaskType +from megatron.core.transformer.identity_op import IdentityOp +from megatron.core.transformer.spec_utils import ModuleSpec +from megatron.core.transformer.transformer_block import TransformerBlockSubmodules +from megatron.core.transformer.transformer_layer import TransformerLayer + +# --------------------------------------------------------------------------- +# Helpers: fake backend and config builders +# --------------------------------------------------------------------------- + + +class _FakeLinear: + pass + + +class _FakeColumnParallelLinear: + pass + + +class _FakeRowParallelLinear: + pass + + +class _FakeLayerNormColumnParallelLinear: + pass + + +class _FakeLayerNorm: + pass + + +class _FakeQKNorm: + pass + + +class _FakeCoreAttention: + pass + + +def _make_backend(fuse_layernorm=True): + """Return a mock BackendSpecProvider with deterministic return values.""" + backend = MagicMock() + backend.linear.return_value = _FakeLinear + backend.column_parallel_linear.return_value = _FakeColumnParallelLinear + backend.row_parallel_linear.return_value = _FakeRowParallelLinear + backend.column_parallel_layer_norm_linear.return_value = _FakeLayerNormColumnParallelLinear + backend.fuse_layernorm_and_linear.return_value = fuse_layernorm + backend.core_attention.return_value = _FakeCoreAttention + + def _layer_norm(rms_norm=False, for_qk=False): + return _FakeQKNorm if for_qk else _FakeLayerNorm + + backend.layer_norm.side_effect = _layer_norm + return backend + + +def _make_config(**overrides): + """Return a mock TransformerConfig with sane defaults.""" + defaults = dict( + num_layers=4, + normalization="RMSNorm", + qk_layernorm=False, + multi_latent_attention=False, + qk_l2_norm=False, + transformer_impl="transformer_engine", + use_kitchen=False, + experimental_attention_variant=None, + linear_attention_freq=None, + moe_layer_freq=1, + num_moe_experts=None, + moe_grouped_gemm=False, + moe_use_legacy_grouped_gemm=False, + use_te_activation_func=False, + pipeline_model_parallel_size=1, + pipeline_model_parallel_layout=None, + use_kitchen_attention=False, + kitchen_attention_backend="sdpa", + fallback_to_eager_attn=False, + ) + defaults.update(overrides) + cfg = MagicMock() + for k, v in defaults.items(): + setattr(cfg, k, v) + return cfg + + +# =================================================================== +# Tests for is_linear_attention_variant +# =================================================================== + + +class TestIsLinearAttentionVariant: + @staticmethod + def _fn(variant): + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + is_linear_attention_variant, + ) + + return is_linear_attention_variant(variant) + + @pytest.mark.parametrize( + "variant, expected", + [("gated_delta_net", True), ("dsa", False), (None, False), ("some_unknown_variant", False)], + ) + def test_variants(self, variant, expected): + """Validate linear-attention variant classification across supported and unsupported names.""" + assert self._fn(variant) is expected + + +# =================================================================== +# Tests for get_moe_layer_pattern +# =================================================================== + + +class TestGetMoeLayerPattern: + @staticmethod + def _fn(config): + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_moe_layer_pattern, + ) + + return get_moe_layer_pattern(config) + + @pytest.mark.parametrize( + "num_layers, freq, expected", + [(4, 1, [1, 1, 1, 1]), (6, 2, [1, 0, 1, 0, 1, 0]), (6, 3, [1, 0, 0, 1, 0, 0])], + ) + def test_int_freq(self, num_layers, freq, expected): + """Verify integer moe_layer_freq is expanded into the expected per-layer MoE pattern.""" + cfg = _make_config(num_layers=num_layers, moe_layer_freq=freq) + assert self._fn(cfg) == expected + + def test_list_freq(self): + """Verify an explicit list pattern is used as-is.""" + pattern = [1, 0, 1, 0] + cfg = _make_config(num_layers=4, moe_layer_freq=pattern) + assert self._fn(cfg) == pattern + + def test_list_freq_wrong_length_raises(self): + """Verify a list with mismatched length fails fast.""" + cfg = _make_config(num_layers=4, moe_layer_freq=[1, 0]) + with pytest.raises(AssertionError, match="Invalid length"): + self._fn(cfg) + + def test_invalid_type_raises(self): + """Verify unsupported moe_layer_freq types raise ValueError.""" + cfg = _make_config(num_layers=4, moe_layer_freq="bad") + with pytest.raises(ValueError, match="Invalid moe_layer_freq"): + self._fn(cfg) + + +# =================================================================== +# Tests for get_linear_attention_pattern +# =================================================================== + + +class TestGetLinearAttentionPattern: + @staticmethod + def _fn(config): + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_linear_attention_pattern, + ) + + return get_linear_attention_pattern(config) + + @pytest.mark.parametrize( + "num_layers, freq, expected", + [ + # Every 4th layer (1-indexed) is SDPA (0), the rest are LA (1) + (8, 4, [1, 1, 1, 0, 1, 1, 1, 0]), + (4, 2, [1, 0, 1, 0]), + (3, 1, [0, 0, 0]), + ], + ) + def test_int_freq(self, num_layers, freq, expected): + """Verify integer linear_attention_freq is expanded into the expected LA/SDPA pattern.""" + cfg = _make_config(num_layers=num_layers, linear_attention_freq=freq) + assert self._fn(cfg) == expected + + def test_list_freq(self): + """Verify an explicit linear-attention pattern list is used directly.""" + pattern = [1, 0, 1, 0] + cfg = _make_config(num_layers=4, linear_attention_freq=pattern) + assert self._fn(cfg) == pattern + + def test_list_freq_wrong_length_raises(self): + """Verify list length validation for linear_attention_freq.""" + cfg = _make_config(num_layers=4, linear_attention_freq=[1, 0, 1]) + with pytest.raises(AssertionError, match="Invalid length"): + self._fn(cfg) + + def test_none_for_non_linear_variant(self): + """Verify non-linear variants default to all-standard attention when freq is None.""" + cfg = _make_config( + num_layers=4, linear_attention_freq=None, experimental_attention_variant="dsa" + ) + assert self._fn(cfg) == [0, 0, 0, 0] + + def test_none_for_linear_variant_raises(self): + """Verify linear variants require linear_attention_freq to be explicitly set.""" + cfg = _make_config( + num_layers=4, + linear_attention_freq=None, + experimental_attention_variant="gated_delta_net", + ) + with pytest.raises(ValueError, match="linear_attention_freq is None"): + self._fn(cfg) + + def test_invalid_type_raises(self): + """Verify unsupported linear_attention_freq types raise ValueError.""" + cfg = _make_config(num_layers=4, linear_attention_freq=3.14) + with pytest.raises(ValueError, match="Invalid linear_attention_freq"): + self._fn(cfg) + + +# =================================================================== +# Tests for get_gated_delta_net_module_spec +# =================================================================== + + +class TestGetGatedDeltaNetModuleSpec: + def test_returns_correct_module_spec(self): + """Verify the top-level module spec targets GatedDeltaNet with expected metainfo.""" + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_gated_delta_net_module_spec, + ) + from megatron.core.ssm.gated_delta_net import GatedDeltaNet + + backend = _make_backend() + cfg = _make_config(normalization="RMSNorm") + spec = get_gated_delta_net_module_spec(cfg, backend=backend) + + assert isinstance(spec, ModuleSpec) + assert spec.module is GatedDeltaNet + assert spec.metainfo == {"fuse_input_layernorm": True} + + def test_submodules_use_backend_modules(self): + """Verify backend-provided projection/norm modules are wired into submodules.""" + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_gated_delta_net_module_spec, + ) + + backend = _make_backend() + cfg = _make_config(normalization="RMSNorm") + spec = get_gated_delta_net_module_spec(cfg, backend=backend) + + subs = spec.submodules + assert subs.in_proj == _FakeLayerNormColumnParallelLinear + assert subs.out_proj == _FakeRowParallelLinear + backend.layer_norm.assert_any_call(rms_norm=True, for_qk=False) + + def test_layer_norm_normalization(self): + """Verify LayerNorm mode passes rms_norm=False to backend.layer_norm.""" + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_gated_delta_net_module_spec, + ) + + backend = _make_backend() + cfg = _make_config(normalization="LayerNorm") + get_gated_delta_net_module_spec(cfg, backend=backend) + backend.layer_norm.assert_any_call(rms_norm=False, for_qk=False) + + def test_backend_auto_resolved_when_none(self): + """Verify backend is auto-resolved when caller does not pass one.""" + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_gated_delta_net_module_spec, + ) + + cfg = _make_config(normalization="RMSNorm") + with patch( + "megatron.core.models.gpt.experimental_attention_variant_module_specs" + "._get_backend_spec_provider", + return_value=_make_backend(), + ): + spec = get_gated_delta_net_module_spec(cfg, backend=None) + assert isinstance(spec, ModuleSpec) + + +# =================================================================== +# Tests for get_dsa_module_spec_for_backend +# =================================================================== + + +class TestGetDsaModuleSpec: + def _call(self, cfg=None, backend=None): + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_dsa_module_spec_for_backend, + ) + + if cfg is None: + cfg = _make_config(multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True) + if backend is None: + backend = _make_backend() + return get_dsa_module_spec_for_backend(cfg, backend=backend) + + def test_requires_multi_latent_attention(self): + """Verify DSA path rejects configs without MLA enabled.""" + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_dsa_module_spec_for_backend, + ) + + cfg = _make_config(multi_latent_attention=False, qk_l2_norm=False) + with pytest.raises(AssertionError, match="only MLA supports"): + get_dsa_module_spec_for_backend(cfg, backend=_make_backend()) + + def test_rejects_qk_l2_norm(self): + """Verify unsupported qk_l2_norm setting is rejected for DSA+MLA.""" + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_dsa_module_spec_for_backend, + ) + + cfg = _make_config(multi_latent_attention=True, qk_l2_norm=True) + with pytest.raises(AssertionError, match="qk_l2_norm is not supported"): + get_dsa_module_spec_for_backend(cfg, backend=_make_backend()) + + def test_returns_mla_self_attention_spec(self): + """Verify the returned attention module is MLA self-attention with causal mask.""" + from megatron.core.transformer.multi_latent_attention import MLASelfAttention + + spec = self._call() + assert spec.module is MLASelfAttention + assert spec.params == {"attn_mask_type": AttnMaskType.causal} + assert spec.metainfo == {"fuse_input_layernorm": False} + + def test_core_attention_is_dsa(self): + """Verify MLA core_attention is wrapped with DSAttention.""" + from megatron.core.transformer.experimental_attention_variant.dsa import DSAttention + + spec = self._call() + core = spec.submodules.core_attention + assert core.module is DSAttention + + def test_dsa_indexer_structure(self): + """Verify DSA indexer wiring uses expected backend linear/norm modules.""" + from megatron.core.transformer.experimental_attention_variant.dsa import DSAIndexer + + spec = self._call() + indexer = spec.submodules.core_attention.submodules.indexer + assert indexer.module is DSAIndexer + subs = indexer.submodules + assert subs.linear_wq_b == _FakeLinear + assert subs.linear_wk == _FakeLinear + assert subs.k_norm == _FakeQKNorm + assert subs.linear_weights_proj == _FakeLinear + + def test_qk_layernorm_enabled(self): + """Verify q/kv layernorm submodules are enabled when qk_layernorm=True.""" + cfg = _make_config( + multi_latent_attention=True, + qk_l2_norm=False, + qk_layernorm=True, + normalization="RMSNorm", + ) + spec = self._call(cfg=cfg) + assert spec.submodules.q_layernorm == _FakeQKNorm + assert spec.submodules.kv_layernorm == _FakeQKNorm + + def test_qk_layernorm_disabled(self): + """Verify q/kv layernorm submodules become IdentityOp when qk_layernorm=False.""" + cfg = _make_config(multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=False) + spec = self._call(cfg=cfg) + assert spec.submodules.q_layernorm is IdentityOp + assert spec.submodules.kv_layernorm is IdentityOp + + def test_linear_projections(self): + """Verify all major Q/KV projection slots map to the expected backend modules.""" + spec = self._call() + subs = spec.submodules + assert subs.linear_q_proj == _FakeColumnParallelLinear + assert subs.linear_q_down_proj == _FakeLinear + assert subs.linear_q_up_proj == _FakeColumnParallelLinear + assert subs.linear_kv_down_proj == _FakeLinear + assert subs.linear_kv_up_proj == _FakeColumnParallelLinear + assert subs.linear_proj == _FakeRowParallelLinear + + +# =================================================================== +# Tests for get_experimental_attention_variant_module_spec +# =================================================================== + + +class TestGetExperimentalAttentionVariantModuleSpec: + MODULE = "megatron.core.models.gpt.experimental_attention_variant_module_specs" + + @pytest.mark.parametrize( + "variant, target_fn", + [ + ("gated_delta_net", "get_gated_delta_net_module_spec"), + ("dsa", "get_dsa_module_spec_for_backend"), + ], + ) + def test_dispatches_to_variant_handler(self, variant, target_fn): + """Verify dispatcher routes each variant name to its corresponding builder function.""" + backend = _make_backend() + cfg = _make_config(experimental_attention_variant=variant, normalization="RMSNorm") + with patch(f"{self.MODULE}.{target_fn}") as mock_fn: + mock_fn.return_value = ModuleSpec(module=MagicMock) + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_experimental_attention_variant_module_spec, + ) + + result = get_experimental_attention_variant_module_spec(cfg, backend=backend) + mock_fn.assert_called_once_with(config=cfg, backend=backend) + assert result is mock_fn.return_value + + def test_invalid_variant_raises(self): + """Verify unknown variant names raise a clear ValueError.""" + cfg = _make_config(experimental_attention_variant="unknown") + with pytest.raises(ValueError, match="Invalid experimental attention variant"): + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_experimental_attention_variant_module_spec, + ) + + get_experimental_attention_variant_module_spec(cfg, backend=_make_backend()) + + +# =================================================================== +# Tests for get_transformer_layer_with_experimental_attention_variant_spec +# =================================================================== + + +class TestGetTransformerLayerWithExperimentalAttentionVariantSpec: + MODULE = "megatron.core.models.gpt.experimental_attention_variant_module_specs" + + def _make_attention_spec(self, fuse_input_layernorm=True): + """Construct a mock attention spec with configurable fuse metadata.""" + return ModuleSpec(module=MagicMock, metainfo={"fuse_input_layernorm": fuse_input_layernorm}) + + def _make_mlp_spec(self, fuse_pre_mlp_layernorm=True): + """Construct a mock MLP spec with configurable fuse metadata.""" + return ModuleSpec( + module=MagicMock, metainfo={"fuse_pre_mlp_layernorm": fuse_pre_mlp_layernorm} + ) + + def test_all_experimental_no_moe(self): + """Verify all layers use experimental attention and dense MLP when no MoE is configured.""" + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_transformer_layer_with_experimental_attention_variant_spec, + ) + + cfg = _make_config( + num_layers=4, + experimental_attention_variant="dsa", + num_moe_experts=None, + normalization="RMSNorm", + ) + backend = _make_backend() + attn_spec = self._make_attention_spec(fuse_input_layernorm=False) + mlp_spec = self._make_mlp_spec(fuse_pre_mlp_layernorm=True) + + with ( + patch( + f"{self.MODULE}.get_experimental_attention_variant_module_spec", + return_value=attn_spec, + ), + patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=mlp_spec), + ): + specs = get_transformer_layer_with_experimental_attention_variant_spec( + cfg, backend=backend + ) + + assert len(specs) == 4 + for s in specs: + # Each layer should share the same selected module specs in this setup. + assert s.module is TransformerLayer + assert s.submodules.self_attention is attn_spec + assert s.submodules.mlp is mlp_spec + + def test_hybrid_attention_pattern(self): + """Verify attention alternates between experimental and standard specs per pattern.""" + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_transformer_layer_with_experimental_attention_variant_spec, + ) + + cfg = _make_config( + num_layers=4, + experimental_attention_variant="gated_delta_net", + linear_attention_freq=2, + num_moe_experts=None, + normalization="RMSNorm", + ) + backend = _make_backend() + exp_attn_spec = self._make_attention_spec(fuse_input_layernorm=True) + std_attn_spec = self._make_attention_spec(fuse_input_layernorm=False) + mlp_spec = self._make_mlp_spec(fuse_pre_mlp_layernorm=True) + + with ( + patch( + f"{self.MODULE}.get_experimental_attention_variant_module_spec", + return_value=exp_attn_spec, + ), + patch(f"{self.MODULE}._get_self_attention_module_spec", return_value=std_attn_spec), + patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=mlp_spec), + ): + specs = get_transformer_layer_with_experimental_attention_variant_spec( + cfg, backend=backend + ) + + assert len(specs) == 4 + # Pattern for linear_attention_freq=2: [1, 0, 1, 0] + assert specs[0].submodules.self_attention is exp_attn_spec + assert specs[1].submodules.self_attention is std_attn_spec + assert specs[2].submodules.self_attention is exp_attn_spec + assert specs[3].submodules.self_attention is std_attn_spec + + def test_hybrid_moe_pattern(self): + """Verify MLP alternates between MoE and dense specs per moe_layer_freq pattern.""" + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_transformer_layer_with_experimental_attention_variant_spec, + ) + + cfg = _make_config( + num_layers=4, + experimental_attention_variant="dsa", + num_moe_experts=8, + moe_layer_freq=2, + normalization="RMSNorm", + ) + backend = _make_backend() + attn_spec = self._make_attention_spec(fuse_input_layernorm=False) + moe_spec = self._make_mlp_spec(fuse_pre_mlp_layernorm=False) + dense_spec = self._make_mlp_spec(fuse_pre_mlp_layernorm=True) + + with ( + patch( + f"{self.MODULE}.get_experimental_attention_variant_module_spec", + return_value=attn_spec, + ), + patch(f"{self.MODULE}._get_moe_module_spec", return_value=moe_spec), + patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=dense_spec), + ): + specs = get_transformer_layer_with_experimental_attention_variant_spec( + cfg, backend=backend + ) + + # moe_layer_freq=2 -> [1, 0, 1, 0] + assert specs[0].submodules.mlp is moe_spec + assert specs[1].submodules.mlp is dense_spec + assert specs[2].submodules.mlp is moe_spec + assert specs[3].submodules.mlp is dense_spec + + +# =================================================================== +# Tests for get_transformer_block_with_experimental_attention_variant_spec +# =================================================================== + + +class TestGetTransformerBlockWithExperimentalAttentionVariantSpec: + MODULE = "megatron.core.models.gpt.experimental_attention_variant_module_specs" + + @pytest.mark.parametrize( + "pp_size,vp_stage,pp_rank,use_layout,offset,num_layers_to_build,layout_ids,expected_ids", + [ + # no pipeline split + (1, None, None, False, 0, 4, None, [0, 1, 2, 3]), + # pp split (rank 1 gets [4,5,6,7]) + (2, None, 1, False, 4, 4, None, [4, 5, 6, 7]), + # vpp + pp split (example stage) + (2, 1, 0, False, 2, 2, None, [2, 3]), + # explicit pipeline layout wins over offset/num_layers + (2, 0, 0, True, None, None, [0, 3, 5], [0, 3, 5]), + ], + ) + def test_get_transformer_block_with_experimental_attention_variant_spec( + self, + pp_size, + vp_stage, + pp_rank, + use_layout, + offset, + num_layers_to_build, + layout_ids, + expected_ids, + ): + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_transformer_block_with_experimental_attention_variant_spec, + ) + + """Verify transformer block layer slicing across pp/vpp/layout combinations.""" + # Keep layer pool just large enough for expected ids in each parameterized case. + num_layers = 8 if expected_ids and max(expected_ids) >= 4 else 4 + mock_layout = MagicMock() if use_layout else None + if mock_layout is not None: + # When layout is provided, it should fully control local layer selection. + mock_layout.get_layer_id_list.return_value = layout_ids + + cfg = _make_config( + num_layers=num_layers, + pipeline_model_parallel_size=pp_size, + pipeline_model_parallel_layout=mock_layout, + normalization="RMSNorm", + ) + backend = _make_backend() + fake_layer_specs = [ + ModuleSpec(module=TransformerLayer, submodules=MagicMock()) for _ in range(num_layers) + ] + + with ( + patch(f"{self.MODULE}._get_backend_spec_provider", return_value=backend), + patch( + f"{self.MODULE}.get_transformer_layer_with_experimental_attention_variant_spec", + return_value=fake_layer_specs, + ), + ): + if use_layout: + result = get_transformer_block_with_experimental_attention_variant_spec( + cfg, vp_stage=vp_stage, pp_rank=pp_rank + ) + else: + # Without explicit layout, slicing comes from offset + num_layers_to_build. + with ( + patch(f"{self.MODULE}.get_transformer_layer_offset", return_value=offset), + patch( + f"{self.MODULE}.get_num_layers_to_build", return_value=num_layers_to_build + ), + ): + result = get_transformer_block_with_experimental_attention_variant_spec( + cfg, vp_stage=vp_stage, pp_rank=pp_rank + ) + + assert isinstance(result, TransformerBlockSubmodules) + assert result.layer_specs == [fake_layer_specs[i] for i in expected_ids] From feef7493d3f803867f719124b1c4c67f8e01df2a Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Tue, 3 Mar 2026 23:25:23 -0800 Subject: [PATCH 09/18] update exp spec UT --- ...rimental_attention_variant_module_specs.py | 63 +++++++++++++------ 1 file changed, 44 insertions(+), 19 deletions(-) diff --git a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py index 8dcb88a3b31..af43a5e7aa4 100644 --- a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py +++ b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py @@ -4,7 +4,7 @@ import pytest -from megatron.core.transformer.enums import AttnMaskType +from megatron.core.transformer.enums import AttnMaskType, LayerType from megatron.core.transformer.identity_op import IdentityOp from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.transformer.transformer_block import TransformerBlockSubmodules @@ -349,28 +349,44 @@ def test_dsa_indexer_structure(self): assert subs.k_norm == _FakeQKNorm assert subs.linear_weights_proj == _FakeLinear - def test_qk_layernorm_enabled(self): - """Verify q/kv layernorm submodules are enabled when qk_layernorm=True.""" + @pytest.mark.parametrize("normalization", ["RMSNorm", "LayerNorm"]) + def test_qk_layernorm_enabled(self, normalization): + """Verify q/kv layernorm uses backend.layer_norm(rms_norm=..., for_qk=True).""" + backend = _make_backend() cfg = _make_config( multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True, - normalization="RMSNorm", + normalization=normalization, ) - spec = self._call(cfg=cfg) + spec = self._call(cfg=cfg, backend=backend) + expected_rms = normalization == "RMSNorm" assert spec.submodules.q_layernorm == _FakeQKNorm assert spec.submodules.kv_layernorm == _FakeQKNorm + # Both point to the same qk_norm object + assert spec.submodules.q_layernorm is spec.submodules.kv_layernorm + backend.layer_norm.assert_any_call(rms_norm=expected_rms, for_qk=True) def test_qk_layernorm_disabled(self): - """Verify q/kv layernorm submodules become IdentityOp when qk_layernorm=False.""" + """Verify q/kv layernorm becomes IdentityOp, skipping backend.layer_norm for qk.""" + backend = _make_backend() cfg = _make_config(multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=False) - spec = self._call(cfg=cfg) + spec = self._call(cfg=cfg, backend=backend) assert spec.submodules.q_layernorm is IdentityOp assert spec.submodules.kv_layernorm is IdentityOp + # backend.layer_norm is still called for the indexer k_norm (for_qk=True at line 94), + # but NOT for the outer qk_norm (line 105-107 takes the else branch). + # Exactly one for_qk=True call should exist (from the indexer, not from qk_norm). + qk_calls = [c for c in backend.layer_norm.call_args_list if c.kwargs.get("for_qk")] + assert ( + len(qk_calls) == 1 + ), f"Expected 1 for_qk=True call (indexer only), got {len(qk_calls)}" def test_linear_projections(self): - """Verify all major Q/KV projection slots map to the expected backend modules.""" - spec = self._call() + """Verify Q/KV projection slots and backend.column_parallel_linear call count.""" + backend = _make_backend() + cfg = _make_config(multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True) + spec = self._call(cfg=cfg, backend=backend) subs = spec.submodules assert subs.linear_q_proj == _FakeColumnParallelLinear assert subs.linear_q_down_proj == _FakeLinear @@ -378,6 +394,9 @@ def test_linear_projections(self): assert subs.linear_kv_down_proj == _FakeLinear assert subs.linear_kv_up_proj == _FakeColumnParallelLinear assert subs.linear_proj == _FakeRowParallelLinear + # column_parallel_linear() is called exactly 3 times (q_proj, q_up_proj, kv_up_proj) + assert backend.column_parallel_linear.call_count == 3 + assert backend.row_parallel_linear.call_count == 1 # =================================================================== @@ -555,20 +574,21 @@ class TestGetTransformerBlockWithExperimentalAttentionVariantSpec: MODULE = "megatron.core.models.gpt.experimental_attention_variant_module_specs" @pytest.mark.parametrize( - "pp_size,vp_stage,pp_rank,use_layout,offset,num_layers_to_build,layout_ids,expected_ids", + "num_layers,pp_size,vp_stage,pp_rank,use_layout,offset,num_layers_to_build,layout_ids,expected_ids", [ # no pipeline split - (1, None, None, False, 0, 4, None, [0, 1, 2, 3]), + (4, 1, None, None, False, 0, 4, None, [0, 1, 2, 3]), # pp split (rank 1 gets [4,5,6,7]) - (2, None, 1, False, 4, 4, None, [4, 5, 6, 7]), + (8, 2, None, 1, False, 4, 4, None, [4, 5, 6, 7]), # vpp + pp split (example stage) - (2, 1, 0, False, 2, 2, None, [2, 3]), + (8, 2, 1, 0, False, 2, 2, None, [2, 3]), # explicit pipeline layout wins over offset/num_layers - (2, 0, 0, True, None, None, [0, 3, 5], [0, 3, 5]), + (8, 2, 0, 0, True, None, None, [0, 3, 5], [0, 3, 5]), ], ) def test_get_transformer_block_with_experimental_attention_variant_spec( self, + num_layers, pp_size, vp_stage, pp_rank, @@ -578,13 +598,11 @@ def test_get_transformer_block_with_experimental_attention_variant_spec( layout_ids, expected_ids, ): + """Verify transformer block layer slicing and vp/pp argument forwarding.""" from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( get_transformer_block_with_experimental_attention_variant_spec, ) - """Verify transformer block layer slicing across pp/vpp/layout combinations.""" - # Keep layer pool just large enough for expected ids in each parameterized case. - num_layers = 8 if expected_ids and max(expected_ids) >= 4 else 4 mock_layout = MagicMock() if use_layout else None if mock_layout is not None: # When layout is provided, it should fully control local layer selection. @@ -612,17 +630,24 @@ def test_get_transformer_block_with_experimental_attention_variant_spec( result = get_transformer_block_with_experimental_attention_variant_spec( cfg, vp_stage=vp_stage, pp_rank=pp_rank ) + mock_layout.get_layer_id_list.assert_called_once_with( + layer_type=LayerType.decoder, vp_stage=vp_stage, pp_rank=pp_rank + ) else: # Without explicit layout, slicing comes from offset + num_layers_to_build. with ( - patch(f"{self.MODULE}.get_transformer_layer_offset", return_value=offset), + patch( + f"{self.MODULE}.get_transformer_layer_offset", return_value=offset + ) as mock_offset, patch( f"{self.MODULE}.get_num_layers_to_build", return_value=num_layers_to_build - ), + ) as mock_num_layers, ): result = get_transformer_block_with_experimental_attention_variant_spec( cfg, vp_stage=vp_stage, pp_rank=pp_rank ) + mock_offset.assert_called_once_with(cfg, vp_stage=vp_stage, pp_rank=pp_rank) + mock_num_layers.assert_called_once_with(cfg, vp_stage=vp_stage, pp_rank=pp_rank) assert isinstance(result, TransformerBlockSubmodules) assert result.layer_specs == [fake_layer_specs[i] for i in expected_ids] From 06f591375fba0b16e657210ed59fb3c830554db1 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Fri, 13 Feb 2026 00:03:15 -0800 Subject: [PATCH 10/18] add a functional test --- .../model_config.yaml | 65 +++++++++++++++++++ tests/test_utils/recipes/h100/gpt.yaml | 5 ++ 2 files changed, 70 insertions(+) create mode 100644 tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml new file mode 100644 index 00000000000..b54637b3b6e --- /dev/null +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml @@ -0,0 +1,65 @@ +ENV_VARS: + CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 + NCCL_ALGO: Ring + CUBLAS_WORKSPACE_CONFIG: :4096:8 + ENABLE_LIGHTWEIGHT_MODE: true +MODEL_ARGS: + --num-layers: 4 + --hidden-size: 512 + --num-attention-heads: 8 + --multi-latent-attention: true + --q-lora-rank: 192 + --kv-lora-rank: 64 + --qk-head-dim: 16 + --qk-pos-emb-head-dim: 8 + --v-head-dim: 16 + --experimental-attention-variant: dsa + --dsa-indexer-n-heads: 64 + --dsa-indexer-head-dim: 128 + --dsa-indexer-topk: 2048 + --attention-backend: fused + --log-params-norm: true + --log-num-zeros-in-grad: true + --log-validation-ppl-to-tensorboard: true + --log-timers-to-tensorboard: true + --tensorboard-dir: ${TENSORBOARD_PATH} + --micro-batch-size: 4 + --global-batch-size: 32 + --seq-length: 1024 + --max-position-embeddings: 1024 + --train-iters: 50 + --timing-log-level: 0 + --lr-decay-iters: 320000 + --save: ${CHECKPOINT_SAVE_PATH} + --load: ${CHECKPOINT_LOAD_PATH} + --data-path: ${DATA_PATH}/text/the_pile/shard00/my-gpt3_00_text_document + --vocab-file: ${DATA_PATH}/text/the_pile/shard00/bpe/vocab.json + --merge-file: ${DATA_PATH}/text/the_pile/shard00/bpe/merges.txt + --split: 949,50,1 + --distributed-backend: nccl + --lr: 0.00015 + --lr-decay-style: cosine + --min-lr: 1.0e-5 + --weight-decay: 1e-2 + --clip-grad: 1.0 + --lr-warmup-fraction: .01 + --log-interval: 1 + --save-interval: 25 + --eval-interval: 1000 + --eval-iters: 10 + --transformer-impl: transformer_engine + --tensor-model-parallel-size: 2 + --pipeline-model-parallel-size: 2 + --sequence-parallel: true + --untie-embeddings-and-output-weights: true + --deterministic-mode: true + --no-gradient-accumulation-fusion: true + --attention-softmax-in-fp32: true + --use-mcore-models: true + --ckpt-format: torch_dist + --data-cache-path: ${DATA_CACHE_PATH} + --bf16: true + --attention-backend: unfused + --log-memory-to-tensorboard: true +TEST_TYPE: ckpt-resume diff --git a/tests/test_utils/recipes/h100/gpt.yaml b/tests/test_utils/recipes/h100/gpt.yaml index a37054ce015..98ad9e318a2 100644 --- a/tests/test_utils/recipes/h100/gpt.yaml +++ b/tests/test_utils/recipes/h100/gpt.yaml @@ -369,6 +369,11 @@ products: - environment: [dev] scope: [mr, mr-github] platforms: [dgx_h100] + - test_case: [gpt3_mcore_te_tp2_pp2_dsa] + products: + - environment: [dev] + scope: [mr, mr-github, mr-github-slim] + platforms: [dgx_h100] - test_case: [gpt3_mcore_te_tp2_pp2_resume_torch_dist_ddp_average_in_collective] products: - environment: [dev] From 7aed8304228441ff8c3236d5999f2c76a36dd837 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Fri, 13 Feb 2026 00:39:35 -0800 Subject: [PATCH 11/18] fix --- .../test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml index b54637b3b6e..63a0933313c 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml @@ -18,6 +18,7 @@ MODEL_ARGS: --dsa-indexer-n-heads: 64 --dsa-indexer-head-dim: 128 --dsa-indexer-topk: 2048 + --dsa-indexer-loss-coeff: 0.01 --attention-backend: fused --log-params-norm: true --log-num-zeros-in-grad: true From 2fb2d1af091c06525c7b3214926ad1ff51b5be3c Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Wed, 4 Mar 2026 01:26:49 -0800 Subject: [PATCH 12/18] update dependency --- pyproject.toml | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index 77ea81bf124..b8f45b6f55b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -110,6 +110,7 @@ dev = [ "quart", "openai[aiohttp]", "orjson", + "fast-hadamard-transform", ] ### 'lts' extra is deprecated ### @@ -180,6 +181,7 @@ no-build-isolation-package = [ "mamba-ssm", "transformer-engine", "transformer-engine-torch", + "fast-hadamard-transform", ] link-mode = "copy" # We don't define override-dependencies globally but rather locally where we need it. @@ -191,6 +193,30 @@ override-dependencies = [ "triton; sys_platform == 'never'", ] +[[tool.uv.dependency-metadata]] +name = "flash-mla" +version = "1.0.0+9edee0c" +requires-dist = [] + +[[tool.uv.dependency-metadata]] +name = "transformer-engine" +version = "2.15.0+42b84005" +requires-dist = [ + "pydantic", + "importlib-metadata>=1.0", + "packaging", + "torch>=2.1", + "einops", + "onnxscript", + "onnx", + "nvdlfw-inspect", +] + +[[tool.uv.dependency-metadata]] +name = "fast-hadamard-transform" +version = "1.0.4.post1" +requires-dist = ["torch", "packaging", "ninja"] + [tool.uv.sources] flash_mla = [ @@ -199,6 +225,7 @@ flash_mla = [ transformer-engine = { git = "https://github.com/NVIDIA/TransformerEngine.git", rev = "f8a01cd5ac55a6b669bacc4b222242bfff2c822a" } nemo-run = { git = "https://github.com/NVIDIA-NeMo/Run.git", rev = "17ae86b64d7f75653351664f5d8c9e466faede00" } emerging_optimizers = { git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git", rev = "v0.2.0" } +fast-hadamard-transform = { git = "https://github.com/Dao-AILab/fast-hadamard-transform.git", rev = "f134af63deb2df17e1171a9ec1ea4a7d8604d5ca" } [tool.isort] profile = "black" # black-compatible From c72087f600d9bdbabbe0e69f51690eca0b1640ce Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Tue, 7 Apr 2026 21:35:33 -0700 Subject: [PATCH 13/18] fix absorbed mla --- .../transformer/experimental_attention_variant/absorbed_mla.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py index 00a725378f0..48e6c76ea2f 100644 --- a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py +++ b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py @@ -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( @@ -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)] From 082c0bb14b69d0651b72b97b9dc32c6f6c36b329 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Thu, 28 May 2026 06:09:50 -0700 Subject: [PATCH 14/18] test: fix experimental attention spec mocks --- .../test_experimental_attention_variant_module_specs.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py index af43a5e7aa4..7cc406a198a 100644 --- a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py +++ b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py @@ -478,7 +478,7 @@ def test_all_experimental_no_moe(self): f"{self.MODULE}.get_experimental_attention_variant_module_spec", return_value=attn_spec, ), - patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=mlp_spec), + patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=(mlp_spec, True)), ): specs = get_transformer_layer_with_experimental_attention_variant_spec( cfg, backend=backend @@ -515,7 +515,7 @@ def test_hybrid_attention_pattern(self): return_value=exp_attn_spec, ), patch(f"{self.MODULE}._get_self_attention_module_spec", return_value=std_attn_spec), - patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=mlp_spec), + patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=(mlp_spec, True)), ): specs = get_transformer_layer_with_experimental_attention_variant_spec( cfg, backend=backend @@ -551,8 +551,8 @@ def test_hybrid_moe_pattern(self): f"{self.MODULE}.get_experimental_attention_variant_module_spec", return_value=attn_spec, ), - patch(f"{self.MODULE}._get_moe_module_spec", return_value=moe_spec), - patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=dense_spec), + patch(f"{self.MODULE}._get_moe_module_spec", return_value=(moe_spec, False)), + patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=(dense_spec, True)), ): specs = get_transformer_layer_with_experimental_attention_variant_spec( cfg, backend=backend From 8eb22b723a94cff466a7d032efeb81acd697f42d Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Thu, 28 May 2026 18:42:39 -0700 Subject: [PATCH 15/18] build: refresh uv lock --- uv.lock | 42 ++++++++++++++++++++++++++++++++++-------- 1 file changed, 34 insertions(+), 8 deletions(-) diff --git a/uv.lock b/uv.lock index 43804a624ac..2e702f82e0f 100644 --- a/uv.lock +++ b/uv.lock @@ -1,5 +1,5 @@ version = 1 -revision = 2 +revision = 3 requires-python = ">=3.12" resolution-markers = [ "python_full_version >= '3.14' and platform_machine != 's390x' and sys_platform == 'win32'", @@ -29,6 +29,20 @@ overrides = [ { name = "triton", marker = "sys_platform == 'never'" }, ] +[[manifest.dependency-metadata]] +name = "fast-hadamard-transform" +version = "1.0.4.post1" +requires-dist = ["torch", "packaging", "ninja"] + +[[manifest.dependency-metadata]] +name = "flash-mla" +version = "1.0.0+9edee0c" + +[[manifest.dependency-metadata]] +name = "transformer-engine" +version = "2.15.0+42b84005" +requires-dist = ["pydantic", "importlib-metadata>=1.0", "packaging", "torch>=2.1", "einops", "onnxscript", "onnx", "nvdlfw-inspect"] + [[package]] name = "absl-py" version = "2.4.0" @@ -1212,6 +1226,16 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/37/f9/f8497ef8b873a8bb2a750ee2a6c5f0fc22258e1acb6245fd237042a6c279/fabric-3.2.3-py3-none-any.whl", hash = "sha256:ce61917f4f398018337ce279b357650a3a74baecf3fdd53a5839013944af965e", size = 59502, upload-time = "2026-04-06T00:00:10.176Z" }, ] +[[package]] +name = "fast-hadamard-transform" +version = "1.0.4.post1" +source = { git = "https://github.com/Dao-AILab/fast-hadamard-transform.git?rev=f134af63deb2df17e1171a9ec1ea4a7d8604d5ca#f134af63deb2df17e1171a9ec1ea4a7d8604d5ca" } +dependencies = [ + { name = "ninja" }, + { name = "packaging" }, + { name = "torch", marker = "sys_platform == 'never'" }, +] + [[package]] name = "fastapi" version = "0.136.3" @@ -2261,6 +2285,7 @@ dev = [ { name = "datasets" }, { name = "einops" }, { name = "emerging-optimizers" }, + { name = "fast-hadamard-transform" }, { name = "fastapi" }, { name = "flash-linear-attention" }, { name = "flashinfer-python" }, @@ -2364,6 +2389,7 @@ requires-dist = [ { name = "datasets", marker = "extra == 'dev'" }, { name = "einops", marker = "extra == 'dev'", specifier = "~=0.8" }, { name = "emerging-optimizers", marker = "extra == 'dev'", git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git?rev=v0.2.0" }, + { name = "fast-hadamard-transform", marker = "extra == 'dev'", git = "https://github.com/Dao-AILab/fast-hadamard-transform.git?rev=f134af63deb2df17e1171a9ec1ea4a7d8604d5ca" }, { name = "fastapi", marker = "extra == 'dev'", specifier = "~=0.50" }, { name = "flash-linear-attention", marker = "extra == 'dev'", specifier = "~=0.4.0" }, { name = "flashinfer-python", marker = "extra == 'dev'", specifier = ">=0.5.0,<0.7.0" }, @@ -5276,14 +5302,14 @@ version = "2.12.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "cuda-bindings", marker = "sys_platform == 'linux'" }, - { name = "filelock" }, - { name = "fsspec" }, - { name = "jinja2" }, - { name = "networkx" }, - { name = "setuptools" }, - { name = "sympy" }, + { name = "filelock", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "fsspec", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "jinja2", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "networkx", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "setuptools", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "sympy", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, { name = "triton", marker = "sys_platform == 'never'" }, - { name = "typing-extensions" }, + { name = "typing-extensions", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, ] [[package]] From c0f93cfb90bdd17ea8a9ec6e9bb0fbe3aa792dde Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Sun, 31 May 2026 18:03:12 -0700 Subject: [PATCH 16/18] test: stabilize weighted squared relu fusion --- tests/unit_tests/fusions/test_weighted_squared_relu_fusion.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/unit_tests/fusions/test_weighted_squared_relu_fusion.py b/tests/unit_tests/fusions/test_weighted_squared_relu_fusion.py index 85755ac1de7..6de1955e55d 100644 --- a/tests/unit_tests/fusions/test_weighted_squared_relu_fusion.py +++ b/tests/unit_tests/fusions/test_weighted_squared_relu_fusion.py @@ -11,6 +11,8 @@ @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @pytest.mark.parametrize("input_dtype", [torch.bfloat16, torch.float32]) def test_weighted_squared_relu_fusion(input_dtype): + torch.manual_seed(0) + # Tolerances depend on dtype precision if input_dtype == torch.float32: tols = dict(rtol=1.0e-6, atol=1.0e-6) From a5ea4ea5170ac210a0e84b5f4165f9b748ef9c81 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Mon, 1 Jun 2026 18:02:05 -0700 Subject: [PATCH 17/18] fix: accept stage args in decoder layer specs --- megatron/core/models/gpt/gpt_layer_specs.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py index 28107556b21..7711a92bda6 100755 --- a/megatron/core/models/gpt/gpt_layer_specs.py +++ b/megatron/core/models/gpt/gpt_layer_specs.py @@ -565,8 +565,11 @@ def get_gpt_decoder_layer_specs( use_transformer_engine: bool, normalization: Optional[str] = None, qk_l2_norm: Optional[bool] = False, + vp_stage: Optional[int] = None, + pp_rank: Optional[int] = None, ) -> TransformerBlockSubmodules: """GPT block spec.""" + del vp_stage, pp_rank # Accepted for API compatibility with stage-aware callers. 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=}." From 84771311b3b87d6cb0a723a8e865439825a04839 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Tue, 2 Jun 2026 09:13:03 +0800 Subject: [PATCH 18/18] Update gpt_layer_specs.py --- megatron/core/models/gpt/gpt_layer_specs.py | 1 - 1 file changed, 1 deletion(-) diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py index 7711a92bda6..984840b3a87 100755 --- a/megatron/core/models/gpt/gpt_layer_specs.py +++ b/megatron/core/models/gpt/gpt_layer_specs.py @@ -569,7 +569,6 @@ def get_gpt_decoder_layer_specs( pp_rank: Optional[int] = None, ) -> TransformerBlockSubmodules: """GPT block spec.""" - del vp_stage, pp_rank # Accepted for API compatibility with stage-aware callers. 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=}."