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 枚举模式。