Skip to content
Open
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
27 changes: 22 additions & 5 deletions src/mcore_bridge/model/gpts/qwen3_next_gdn.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,11 +115,8 @@ def _set_linear_decoupled_in_proj(self, mg_attn, hf_state_dict, to_mcore: bool):
class Qwen3NextLoader(ModelLoader):
gated_delta_net = GatedDeltaNet

def get_transformer_layer_spec(self, vp_stage: Optional[int] = None):
from megatron.core.models.gpt.experimental_attention_variant_module_specs import \
get_transformer_block_with_experimental_attention_variant_spec
layer_specs = get_transformer_block_with_experimental_attention_variant_spec(self.config, vp_stage)
for layer_spec in layer_specs.layer_specs:
def _replace_transformer_layer_specs(self, layer_specs):
for layer_spec in layer_specs:
attn_module = layer_spec.submodules.self_attention.module
if issubclass(attn_module, SelfAttention):
layer_spec.submodules.self_attention.module = GatedSelfAttention
Expand All @@ -128,8 +125,28 @@ def get_transformer_layer_spec(self, vp_stage: Optional[int] = None):
if self.config.linear_decoupled_in_proj:
layer_spec.submodules.input_layernorm = TENorm
layer_spec.submodules.self_attention.submodules.in_proj = TEColumnParallelLinear

def get_transformer_layer_spec(self, vp_stage: Optional[int] = None):
from megatron.core.models.gpt.experimental_attention_variant_module_specs import \
get_transformer_block_with_experimental_attention_variant_spec
layer_specs = get_transformer_block_with_experimental_attention_variant_spec(self.config, vp_stage)
self._replace_transformer_layer_specs(layer_specs.layer_specs)
return layer_specs

def get_mtp_block_spec(self, transformer_layer_spec, vp_stage: Optional[int] = None):
from megatron.core.models.gpt.experimental_attention_variant_module_specs import \
get_transformer_layer_with_experimental_attention_variant_spec
decoder_layer_specs = get_transformer_layer_with_experimental_attention_variant_spec(self.config)
self._replace_transformer_layer_specs(decoder_layer_specs)

transformer_layer_spec = copy.deepcopy(transformer_layer_spec)
transformer_layer_spec.layer_specs = [decoder_layer_specs[-1]]
self._set_shared_expert_gate(transformer_layer_spec)
self._set_transformer_layer(transformer_layer_spec)
self._replace_mla_attention(transformer_layer_spec)
self._replace_router(transformer_layer_spec)
return super().get_mtp_block_spec(transformer_layer_spec, vp_stage=vp_stage)

def build_model(
self,
pre_process=True,
Expand Down