Files
ModelHub XC d4e0a1af66 初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD
Source: Original Platform
2026-08-27 23:50:14 +08:00

105 lines
4.2 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import torch
from mbridge.core import register_model
from mbridge.models import Qwen2MoEBridge
@register_model("qwen3_next")
class Qwen3NextBridge(Qwen2MoEBridge):
_ATTENTION_MAPPING = (
Qwen2MoEBridge._ATTENTION_MAPPING
| {
f"self_attention.{weight_name}": ["model.layers.{layer_number}." + weight_name]
for weight_name in [
"input_layernorm.weight",
# linear attn
"linear_attn.A_log",
"linear_attn.conv1d.weight",
"linear_attn.dt_bias",
"linear_attn.in_proj_ba.weight",
"linear_attn.in_proj_qkvz.weight",
"linear_attn.norm.weight",
"linear_attn.out_proj.weight",
# gated attn
"self_attn.k_norm.weight",
"self_attn.k_proj.weight",
"self_attn.o_proj.weight",
"self_attn.q_norm.weight",
"self_attn.q_proj.weight",
"self_attn.v_proj.weight",
]
}
| {
"self_attention.linear_qkv.layer_norm_weight": ["model.layers.{layer_number}.input_layernorm.weight"],
"self_attention.linear_qkv.weight": [
"model.layers.{layer_number}.self_attn.q_proj.weight",
"model.layers.{layer_number}.self_attn.k_proj.weight",
"model.layers.{layer_number}.self_attn.v_proj.weight",
],
}
)
def _weight_to_mcore_format(
self, mcore_weights_name: str, hf_weights: list[torch.Tensor]
) -> tuple[list[str], list[torch.Tensor]]:
if "self_attention.linear_qkv." in mcore_weights_name and "layer_norm" not in mcore_weights_name:
# merge qkv
assert len(hf_weights) == 3
num_key_value_heads = self.hf_config.num_key_value_heads
hidden_dim = self.hf_config.hidden_size
num_attention_heads = self.hf_config.num_attention_heads
num_querys_per_group = num_attention_heads // self.hf_config.num_key_value_heads
head_dim = getattr(self.hf_config, "head_dim", hidden_dim // num_attention_heads)
group_dim = head_dim * num_attention_heads // num_key_value_heads
q, k, v = hf_weights
# q k v might be tp split
real_num_key_value_heads = q.shape[0] // (2 * group_dim)
q = (
q.view(
[
real_num_key_value_heads,
num_querys_per_group,
2,
head_dim,
-1,
]
)
.transpose(1, 2)
.flatten(1, 3)
)
k = k.view([real_num_key_value_heads, head_dim, -1])
v = v.view([real_num_key_value_heads, head_dim, -1])
out_shape = [-1, hidden_dim] if ".bias" not in mcore_weights_name else [-1]
qgkv = torch.cat([q, k, v], dim=1).view(*out_shape).contiguous()
return qgkv
return super()._weight_to_mcore_format(mcore_weights_name, hf_weights)
def _build_config(self):
return self._build_base_config(
use_cpu_initialization=False,
# MoE specific
moe_ffn_hidden_size=self.hf_config.moe_intermediate_size,
moe_router_bias_update_rate=0.001,
moe_router_topk=self.hf_config.num_experts_per_tok,
num_moe_experts=self.hf_config.num_experts,
moe_aux_loss_coeff=self.hf_config.router_aux_loss_coef,
# moe_router_load_balancing_type="aux_loss",
moe_router_load_balancing_type="none", # default None for RL
moe_grouped_gemm=True,
moe_router_score_function="softmax",
# Other optimizations
persist_layer_norm=True,
bias_activation_fusion=True,
bias_dropout_fusion=True,
# Qwen specific
moe_router_pre_softmax=False,
qk_layernorm=True,
# Qwen3 Next specific
attention_output_gate=True,
moe_shared_expert_gate=True,
)