llama_loader_depracated.py — LLaMA 权重加载(本地版)¶
文件路径¶
verl/models/llama/megatron/checkpoint_utils/llama_loader_depracated.py
文件概述¶
LLaMA 权重加载的旧版实现。与广播版 llama_loader.py 不同,这个版本假设每个 rank 都已有完整的 state_dict(例如通过共享文件系统),直接在本地进行 TP 切分,无需跨 rank 广播。
关键代码讲解¶
与广播版的核心区别¶
def _fetch_tp_shard_tensor(tensor, name, chunk_dim=0):
"""本地 fetch:直接从 state_dict 切分"""
tp_rank = mpu.get_tensor_model_parallel_rank()
tp_size = mpu.get_tensor_model_parallel_world_size()
full_weight = state_dict[name]
tensor_chunk = torch.chunk(full_weight, tp_size, dim=chunk_dim)
tensor.data.copy_(tensor_chunk[tp_rank])
对比广播版:
def _broadcast_tp_shard_tensor(tensor, name):
"""广播版:rank 0 切分后逐个广播"""
for i in range(tp_size):
dist.broadcast(sync_tensor, src=0, group=mp_group)
QKV 本地切分¶
def _fetch_tp_shard_tensor_qkv(tensor, q_name, k_name, v_name):
"""直接从本地 state_dict 读取 Q/K/V,合并后按 TP 切分"""
full_weight_q = state_dict[q_name]
full_weight_k = state_dict[k_name]
full_weight_v = state_dict[v_name]
# 按 TP 交错合并(逻辑同广播版)
for i in range(tp_size):
new_weight_qkv[i*total_size:(i+1)*total_size] = cat([Q_i, K_i, V_i])
tensor.data.copy_(tensor_chunk[tp_rank])
加载流程¶
与广播版相同的结构,但不需要 dist.broadcast 和 dist.broadcast_object_list 调用。只在需要的 pp_rank 上加载对应层。
使用场景¶
- 适用于所有 rank 都能访问模型权重文件(如 NFS 共享存储)
- 加载速度更快(无网络广播开销)
- 但内存开销更大(每个 rank 都加载完整 state_dict)
小结¶
这是旧版的本地切分加载器,适合所有 rank 可以访问完整权重的场景。新版广播加载器更适合只有 rank 0 有权重的分布式场景。