跳转至

weight_loader_registry.py — 权重加载/保存注册表

文件路径

verl/models/weight_loader_registry.py

文件概述

这个文件提供了两个工厂函数,用于获取将 HuggingFace 格式权重加载到 Megatron 模型(loader)和将 Megatron 分片权重合并保存回 HuggingFace 格式(saver)的函数。它充当一个简单的注册表,根据模型架构名返回对应的转换函数。

关键代码讲解

get_weight_loader -- 获取权重加载器

def get_weight_loader(arch: str):
    from verl.models.mcore.loader import load_state_dict_to_megatron_gptmodel

    _MODEL_WEIGHT_MEGATRON_LOADER_REGISTRY = {
        "LlamaForCausalLM": load_state_dict_to_megatron_gptmodel,
        "Qwen2ForCausalLM": load_state_dict_to_megatron_gptmodel,
    }

    if arch in _MODEL_WEIGHT_MEGATRON_LOADER_REGISTRY:
        return _MODEL_WEIGHT_MEGATRON_LOADER_REGISTRY[arch]
    raise ValueError(...)

目前 Llama 和 Qwen2 都使用同一个加载函数 load_state_dict_to_megatron_gptmodel。

get_weight_saver -- 获取权重保存器

def get_weight_saver(arch: str):
    from verl.models.mcore.saver import (
        merge_megatron_ckpt_gptmodel,
        merge_megatron_ckpt_gptmodel_dpskv3,
        merge_megatron_ckpt_gptmodel_mixtral,
        merge_megatron_ckpt_gptmodel_qwen2_5_vl,
        merge_megatron_ckpt_gptmodel_qwen_moe,
    )

    _MODEL_WEIGHT_MEGATRON_SAVER_REGISTRY = {
        "LlamaForCausalLM": merge_megatron_ckpt_gptmodel,
        "Qwen2ForCausalLM": merge_megatron_ckpt_gptmodel,
        "MixtralForCausalLM": merge_megatron_ckpt_gptmodel_mixtral,
        "DeepseekV3ForCausalLM": merge_megatron_ckpt_gptmodel_dpskv3,
        # ... 更多模型
    }

不同模型架构使用不同的合并函数,因为它们的权重结构不同(如 MoE 模型有专家权重需要特殊处理)。

核心函数列表

函数名 作用
get_weight_loader(arch) 根据架构名返回权重加载函数
get_weight_saver(arch) 根据架构名返回权重保存函数

与其他模块的关系

  • 调用 mcore/loader.py 中的加载函数
  • 调用 mcore/saver.py 中的保存函数
  • 被训练流程中的 checkpoint 管理代码调用

小结

一个轻量级的注册表,将模型架构名映射到对应的权重 loader 和 saver 函数。使用了延迟导入(lazy import)避免不必要的依赖加载。