跳转至

llama_loader.py — LLaMA 权重加载(广播版)

文件路径

verl/models/llama/megatron/checkpoint_utils/llama_loader.py

文件概述

将 HuggingFace 格式的完整 state_dict 加载到 Megatron 分片模型中。这是广播版本(deprecated 版本的升级),rank 0 广播权重到所有 rank,然后各 rank 根据自己的 TP/PP 位置截取对应分片。

关键代码讲解

主函数

def load_state_dict_to_megatron_llama(
    state_dict, wrapped_models, config, params_dtype,
    is_value_model=False, tie_word_embeddings=False
):

核心辅助函数

1. 层映射计算

def _megatron_calc_layer_map(config):
    """计算全局层号 -> (pp_rank, vpp_rank, local_layer_idx) 的映射"""
    for pp_rank_idx in range(pp_size):
        for virtual_pp_rank_idx in range(virtual_pp_size):
            layer_offset = vpp_rank * (num_layers // vpp_size) + pp_rank * layers_per_model
            for layer_idx in range(layers_per_model):
                layer_map[layer_offset + layer_idx] = (pp_rank_idx, vpp_rank_idx, layer_idx)

2. 普通张量广播

def _broadcast_tensor(tensor, name):
    """rank 0 广播完整张量到所有 rank"""
    # 1. rank 0 获取 shape,广播给其他 rank
    dist.broadcast_object_list([tensor_shape], src=0, group=mp_group)
    # 2. 广播权重数据
    dist.broadcast(tensor, src=0, group=mp_group)

3. QKV 权重的 TP 切分广播

def _broadcast_tp_shard_tensor_qkv(tensor, q_name, k_name, v_name):
    """rank 0 将 Q/K/V 合并、按 TP 交错排列后,逐个广播给各 TP rank"""
    # GQA 情况:num_key_value_heads >= tp_size
    q_size_tp = config.hidden_size // tp_size
    kv_size_tp = head_dim * config.num_key_value_heads // tp_size
    total_size = q_size_tp + 2 * kv_size_tp

    # 按 TP rank 交错排列 [Q_tp0, K_tp0, V_tp0, Q_tp1, K_tp1, V_tp1, ...]
    for i in range(tp_size):
        new_weight_qkv[i*total_size : (i+1)*total_size] = cat([Q_i, K_i, V_i])

    # 逐个 TP rank 广播
    for i in range(tp_size):
        sync_tensor.copy_(tensor_chunk[i])
        dist.broadcast(sync_tensor, src=0, group=mp_group)
        if i == tp_rank and tensor is not None:
            tensor.data.copy_(sync_tensor)

4. gate+up 权重的 TP 切分广播

def _broadcast_tp_shard_tensor_gate_up(tensor, gate_name, up_name):
    """将 gate 和 up 按 TP 交错合并后广播"""
    for i in range(tp_size):
        gate_tp = gate_weight[i * intermediate_tp : (i+1) * intermediate_tp]
        up_tp = up_weight[i * intermediate_tp : (i+1) * intermediate_tp]
        new_gate_up[i*2*intermediate_tp : (i+1)*2*intermediate_tp] = cat([gate_tp, up_tp])

加载流程

1. rank 0 上有完整 state_dict
2. 加载 embedding(按 vocab 维度 TP 切分)
3. 逐层加载 Transformer:
   - input_layernorm.weight: 广播(不切分)
   - qkv_proj.weight: Q+K+V 合并 + TP 切分广播
   - o_proj.weight: 按 dim=1 TP 切分广播
   - post_attention_layernorm.weight: 广播
   - gate_up_proj.weight: gate+up 合并 + TP 切分广播
   - down_proj.weight: 按 dim=1 TP 切分广播
4. 加载 final norm + lm_head
5. DP 内广播: broadcast_params(wrapped_model)

与 deprecated 版本的区别

方面 llama_loader.py (deprecated) llama_loader_depracated.py
权重分发 本地 fetch(单进程) rank 0 广播到所有 rank
适用场景 每个 rank 都有 state_dict 只有 rank 0 有 state_dict
DP 同步 无 最后广播到 DP 组

与其他模块的关系

  • 被 weight_loader_registry.py 注册为 LLaMA 的权重加载器
  • 与 llama_saver.py 互为逆操作

小结

这个加载器实现了从 HF 到 Megatron 分片模型的权重分发。核心挑战是 QKV 合并权重和 gate+up 合并权重的 TP 切分,需要按 TP rank 交错排列以匹配 QKVParallelLinear 和 MergedColumnParallelLinear 的权重布局。