跳转至

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

if tie_word_embeddings:
    print_rank_0("tie_word_embeddings skip load lm_head")

小结

Qwen2 的本地版加载器,适合所有 rank 都能访问完整权重的场景。增加了 QKV bias 和 tie_word_embeddings 的处理。