跳转至

registry.py — 旧版 Megatron 模型注册表

文件路径

verl/models/registry.py

文件概述

这个文件是 旧版 Megatron-LM 并行模型 的注册表。它维护了一个模型架构名称到具体并行模型类的映射字典,提供 ModelRegistry 类来按需动态加载模型类。

注意:这套旧版注册表正在被 mcore/registry.py 中的新版注册表逐步取代。

关键代码讲解

模型映射字典

_MODELS = {
    "LlamaForCausalLM": (
        "llama",
        ("ParallelLlamaForCausalLMRmPadPP", "ParallelLlamaForValueRmPadPP", "ParallelLlamaForCausalLMRmPad"),
    ),
    "Qwen2ForCausalLM": (
        "qwen2",
        ("ParallelQwen2ForCausalLMRmPadPP", "ParallelQwen2ForValueRmPadPP", "ParallelQwen2ForCausalLMRmPad"),
    ),
    # ... 还有 Mistral 和 Apertus
}

每个条目的结构是: - Key:HuggingFace 的模型架构名称(如 "LlamaForCausalLM") - Value:一个元组 (模块名, (Actor模型类, Value模型类, 无PP模型类))

其中三个模型类分别对应: 1. Actor/Reference 模型(用于生成和计算 log prob) 2. Value/Reward 模型(用于估计状态价值) 3. 无流水线并行版本

ModelRegistry 类

class ModelRegistry:
    @staticmethod
    def load_model_cls(model_arch: str, value=False) -> Optional[type[nn.Module]]:
        module_name, model_cls_name = _MODELS[model_arch]
        if not value:  # actor/ref
            model_cls_name = model_cls_name[0]
        elif value:  # critic/rm
            model_cls_name = model_cls_name[1]

        module = importlib.import_module(f"verl.models.{module_name}.megatron.modeling_{module_name}_megatron")
        return getattr(module, model_cls_name, None)

核心逻辑: 1. 根据 model_arch 查找模块名和类名 2. 根据 value 参数选择 Actor 类还是 Value 类 3. 使用 importlib 动态导入对应模块 4. 返回模型类

核心类/函数列表

名称 作用
_MODELS 模型架构到(模块名, 类名元组)的映射字典
ModelRegistry.load_model_cls() 根据架构名和角色(actor/value)动态加载模型类
ModelRegistry.get_supported_archs() 返回所有支持的架构名列表

与其他模块的关系

  • 被上层的训练代码调用,用于创建旧版 Megatron 并行模型
  • 加载的模型类来自 verl/models/llama/megatron/ 和 verl/models/qwen2/megatron/

小结

这是一个简单的工厂模式注册表,将 HuggingFace 架构名映射到 verl 自定义的 Megatron 并行模型类。正在被 mcore 新架构取代。