跳转至

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 有权重的分布式场景。