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 新架构取代。