跳转至

registry.py — Megatron-Core 模型注册表

文件路径

verl/models/mcore/registry.py

文件概述

Megatron-Core 子模块的中央注册表,维护了多个字典,将模型架构名称映射到配置转换器、模型初始化器、前向函数和权重转换器。这是使用 mcore 路线时的核心枢纽。

关键代码讲解

SupportedModel 枚举

class SupportedModel(Enum):
    LLAMA = "LlamaForCausalLM"
    QWEN2 = "Qwen2ForCausalLM"
    QWEN2_MOE = "Qwen2MoeForCausalLM"
    DEEPSEEK_V3 = "DeepseekV3ForCausalLM"
    MIXTRAL = "MixtralForCausalLM"
    QWEN3 = "Qwen3ForCausalLM"
    QWEN3_MOE = "Qwen3MoeForCausalLM"
    # ... 更多模型

四大注册表

# 1. 配置转换器:HF config -> mcore TransformerConfig
MODEL_CONFIG_CONVERTER_REGISTRY = {
    SupportedModel.LLAMA: hf_to_mcore_config_dense,
    SupportedModel.QWEN2: hf_to_mcore_config_dense,
    SupportedModel.DEEPSEEK_V3: hf_to_mcore_config_dpskv3,
    SupportedModel.MIXTRAL: hf_to_mcore_config_mixtral,
    # ...
}

# 2. 模型初始化器:创建 mcore GPTModel
MODEL_INITIALIZER_REGISTRY = {
    SupportedModel.LLAMA: DenseModel,
    SupportedModel.QWEN2_MOE: Qwen2MoEModel,
    SupportedModel.DEEPSEEK_V3: DeepseekV3Model,
    # ...
}

# 3. 前向函数:标准 / 融合 / 无 padding 三种
MODEL_FORWARD_REGISTRY = { ... }
MODEL_FORWARD_FUSED_REGISTRY = { ... }
MODEL_FORWARD_NOPAD_REGISTRY = { ... }

# 4. 权重转换器:mcore <-> HF 格式
MODEL_WEIGHT_CONVERTER_REGISTRY = {
    SupportedModel.LLAMA: McoreToHFWeightConverterDense,
    SupportedModel.DEEPSEEK_V3: McoreToHFWeightConverterDpskv3,
    # ...
}

工厂函数

def hf_to_mcore_config(hf_config, dtype, **override_kwargs):
    model = get_supported_model(hf_config.architectures[0])
    return MODEL_CONFIG_CONVERTER_REGISTRY[model](hf_config, dtype, **override_kwargs)

def init_mcore_model(tfconfig, hf_config, pre_process=True, post_process=None, ...):
    model = get_supported_model(hf_config.architectures[0])
    initializer_cls = MODEL_INITIALIZER_REGISTRY[model]
    initializer = initializer_cls(tfconfig, hf_config)
    return initializer.initialize(...)

def get_mcore_weight_converter(hf_config, dtype):
    model = get_supported_model(hf_config.architectures[0])
    return MODEL_WEIGHT_CONVERTER_REGISTRY[model](hf_config, tfconfig)

新版 API(非 deprecated)

def get_mcore_forward_fn(hf_config):
    if hf_config.architectures[0] in supported_vlm:
        return model_forward_gen(True)   # VLM
    else:
        return model_forward_gen(False)  # 语言模型

新版 API 更简洁,根据架构名直接判断是否为 VLM,不再依赖 SupportedModel 枚举。

核心类/函数列表

名称 作用
SupportedModel 支持的模型架构枚举
SupportedVLM 支持的 VLM 架构枚举
hf_to_mcore_config() 配置转换工厂函数
init_mcore_model() 模型初始化工厂函数
get_mcore_weight_converter() 权重转换器工厂函数
get_mcore_forward_fn() 前向函数工厂
get_mcore_forward_fused_fn() 融合前向函数工厂

与其他模块的关系

  • 调用 config_converter.py 的配置转换函数
  • 调用 model_initializer.py 的模型初始化类
  • 调用 model_forward.py 和 model_forward_fused.py 的前向函数
  • 调用 weight_converter.py 的权重转换类
  • 被上层训练代码通过 __init__.py 调用

小结

这是 mcore 子模块的核心枢纽,采用注册表模式将模型架构名映射到各种组件。新版 API 正在简化旧版的 SupportedModel 枚举模式。