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
Original file line number Diff line number Diff line change
Expand Up @@ -27,8 +27,7 @@
)
from megatron.bridge.models.qwen3_asr.modeling_qwen3_asr.rope import get_rope_index
from megatron.bridge.models.qwen3_asr.modeling_qwen3_asr.transformer_config import Qwen3ASRTransformerConfig
from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.attention import Qwen3VLSelfAttention
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.utils import (
split_data_cp_rank,
)
Expand Down Expand Up @@ -56,8 +55,6 @@ def __init__(
) -> None:
super().__init__(config=language_transformer_config)

language_transformer_layer_spec.submodules.self_attention.module = Qwen3VLSelfAttention

self.pre_process = pre_process
self.post_process = post_process
self.add_encoder = add_encoder
Expand Down Expand Up @@ -95,13 +92,13 @@ def __init__(
self.audio_model = Qwen3ASRAudioEncoderHF._from_config(thinker_transformer_config.audio_config)
hook_hf_module_setattr_for_tp_grad_sync(self.audio_model)

self.language_model = Qwen3VLGPTModel(
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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ class Qwen3ASRTransformerConfig(TransformerConfig):
share_embeddings_and_output_weights: bool = False
rotary_percent: float = 1.0
rotary_base: float = 5000000.0
hybrid_layer_pattern: str | None = None

# Multimodal rope section for 3 dimensions (same position IDs across all dims for ASR)
mrope_section: list[int] = field(default_factory=lambda: [24, 20, 20])
Expand Down
93 changes: 63 additions & 30 deletions src/megatron/bridge/models/qwen3_asr/qwen3_asr_bridge.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
# limitations under the License.

import torch
from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols

from megatron.bridge.models.conversion.mapping_registry import MegatronMappingRegistry
from megatron.bridge.models.conversion.model_bridge import MegatronModelBridge
Expand All @@ -23,14 +24,24 @@
ReplicatedMapping,
)
from megatron.bridge.models.hf_pretrained.causal_lm import PreTrainedCausalLM
from megatron.bridge.models.qwen.qwen_hybrid import (
configure_qwen_hybrid_layers,
qwen_logical_layer_count,
qwen_physical_layer_indices,
)
from megatron.bridge.models.qwen3_asr.modeling_qwen3_asr.model import Qwen3ASRModel
from megatron.bridge.models.qwen3_asr.qwen3_asr_provider import Qwen3ASRModelProvider


# Use string-based registration because Qwen3ASRForConditionalGeneration is not in
# the standard transformers library (it's a custom model in qwen_asr package).
# auto_bridge.py resolves custom architectures via config.auto_map or string fallback.
@MegatronModelBridge.register_bridge(source="Qwen3ASRForConditionalGeneration", target=Qwen3ASRModel)
@MegatronModelBridge.register_bridge(
source="Qwen3ASRForConditionalGeneration",
target=Qwen3ASRModel,
provider=Qwen3ASRModelProvider,
model_type="qwen3_asr",
)
class Qwen3ASRBridge(MegatronModelBridge):
"""
Megatron Bridge for Qwen3-ASR Conditional Generation.
Expand Down Expand Up @@ -79,31 +90,57 @@ def provider_bridge(self, hf_pretrained: PreTrainedCausalLM) -> Qwen3ASRModelPro
audio_start_token_id=getattr(thinker_config, "audio_start_token_id", 151647),
mrope_section=(getattr(text_config, "rope_scaling", None) or {}).get("mrope_section", [24, 20, 20]),
)
configure_qwen_hybrid_layers(
provider,
num_logical_layers=text_config.num_hidden_layers,
mlp_symbols=Symbols.MLP,
mtp_mlp_symbol=Symbols.MLP,
)
return provider

def mapping_registry(self) -> MegatronMappingRegistry:
"""Return MegatronMappingRegistry containing parameter mappings for Qwen3-ASR models."""
# LLM parameter mappings (Qwen3-style with QK layernorm, prefixed with thinker.)
param_mappings = {
# Embeddings and output layers
"thinker.language_model.embedding.word_embeddings.weight": "thinker.model.embed_tokens.weight",
"thinker.language_model.output_layer.weight": "thinker.lm_head.weight",
"thinker.language_model.decoder.final_layernorm.weight": "thinker.model.norm.weight",
# Layer normalization
"thinker.language_model.decoder.layers.*.self_attention.linear_qkv.layer_norm_weight": "thinker.model.layers.*.input_layernorm.weight",
"thinker.language_model.decoder.layers.*.mlp.linear_fc1.layer_norm_weight": "thinker.model.layers.*.post_attention_layernorm.weight",
# QK layernorm (Qwen3-specific)
"thinker.language_model.decoder.layers.*.self_attention.q_layernorm.weight": "thinker.model.layers.*.self_attn.q_norm.weight",
"thinker.language_model.decoder.layers.*.self_attention.k_layernorm.weight": "thinker.model.layers.*.self_attn.k_norm.weight",
# Attention output projection
"thinker.language_model.decoder.layers.*.self_attention.linear_proj.weight": "thinker.model.layers.*.self_attn.o_proj.weight",
# MLP down projection
"thinker.language_model.decoder.layers.*.mlp.linear_fc2.weight": "thinker.model.layers.*.mlp.down_proj.weight",
"thinker.language_model.decoder.final_norm.weight": "thinker.model.norm.weight",
}

mapping_list = []
for megatron_param, hf_param in param_mappings.items():
mapping_list.append(AutoMapping(megatron_param=megatron_param, hf_param=hf_param))
mapping_list = [AutoMapping(k, v) for k, v in param_mappings.items()]

num_layers = self.hf_config.thinker_config.text_config.num_hidden_layers
for logical_layer_idx in range(num_layers):
attention_layer_idx, mlp_layer_idx = qwen_physical_layer_indices(logical_layer_idx)
hf_layer = f"thinker.model.layers.{logical_layer_idx}"
attention_layer = f"thinker.language_model.decoder.layers.{attention_layer_idx}.self_attention"
mlp_layer = f"thinker.language_model.decoder.layers.{mlp_layer_idx}.mlp"
mapping_list.extend(
[
AutoMapping(
f"{attention_layer}.linear_qkv.layer_norm_weight",
f"{hf_layer}.input_layernorm.weight",
),
AutoMapping(f"{attention_layer}.q_layernorm.weight", f"{hf_layer}.self_attn.q_norm.weight"),
AutoMapping(f"{attention_layer}.k_layernorm.weight", f"{hf_layer}.self_attn.k_norm.weight"),
AutoMapping(f"{attention_layer}.linear_proj.weight", f"{hf_layer}.self_attn.o_proj.weight"),
AutoMapping(
f"{mlp_layer}.linear_fc1.layer_norm_weight",
f"{hf_layer}.post_attention_layernorm.weight",
),
AutoMapping(f"{mlp_layer}.linear_fc2.weight", f"{hf_layer}.mlp.down_proj.weight"),
QKVMapping(
megatron_param=f"{attention_layer}.linear_qkv.weight",
q=f"{hf_layer}.self_attn.q_proj.weight",
k=f"{hf_layer}.self_attn.k_proj.weight",
v=f"{hf_layer}.self_attn.v_proj.weight",
),
GatedMLPMapping(
megatron_param=f"{mlp_layer}.linear_fc1.weight",
gate=f"{hf_layer}.mlp.gate_proj.weight",
up=f"{hf_layer}.mlp.up_proj.weight",
),
]
)

mapping_list.extend(
[
Expand All @@ -113,20 +150,16 @@ def mapping_registry(self) -> MegatronMappingRegistry:
megatron_param="thinker.audio_model.**",
hf_param="thinker.audio_tower.**",
),
# QKV weight: Combine separate Q, K, V weights into single QKV matrix (no bias for Qwen3)
QKVMapping(
megatron_param="thinker.language_model.decoder.layers.*.self_attention.linear_qkv.weight",
q="thinker.model.layers.*.self_attn.q_proj.weight",
k="thinker.model.layers.*.self_attn.k_proj.weight",
v="thinker.model.layers.*.self_attn.v_proj.weight",
),
# Gated MLP: Combine gate and up projection matrices into single FC1 matrix
GatedMLPMapping(
megatron_param="thinker.language_model.decoder.layers.*.mlp.linear_fc1.weight",
gate="thinker.model.layers.*.mlp.gate_proj.weight",
up="thinker.model.layers.*.mlp.up_proj.weight",
),
]
)

return MegatronMappingRegistry(*mapping_list)

@classmethod
def megatron_to_hf_config(cls, provider) -> dict:
"""Restore the logical Qwen layer count when exporting HybridModel config."""
hf_config = super().megatron_to_hf_config(provider)
logical_layer_count = qwen_logical_layer_count(provider.hybrid_layer_pattern)
if logical_layer_count is not None:
hf_config["num_hidden_layers"] = logical_layer_count
return hf_config
50 changes: 36 additions & 14 deletions src/megatron/bridge/models/qwen3_asr/qwen3_asr_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,21 +23,26 @@
from typing import Callable

import torch.nn.functional as F
from megatron.core.models.gpt import GPTModel as MCoreGPTModel
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec
from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols
from megatron.core.transformer.spec_utils import ModuleSpec

from megatron.bridge.models.gpt_provider import GPTModelProvider
from megatron.bridge.models.qwen.qwen_hybrid import QwenHybridModelProvider, configure_qwen_hybrid_layers
from megatron.bridge.models.qwen3_asr.hf_qwen3_asr.configuration_qwen3_asr import (
Qwen3ASRThinkerConfig,
)
from megatron.bridge.models.qwen3_asr.modeling_qwen3_asr.model import Qwen3ASRModel
from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.text_model import (
Qwen3VLHybridModel,
get_qwen3_vl_hybrid_stack_spec,
)
from megatron.bridge.models.qwen_vl.qwen3_vl_provider import _provide_qwen3_vl_language_model


@dataclass
class Qwen3ASRModelProvider(GPTModelProvider):
class Qwen3ASRModelProvider(QwenHybridModelProvider):
"""
Base model provider for Qwen3-ASR Models.
Inherits language model configuration from GPTModelProvider with Qwen3-specific defaults.
Inherits language model configuration from HybridModelProvider with Qwen3-specific defaults.

Key characteristics:
- Audio-only (no vision, no video)
Expand Down Expand Up @@ -85,19 +90,36 @@ class Qwen3ASRModelProvider(GPTModelProvider):
distribute_saved_activations: bool = False
cp_comm_type: str = "p2p"
gradient_accumulation_fusion: bool = False
mtp_num_layers: int | None = None
hybrid_stack_spec: ModuleSpec | Callable = get_qwen3_vl_hybrid_stack_spec

def finalize(self) -> None:
if self.hybrid_layer_pattern is None:
if self.num_layers is None:
raise ValueError("num_layers must be configured for Qwen3-ASR")
configure_qwen_hybrid_layers(
self,
num_logical_layers=self.num_layers,
mlp_symbols=Symbols.MLP,
mtp_mlp_symbol=Symbols.MLP,
)
super().finalize()

def provide(self, pre_process=None, post_process=None, vp_stage=None):
"""Provide a Qwen3-ASR model instance with audio and language components."""
if self.hybrid_layer_pattern is None:
if self.num_layers is None:
raise ValueError("num_layers must be configured for Qwen3-ASR")
configure_qwen_hybrid_layers(
self,
num_logical_layers=self.num_layers,
mlp_symbols=Symbols.MLP,
mtp_mlp_symbol=Symbols.MLP,
)
language_transformer_config = self
thinker_config = self.thinker_config

# Qwen3 GPT layer spec with QK layernorm
language_transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec(
num_experts=None,
moe_grouped_gemm=False,
qk_layernorm=self.qk_layernorm,
fp8=False,
)
language_transformer_layer_spec = self._resolve_hybrid_stack_spec()

model = Qwen3ASRModel(
language_transformer_config=language_transformer_config,
Expand All @@ -116,6 +138,6 @@ def provide(self, pre_process=None, post_process=None, vp_stage=None):

return model

def provide_language_model(self, pre_process=None, post_process=None, vp_stage=None) -> MCoreGPTModel:
def provide_language_model(self, pre_process=None, post_process=None, vp_stage=None) -> Qwen3VLHybridModel:
"""Provide just the language model component without audio."""
return GPTModelProvider.provide(self, pre_process=pre_process, post_process=post_process, vp_stage=vp_stage)
return _provide_qwen3_vl_language_model(self, pre_process, post_process, vp_stage)
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@
from megatron.bridge.models.qwen_omni.modeling_qwen3_omni.transformer_config import (
Qwen3OmniTransformerConfig,
)
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.utils import (
split_data_cp_rank,
split_deepstack_embs,
Expand Down Expand Up @@ -215,13 +215,13 @@ def __init__(
_enable_multimodal_gradient_checkpointing(self.visual)
_enable_multimodal_gradient_checkpointing(self.audio_model)

self.language_model = Qwen3VLGPTModel(
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,
Expand Down
Loading
Loading