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
60 changes: 49 additions & 11 deletions megatron/core/models/common/embeddings/rotary_pos_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__(
Expand All @@ -290,13 +293,15 @@ 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__()

dim = kv_channels
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 / (
Expand All @@ -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,
Expand Down Expand Up @@ -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()
Expand Down
1 change: 1 addition & 0 deletions megatron/core/models/gpt/gpt_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down
6 changes: 6 additions & 0 deletions megatron/core/transformer/transformer_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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. """

Expand Down
1 change: 1 addition & 0 deletions tests/unit_tests/models/test_hybrid_moe_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Loading