已合并
[pytorch][bugfix]fix bug of ckpt-v2 #3398
温一盏创建于 2025年9月25日
[pytorch][bugfix]fix bug of ckpt-v2 #3398
已合并
温一盏创建于 2025年9月25日
从refs/pull/3398/head合入到master
共 2 个文件变更+3-7
@@ -40,8 +40,6 @@ def get_args():
40 help="use tp group to extend experts parallism instead of sharding weight tensor of experts in tp group")40 help="use tp group to extend experts parallism instead of sharding weight tensor of experts in tp group")
41 parser.add_argument('--mla-mm-split', action='store_true', default=False,41 parser.add_argument('--mla-mm-split', action='store_true', default=False,
42 help='Split 2 up-proj matmul into 4 in MLA')42 help='Split 2 up-proj matmul into 4 in MLA')
43- parser.add_argument("--shared-expert-gate", action='store_true',
44- help="moe model has shared expert gate")
45 parser.add_argument('--schedules-method', type=str, default=None, choices=['dualpipev'],43 parser.add_argument('--schedules-method', type=str, default=None, choices=['dualpipev'],
46 help='An innovative bidirectional pipeline parallelism algorithm.')44 help='An innovative bidirectional pipeline parallelism algorithm.')
47 parser.add_argument('--first-k-dense-replace', type=int, default=None,45 parser.add_argument('--first-k-dense-replace', type=int, default=None,
@@ -373,10 +373,10 @@ class Hf2MgConvert(Convert):
373 if mtp_flag:373 if mtp_flag:
374 qkv_key = mg_weight_key["mtp_layers_self_attention_linear_qkv"]374 qkv_key = mg_weight_key["mtp_layers_self_attention_linear_qkv"]
375 dense_key = mg_weight_key["mtp_layers_self_attention_linear_proj"]375 dense_key = mg_weight_key["mtp_layers_self_attention_linear_proj"]
376- q_b_key = mg_weight_key["layers_self_attention_linear_q_up_proj"]376+ q_b_key = mg_weight_key["mtp_layers_self_attention_linear_q_up_proj"]
377- kv_b_key = mg_weight_key["layers_self_attention_linear_kv_up_proj"]377+ kv_b_key = mg_weight_key["mtp_layers_self_attention_linear_kv_up_proj"]
378 q_layernorm_key = mg_weight_key["mtp_layers_self_attention_q_layernorm"]378 q_layernorm_key = mg_weight_key["mtp_layers_self_attention_q_layernorm"]
379- kv_layernorm_key = mg_weight_key["layers_self_attention_kv_layernorm"]379+ kv_layernorm_key = mg_weight_key["mtp_layers_self_attention_kv_layernorm"]
380 else:380 else:
381 qkv_key = mg_weight_key["layers_self_attention_linear_qkv"]381 qkv_key = mg_weight_key["layers_self_attention_linear_qkv"]
382 dense_key = mg_weight_key["layers_self_attention_linear_proj"]382 dense_key = mg_weight_key["layers_self_attention_linear_proj"]
@@ -479,8 +479,6 @@ class Hf2MgConvert(Convert):
479 ], dim=1).reshape(-1)479 ], dim=1).reshape(-1)
480 480 
481 if self.load_model.qkv_type == "pack_mla":481 if self.load_model.qkv_type == "pack_mla":
482- qkv_key, dense_key, q_layernorm_key, kv_layernorm_key, q_b_key, kv_b_key = _generate_mla_attn_layers_key(
483- mtp_layer_flag)
484 hf_q_proj = hf_weight.pop(hf_weight_key["layers_self_attention_linear_q_proj"])482 hf_q_proj = hf_weight.pop(hf_weight_key["layers_self_attention_linear_q_proj"])
485 hf_kv_proj = hf_weight.pop(hf_weight_key["layers_self_attention_linear_kv_proj"])483 hf_kv_proj = hf_weight.pop(hf_weight_key["layers_self_attention_linear_kv_proj"])
486 qkv_weight = torch.cat([hf_q_proj.reshape((-1, self.load_model.hidden_size)),484 qkv_weight = torch.cat([hf_q_proj.reshape((-1, self.load_model.hidden_size)),