diff --git a/megatron/core/models/common/embeddings/rotary_pos_embedding.py b/megatron/core/models/common/embeddings/rotary_pos_embedding.py index 0e560f939f2..804bdb7c537 100644 --- a/megatron/core/models/common/embeddings/rotary_pos_embedding.py +++ b/megatron/core/models/common/embeddings/rotary_pos_embedding.py @@ -280,6 +280,9 @@ class MultimodalRotaryEmbedding(nn.Module): for longer sequences. The value must be a float larger than 1.0. Defaults to None rotary_base (int, optional): Base period for rotary position embeddings. Defaults to 10000. + interleaved_mrope (bool, optional): If True, use the interleaved T/H/W MRoPE layout + (Qwen3.5-VL style). If False (default), use the original section-based layout + (Qwen2-VL style). """ def __init__( @@ -290,6 +293,7 @@ def __init__( seq_len_interpolation_factor: Optional[float] = None, rotary_base: int = 10000, cp_group: Optional[torch.distributed.ProcessGroup] = None, + interleaved_mrope: bool = False, ) -> None: super().__init__() @@ -297,6 +301,7 @@ def __init__( if rotary_percent < 1.0: dim = int(dim * rotary_percent) self.rotary_interleaved = rotary_interleaved + self.interleaved_mrope = interleaved_mrope self.seq_len_interpolation_factor = seq_len_interpolation_factor self.inv_freq = 1.0 / ( @@ -312,6 +317,25 @@ def __init__( else parallel_state.get_context_parallel_group(check_initialized=False) ) + @staticmethod + def _apply_interleaved_mrope(freqs: Tensor, mrope_section: List[int]) -> Tensor: + """Merge T/H/W frequency channels into a single interleaved vector. + + Converts from the per-channel outer-product layout ``(3, bs, seq_len, dim)`` + to the interleaved layout ``(bs, seq_len, dim)`` used by HF + ``Qwen3VLTextRotaryEmbedding.apply_interleaved_mrope`` (unified 2026-02-24) + and Megatron-Bridge ``Qwen3VLMultimodalRotaryEmbedding``. + + H freqs occupy stride-3 positions ``{1, 4, 7, ...}`` and W freqs occupy + ``{2, 5, 8, ...}``, while T freqs remain at ``{0, 3, 6, ...}``. + """ + freqs_out = freqs[0].clone() # start with T channel: shape (bs, seq_len, dim) + for dim_idx, offset in enumerate((1, 2), start=1): # H then W + length = mrope_section[dim_idx] * 3 + idx = slice(offset, length, 3) + freqs_out[..., idx] = freqs[dim_idx, ..., idx] + return freqs_out + def forward( self, position_ids: torch.Tensor, @@ -341,20 +365,34 @@ def forward( seq_expanded = seq[:, :, None, :].float() # shape (3, bs, seq_length, dim) freqs = (inv_freq_expanded @ seq_expanded).transpose(2, 3) + # first part even vector components, second part odd vector components, # 2 * dim in dimension size - if not self.rotary_interleaved: - emb = torch.cat((freqs, freqs), dim=-1) # shape (3, bs, seq_length, 2 * dim) + if self.interleaved_mrope: + # Qwen3.5-VL: merge T/H/W with interleaved layout [T₀,H₀,W₀,T₁,H₁,W₁,...]. + # freqs becomes shape (bs, seq_length, dim). + freqs = self._apply_interleaved_mrope(freqs, mrope_section) + if not self.rotary_interleaved: + emb = torch.cat((freqs, freqs), dim=-1) # shape (bs, seq_length, 2 * dim) + else: + bs = freqs.shape[0] + emb = torch.stack((freqs.view(bs, -1, 1), freqs.view(bs, -1, 1)), dim=-1).view( + bs, freqs.shape[1], -1 + ) else: - bs = freqs.shape[1] - emb = torch.stack((freqs.view(3, bs, -1, 1), freqs.view(3, bs, -1, 1)), dim=-1).view( - 3, bs, freqs.shape[0], -1 - ) - - # generate freqs with mrope_section - # shape (bs, seq_length, 2 * dim) - mrope_section = mrope_section * 2 - emb = torch.cat([m[i % 3] for i, m in enumerate(emb.split(mrope_section, dim=-1))], dim=-1) + # Original section-based layout (Qwen2-VL style). + if not self.rotary_interleaved: + emb = torch.cat((freqs, freqs), dim=-1) # shape (3, bs, seq_length, 2 * dim) + else: + bs = freqs.shape[1] + emb = torch.stack( + (freqs.view(3, bs, -1, 1), freqs.view(3, bs, -1, 1)), dim=-1 + ).view(3, bs, freqs.shape[0], -1) + # generate freqs with mrope_section: cycle T/H/W per section chunk + mrope_section_doubled = list(mrope_section) * 2 + emb = torch.cat( + [m[i % 3] for i, m in enumerate(emb.split(mrope_section_doubled, dim=-1))], dim=-1 + ) # shape (bs, seq_length, 2 * dim) # shape (seq_length, bs, 1, 2 * dim) emb = emb[..., None, :].transpose(0, 1).contiguous() diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index 422a0e71e6a..7bce9d96d2c 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -203,6 +203,7 @@ def __init__( rotary_interleaved=self.config.rotary_interleaved, seq_len_interpolation_factor=seq_len_interpolation_factor, rotary_base=rotary_base, + interleaved_mrope=self.config.mrope_interleaved, ) self.mrope_section = self.config.mrope_section assert ( diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index a7608b9881e..a701d6bd268 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -1150,6 +1150,12 @@ class TransformerConfig(ModelParallelConfig): """ Multimodal rope section is for channel dimension of temporal, height and width in rope calculation. """ + mrope_interleaved: bool = False + """When True, use the interleaved T/H/W MRoPE layout (Qwen3.5-VL style) where + H freqs occupy stride-3 positions {1,4,7,...} and W freqs occupy {2,5,8,...}. + When False (default), use the original section-based layout (Qwen2-VL style) + that cycles through T/H/W per mrope_section chunk.""" + is_hybrid_model: bool = False """ Indicates whether this is a hybrid model. """ diff --git a/tests/unit_tests/models/test_hybrid_moe_model.py b/tests/unit_tests/models/test_hybrid_moe_model.py index 3f5aedf4af3..ad7f91e1a7b 100644 --- a/tests/unit_tests/models/test_hybrid_moe_model.py +++ b/tests/unit_tests/models/test_hybrid_moe_model.py @@ -217,6 +217,7 @@ "moe_token_dropping": False, "moe_z_loss_coeff": None, "moe_enable_routing_replay": False, + "mrope_interleaved": False, "mrope_section": None, "mup_attn_scale_power": 1.0, "mup_base_head_dim": None,