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
6 changes: 2 additions & 4 deletions examples/multimodal_dev/models/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,10 +45,8 @@ def _cp_split_tensor(tensor, seq_dim, cp_size, cp_rank):


class _NoCPGroup:
"""Dummy size-1 process group used to bypass MRoPE's BSHD-style
zigzag of pre-computed THD freqs (Megatron-Core gap:
``MultimodalRotaryEmbedding.forward`` lacks the ``not packed_seq``
skip that plain ``RotaryEmbedding`` has).
"""Dummy size-1 process group used to bypass BSHD-style CP slicing
for THD MRoPE call sites that do not pass ``packed_seq=True``.
"""

def size(self):
Expand Down
13 changes: 13 additions & 0 deletions examples/multimodal_dev/models/qwen35_vl/configuration.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,13 @@ def get_qwen35_vl_vision_config(
if num_layers_override is not None:
num_layers = num_layers_override

vision_head_dim = vcfg["kv_channels"]
assert vision_head_dim % 4 == 0, (
"Qwen3.5-VL vision RoPE expects the per-head dimension to split "
f"evenly across row/column frequencies, got {vision_head_dim}"
)
vision_rope_axis_dim = vision_head_dim // 4

return TransformerConfig(
num_layers=num_layers,
hidden_size=vcfg["hidden_size"],
Expand All @@ -120,6 +127,12 @@ def get_qwen35_vl_vision_config(
bias_activation_fusion=False,
apply_query_key_layer_scaling=False,
apply_rope_fusion=False,
# Vision RoPE is 2D row/column RoPE. Represent it as sectioned raw
# mRoPE with a zero temporal section so the fused mRoPE dispatcher can
# reuse the same Triton kernel when rope fusion is enabled.
mrope_section=[0, vision_rope_axis_dim, vision_rope_axis_dim],
mrope_interleaved=False,
rotary_interleaved=False,
bf16=False,
)

Expand Down
83 changes: 57 additions & 26 deletions examples/multimodal_dev/models/qwen35_vl/specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from megatron.core.transformer.spec_utils import ModuleSpec
from megatron.core.transformer.transformer_block import TransformerBlockSubmodules
from megatron.core.transformer.transformer_config import TransformerConfig
from megatron.core.utils import nvtx_range_pop, nvtx_range_push


def _apply_rope_fp32(t, freqs, config, cu_seqlens=None, mscale=1.0, cp_group=None):
Expand All @@ -25,35 +26,55 @@ def _apply_rope_fp32(t, freqs, config, cu_seqlens=None, mscale=1.0, cp_group=Non
Mirrors ``Qwen3VLSelfAttention.apply_rotary_pos_emb_absolute`` in Megatron-Bridge
with ``apply_rotary_pos_emb_in_fp32=True``.
"""
from megatron.core import parallel_state
from megatron.core.models.common.embeddings.rope_utils import (
_apply_rotary_pos_emb_bshd,
_apply_rotary_pos_emb_thd,
)
from megatron.core.models.common.embeddings import rope_utils
from megatron.core.models.common.embeddings.rope_utils import apply_rotary_pos_emb

orig_dtype = t.dtype
t_fp32 = t.float()

