llama_saver.py — LLaMA 权重合并保存¶
文件路径¶
verl/models/llama/megatron/checkpoint_utils/llama_saver.py
文件概述¶
将分片的 Megatron LLaMA 模型权重收集并合并回 HuggingFace 格式的完整 state_dict。与 llama_loader.py 互为逆操作。
关键代码讲解¶
主函数¶
def merge_megatron_ckpt_llama(wrapped_models, config, dtype,
is_value_model=False, tie_word_embeddings=False):
"""合并所有 TP/PP 分片的权重为完整的 HF state_dict"""
全局 rank 计算¶
def _megatron_calc_global_rank(tp_rank=0, dp_rank=0, pp_rank=0):
"""根据 TP/DP/PP 坐标计算全局 rank"""
# 支持 TP-DP-PP 排列顺序
return (pp_rank * dp_size + dp_rank) * tp_size + tp_rank
QKV 权重的反合并¶
def _broadcast_tp_shard_tensor_qkv(tensor, q_name, k_name, v_name, src_pp_rank):
"""从各 TP rank 收集 QKV 分片,拆分为独立的 Q/K/V 权重"""
# 1. 从每个 TP rank 收集分片
for i in range(tp_size):
cur_src_rank = _megatron_calc_global_rank(tp_rank=i, pp_rank=src_pp_rank)
dist.broadcast(sync_tensor, src=cur_src_rank, group=mp_group)
chunk_tensors[i] = sync_tensor.cpu()
# 2. 在 rank 0 上拼接并拆分为 Q/K/V
full_tensor = torch.concat(chunk_tensors, dim=0)
for i in range(tp_size):
qkv_part = full_tensor[i*total_size : (i+1)*total_size]
q_weight_list.append(qkv_part[:q_size_tp])
k_weight_list.append(qkv_part[q_size_tp : q_size_tp+kv_size_tp])
v_weight_list.append(qkv_part[q_size_tp+kv_size_tp:])
state_dict[q_name] = torch.cat(q_weight_list, dim=0)
state_dict[k_name] = torch.cat(k_weight_list, dim=0)
state_dict[v_name] = torch.cat(v_weight_list, dim=0)
gate+up 权重的反合并¶
def _broadcast_tp_shard_tensor_gate_up(tensor, gate_name, up_name, src_pp_rank):
"""收集合并的 gate_up 分片,拆分为独立的 gate 和 up"""
full_tensor = torch.concat(chunk_tensors, dim=0)
for i in range(tp_size):
gate_up_tp = full_tensor[intermediate_tp*2*i : intermediate_tp*2*(i+1)]
gate_weight_list.append(gate_up_tp[:intermediate_tp])
up_weight_list.append(gate_up_tp[intermediate_tp:])
state_dict[gate_name] = torch.cat(gate_weight_list, dim=0)
state_dict[up_name] = torch.cat(up_weight_list, dim=0)
合并流程¶
1. 从 PP rank 0 收集 embedding(TP 各 rank 的分片拼接)
2. 逐层从对应 PP stage 收集权重:
- layernorm: 直接广播(不拆分)
- qkv_proj: 收集 TP 分片 -> 拆分为 q/k/v
- o_proj: 收集 TP 分片 -> 按 dim=1 拼接
- gate_up_proj: 收集 TP 分片 -> 拆分为 gate/up
- down_proj: 收集 TP 分片 -> 按 dim=1 拼接
3. 从最后一个 PP stage 收集 norm 和 lm_head
4. 只有 rank 0 最终持有完整 state_dict
与其他模块的关系¶
- 被
weight_loader_registry.py注册为 LLaMA 的权重保存器 - 与
llama_loader.py/llama_loader_depracated.py互为逆操作
小结¶
权重合并是加载的逆过程。关键操作是从各 TP rank 收集权重分片、按正确维度拼接、以及将合并的 QKV/gate_up 拆分回独立的 HF 格式权重。最终结果在 rank 0 上形成完整的 HF state_dict。