Skip to content
Draft
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
20 changes: 9 additions & 11 deletions src/megatron/bridge/models/qwen_vl/modelling_qwen3_vl/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@

from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.attention import Qwen3VLSelfAttention
from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.rope import get_rope_index
from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.text_model import Qwen3VLGPTModel
from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.text_model import Qwen3VLHybridModel
from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.transformer_config import (
Qwen3VLTransformerConfig,
get_vision_model_config,
Expand Down Expand Up @@ -123,9 +123,6 @@ def __init__(
) -> None:
super().__init__(config=language_transformer_config)

if hasattr(language_transformer_layer_spec, "submodules"):
language_transformer_layer_spec.submodules.self_attention.module = Qwen3VLSelfAttention

self.vision_transformer_config = vision_transformer_config
self.pre_process = pre_process
self.post_process = post_process
Expand Down Expand Up @@ -203,30 +200,31 @@ def __init__(
pg_collection=pg_collection,
)
if self.add_decoder:
self.language_model = Qwen3VLGPTModel(
if mtp_block_spec is not None:
raise ValueError("Qwen3 multimodal HybridModel does not accept a separate MTP block spec.")
self.language_model = Qwen3VLHybridModel(
config=language_transformer_config,
transformer_layer_spec=language_transformer_layer_spec,
hybrid_stack_spec=language_transformer_layer_spec,
vocab_size=language_transformer_config.vocab_size,
max_sequence_length=language_transformer_config.language_max_sequence_length,
hybrid_layer_pattern=language_transformer_config.hybrid_layer_pattern,
parallel_output=parallel_output,
position_embedding_type="mrope",
rotary_percent=language_transformer_config.rotary_percent,
pre_process=self.pre_process,
post_process=self.post_process,
rotary_base=language_transformer_config.rotary_base,
fp16_lm_cross_entropy=language_transformer_config.fp16_lm_cross_entropy,
share_embeddings_and_output_weights=language_transformer_config.share_embeddings_and_output_weights,
scatter_embedding_sequence_parallel=False,
mtp_block_spec=mtp_block_spec,
vp_stage=vp_stage,
pg_collection=pg_collection,
)
if pre_process:
deepstack_indexes = getattr(vision_transformer_config, "deepstack_visual_indexes", [])
assert len(deepstack_indexes) <= len(self.language_model.decoder.layers), (
logical_layers = len(self.language_model.decoder.layers) // 2
assert len(deepstack_indexes) <= logical_layers, (
"the deepstack_visual_embeds should on the first pp-stage of language model",
f"got {len(deepstack_indexes)} deepstack_visual_indexes, "
f" {len(self.language_model.decoder.layers)} language model layers",
f"got {len(deepstack_indexes)} deepstack_visual_indexes, {logical_layers} language model layers",
)

self.share_embeddings_and_output_weights = self.language_model.share_embeddings_and_output_weights
Expand Down
26 changes: 24 additions & 2 deletions src/megatron/bridge/models/qwen_vl/modelling_qwen3_vl/rope.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,6 +139,22 @@ def __init__(
self.mrope_section = [24, 20, 20]
assert cp_group is not None, "cp_group is required"
self.cp_group = cp_group
self._position_ids_context: torch.Tensor | None = None
self._mrope_section_context: List[int] | None = None

def set_forward_context(
self,
position_ids: torch.Tensor,
mrope_section: List[int] | None,
) -> None:
"""Set explicit positions for HybridModel's standard RoPE call path."""
self._position_ids_context = position_ids
self._mrope_section_context = mrope_section

def clear_forward_context(self) -> None:
"""Release references to per-forward position tensors."""
self._position_ids_context = None
self._mrope_section_context = None

def apply_interleaved_mrope(self, freqs, mrope_section):
"""Apply interleaved MRoPE to 3D rotary embeddings.
Expand All @@ -159,8 +175,8 @@ def apply_interleaved_mrope(self, freqs, mrope_section):

def forward(
self,
position_ids: torch.Tensor,
mrope_section: List[int] | None,
position_ids: torch.Tensor | int,
mrope_section: List[int] | None = None,
packed_seq_params: Optional[PackedSeqParams] = None,
**kwargs,
) -> Tensor:
Expand All @@ -174,6 +190,12 @@ def forward(
Returns:
Tensor: Embeddings after applying RoPE.
"""
if not isinstance(position_ids, torch.Tensor):
if self._position_ids_context is None:
raise RuntimeError("Qwen mRoPE positions must be set before HybridModel forward.")
position_ids = self._position_ids_context
mrope_section = self._mrope_section_context

if position_ids.ndim == 2:
position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1)
# Use fp32 for position indices to avoid precision loss when inv_freq is bf16.
Expand Down
Loading
Loading