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 的权重布局。