if cu_seqlens is None:
out = _apply_rotary_pos_emb_bshd(
t_fp32,
freqs,
rotary_interleaved=config.rotary_interleaved,
multi_latent_attention=getattr(config, 'multi_latent_attention', False),
mscale=mscale,
)
else:
if cp_group is None:
cp_group = parallel_state.get_context_parallel_group()
out = _apply_rotary_pos_emb_thd(
t_fp32,
if (
cu_seqlens is not None
and getattr(config, "apply_rope_fusion", False)
and getattr(config, "mrope_section", None) is not None
and getattr(config, "rotary_interleaved", False) is False
and getattr(config, "multi_latent_attention", False) is False
and mscale == 1.0
and t.dim() == 3
and freqs.dim() == 4
and freqs.shape[0] == 3
and cp_group is not None
and rope_utils.fused_apply_mrope_thd is not None
and rope_utils.get_fused_mrope_thd_unavailable_reason is not None
):
unavailable_reason = rope_utils.get_fused_mrope_thd_unavailable_reason(
t,
cu_seqlens,
freqs,
rotary_interleaved=config.rotary_interleaved,
multi_latent_attention=getattr(config, 'multi_latent_attention', False),
mscale=mscale,
cp_group=cp_group,
cp_size=cp_group.size(),
cp_rank=cp_group.rank(),
)
if unavailable_reason is None:
return rope_utils.fused_apply_mrope_thd(
t,
cu_seqlens,
freqs,
config.mrope_section,
interleaved_mrope=config.mrope_interleaved,
rotary_interleaved=config.rotary_interleaved,
cp_size=cp_group.size(),
cp_rank=cp_group.rank(),
fp32_compute=True,
)

t_fp32 = t.float()
out = apply_rotary_pos_emb(
t_fp32,
freqs,
config=config,
cu_seqlens=cu_seqlens,
mscale=mscale,
cp_group=cp_group,
mla_rotary_interleaved=getattr(config, 'multi_latent_attention', False),
)
return out.to(orig_dtype)


Expand All @@ -65,9 +86,19 @@ def _apply_rope_fp32_no_cp(t, freqs, config, cu_seqlens=None, mscale=1.0, cp_gro
incorrectly split the vision seqlens. This wrapper substitutes a
trivial group so the vision RoPE sees the full packed sequence.
"""
return _apply_rope_fp32(
t, freqs, config, cu_seqlens, mscale, cp_group=_NO_CP_GROUP,
)
range_name = "qwen35_vl.vision_encoder.rope_apply"
nvtx_range_push(range_name)
try:
return _apply_rope_fp32(
t,
freqs,
config,
cu_seqlens,
mscale,
cp_group=_NO_CP_GROUP,
)
finally:
nvtx_range_pop(range_name)


class Qwen35VLVisionSelfAttention(SelfAttention):
Expand Down
48 changes: 32 additions & 16 deletions examples/multimodal_dev/models/qwen35_vl/vision_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,15 +26,10 @@
import torch.nn.functional as F
from torch import Tensor

from megatron.core.models.common.vision_module.vision_module import (
VisionModule,
)
from megatron.core.packed_seq_params import PackedSeqParams
from megatron.core.tensor_parallel.layers import (
ColumnParallelLinear,
RowParallelLinear,
)
from megatron.core.extensions.transformer_engine import TENorm
from megatron.core.models.common.vision_module.vision_module import VisionModule
from megatron.core.packed_seq_params import PackedSeqParams
from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear
from megatron.core.transformer.module import MegatronModule
from megatron.core.transformer.spec_utils import ModuleSpec, build_module
from megatron.core.transformer.transformer_block import TransformerBlock
Expand Down Expand Up @@ -322,9 +317,7 @@ def __init__(

# --- Transformer blocks ---
if transformer_layer_spec is None:
from examples.multimodal_dev.models.qwen35_vl.specs import (
get_qwen35_vl_vision_spec,
)
from examples.multimodal_dev.models.qwen35_vl.specs import get_qwen35_vl_vision_spec
transformer_layer_spec = get_qwen35_vl_vision_spec()

self.decoder = TransformerBlock(
Expand Down Expand Up @@ -455,7 +448,9 @@ def _compute_rotary_pos_emb(self, grid_thw: Tensor) -> Tensor:
grid_thw: ``[num_images, 3]`` (T, H, W) per image.

Returns:
``[total_patches, head_dim // 2]`` raw RoPE frequencies.
Raw sectioned frequencies ``[3, 1, total_patches, head_dim // 2]``
when ``config.mrope_section`` is set. Otherwise returns the legacy
``[total_patches, head_dim // 2]`` row/column frequency tensor.
"""
merge = self.spatial_merge_size
grid_thw_list = grid_thw.tolist()
Expand Down Expand Up @@ -512,7 +507,27 @@ def _compute_rotary_pos_emb(self, grid_thw: Tensor) -> Tensor:

embeddings = freq_table[pos_ids]
embeddings = embeddings.flatten(1)
return embeddings

mrope_section = getattr(self.config, "mrope_section", None)
if mrope_section is None:
return embeddings

sec_t, sec_h, sec_w = (int(section) for section in mrope_section)
if sec_t != 0 or sec_h + sec_w != embeddings.shape[-1]:
raise ValueError(
"Qwen3.5-VL vision RoPE expects mrope_section "
f"[0, row_dim, col_dim] summing to {embeddings.shape[-1]}, "
f"got {mrope_section}"
)

raw_freqs = embeddings.new_zeros(
3, 1, embeddings.shape[0], embeddings.shape[1],
)
raw_freqs[1, 0, :, :sec_h] = embeddings[:, :sec_h]
raw_freqs[2, 0, :, sec_h : sec_h + sec_w] = embeddings[
:, sec_h : sec_h + sec_w
]
return raw_freqs

# ---------------------------------------------------------------
# PackedSeqParams for variable-length attention
Expand Down Expand Up @@ -575,16 +590,17 @@ def forward(

# 3. 2D Vision RoPE
rot_freqs = self._compute_rotary_pos_emb(grid_thw)
emb = torch.cat((rot_freqs, rot_freqs), dim=-1)
rot_freqs_expanded = emb.unsqueeze(1).unsqueeze(1)
if getattr(self.config, "mrope_section", None) is None:
emb = torch.cat((rot_freqs, rot_freqs), dim=-1)
rot_freqs = emb.unsqueeze(1).unsqueeze(1)

# 4. Transformer blocks with PackedSeqParams
packed_seq_params = self._build_packed_seq_params(grid_thw)
hidden_states = hidden_states.unsqueeze(1)
hidden_states = self.decoder(
hidden_states=hidden_states,
attention_mask=None,
rotary_pos_emb=rot_freqs_expanded,
rotary_pos_emb=rot_freqs,
packed_seq_params=packed_seq_params,
)
hidden_states = hidden_states.squeeze(1)
Expand Down
1 change: 1 addition & 0 deletions examples/multimodal_dev/pretrain_multimodal.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,7 @@ def model_provider(
)
vision_config.bf16 = language_config.bf16
vision_config.fp16 = language_config.fp16
vision_config.apply_rope_fusion = language_config.apply_rope_fusion

if getattr(args, "recompute_vision", False):
vision_config.recompute_granularity = "full"
Expand Down
Loading