跳转至

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。