qwen2_loader_depracated.py — Qwen2 权重加载(本地版)¶
文件路径¶
verl/models/qwen2/megatron/checkpoint_utils/qwen2_loader_depracated.py
文件概述¶
Qwen2 权重加载的旧版实现。假设每个 rank 都有完整的 state_dict,直接在本地进行 TP 切分。
与广播版的区别¶
与 LLaMA 的 deprecated 版本逻辑相同:使用 _fetch_* 函数直接从本地 state_dict 读取并切分,不需要跨 rank 广播。
QKV bias 的本地切分¶
def _fetch_tp_shard_tensor_qkv(tensor, q_name, k_name, v_name, bias=False):
"""本地切分 QKV(支持 weight 和 bias)"""
if not bias:
new_weight_qkv = torch.empty(total_size * tp_size, config.hidden_size, ...)
else:
new_weight_qkv = torch.empty(total_size * tp_size, ...)
# 直接从本地 state_dict 读取并切分
tensor.data.copy_(tensor_chunk[tp_rank])
tie_word_embeddings¶
小结¶
Qwen2 的本地版加载器,适合所有 rank 都能访问完整权重的场景。增加了 QKV bias 和 tie_word_embeddings 的处理。