跳转至

base_model_merger.py — 抽象基类与配置

源码路径:verl/model_merger/base_model_merger.py

文件概述

这个文件定义了模型合并器的基础设施,包括:

  1. 命令行参数解析 (parse_args)
  2. 配置数据类 (ModelMergerConfig)
  3. 配置生成函数 (generate_config_from_args)
  4. 抽象基类 (BaseModelMerger),提供公共方法供 FSDP 和 Megatron 两种合并器继承

关键代码讲解

1. 命令行参数解析 parse_args()

def parse_args():
    parser = argparse.ArgumentParser(description="verl model merger")
    subparsers = parser.add_subparsers(dest="operation", required=True,
                                        help="Specify 'merge' or 'test' operation.")

    base_op_parser = argparse.ArgumentParser(add_help=False)
    base_op_parser.add_argument("--backend", type=str, required=True,
                                 choices=["fsdp", "megatron"])
    base_op_parser.add_argument("--local_dir", type=str, default=None)
    base_op_parser.add_argument("--tie-word-embedding", action="store_true")
    base_op_parser.add_argument("--trust-remote-code", action="store_true")
    base_op_parser.add_argument("--is-value-model", action="store_true")
    base_op_parser.add_argument("--use_cpu_initialization", action="store_true")
    # ... merge_parser 和 test_parser 分别继承 base_op_parser

这里使用了 argparse 的 子命令模式(subparsers),支持两种操作: - merge:合并 checkpoint 并保存为 HuggingFace 格式 - test:将合并结果与参考 HuggingFace 模型做对比验证

base_op_parser 是公共参数的父解析器,merge 和 test 子命令各自添加自己特有的参数。

2. 配置数据类 ModelMergerConfig

@dataclass
class ModelMergerConfig:
    operation: str          # 'merge' 或 'test'
    backend: str            # 'fsdp' 或 'megatron'
    target_dir: Optional[str] = "tmp"         # 输出目录
    hf_upload_path: Optional[str] = None      # HuggingFace Hub 仓库 ID
    private: bool = False                      # 是否上传为私有仓库
    test_hf_dir: Optional[str] = None         # 测试用的参考模型路径
    tie_word_embedding: bool = False           # 是否共享词嵌入权重
    trust_remote_code: bool = False            # 是否信任远程代码
    is_value_model: bool = False               # 是否为价值模型
    local_dir: Optional[str] = None           # checkpoint 路径
    hf_model_config_path: Optional[str] = None # HF 模型配置路径
    hf_upload: bool = field(init=False)        # 自动计算,不由用户传入
    use_cpu_initialization: bool = False       # 是否用 CPU 初始化

    def __post_init__(self):
        self.hf_upload = self.operation == "merge" and bool(self.hf_upload_path)
        if self.operation == "test":
            self.target_dir = None
            self.hf_upload_path = None
            self.private = False

重点说明: - hf_upload 是通过 __post_init__ 自动计算的,不需要手动设置 - 如果 operation 为 test,会自动清空不相关的字段 - tie_word_embedding:在某些模型中(如 Qwen),输入的词嵌入层和输出的 lm_head 共享同一组权重

3. 抽象基类 BaseModelMerger

class BaseModelMerger(ABC):
    def __init__(self, config: ModelMergerConfig):
        self.config = config
        self.hf_model_config_path = config.hf_model_config_path
        self.model_config = AutoConfig.from_pretrained(
            self.hf_model_config_path,
            trust_remote_code=self.config.trust_remote_code
        )

初始化时从 checkpoint 目录下的 huggingface/ 子目录加载模型配置。verl 在保存 checkpoint 时会同时保存原始的 HuggingFace 配置文件。

4. 自动模型类选择 get_transformers_auto_model_class()

def get_transformers_auto_model_class(self):
    has_remote_code = hasattr(self.model_config, "auto_map") and any(
        self.model_config.architectures[0] in val
        for val in self.model_config.auto_map.values()
    )
    if has_remote_code:
        auto_class = next(
            k for k, v in self.model_config.auto_map.items()
            if self.model_config.architectures[0] in v
        )
        match auto_class:
            case "AutoModelForCausalLM":
                return AutoModelForCausalLM
            case "AutoModelForTokenClassification":
                return AutoModelForTokenClassification
            case "AutoModelForVision2Seq":
                return AutoModelForVision2Seq
            # ...

这个方法根据模型配置自动决定使用哪个 AutoModel 类。支持: - AutoModelForCausalLM:标准因果语言模型(GPT 系列) - AutoModelForTokenClassification:Token 分类模型(价值模型) - AutoModelForVision2Seq:视觉-语言模型

5. LoRA 适配器保存 save_lora_adapter()

def save_lora_adapter(self, state_dict: dict[str, torch.Tensor]):
    lora_params_names = [name for name in state_dict.keys() if "lora_" in name]
    if len(lora_params_names) == 0:
        return None

    lora_params = OrderedDict()
    target_modules = set()
    for name in lora_params_names:
        lora_key = name.replace(".default.weight", ".weight")
        target_modules.add(lora_key.split(".")[-3])
        lora_params[lora_key] = state_dict.pop(name)

    lora_rank = min(lora_params[lora_key].shape[0], lora_params[lora_key].shape[1])
    # ... 保存 adapter_config.json 和 adapter_model.safetensors

如果 state_dict 中包含 LoRA 参数(名称中含有 lora_),这个方法会: 1. 提取所有 LoRA 参数 2. 自动推断 LoRA rank 和 target_modules 3. 生成 PEFT 配置文件并保存为 safetensors 格式 4. 将剩余的 base model 参数名恢复为标准格式

注意:lora_alpha 被设为 0,需要用户手动设置正确的值。

6. 保存 HuggingFace 模型 save_hf_model_and_tokenizer()

def save_hf_model_and_tokenizer(self, state_dict: dict[str, torch.Tensor]):
    auto_model_class = self.get_transformers_auto_model_class()
    with init_empty_weights():
        model = auto_model_class.from_config(
            self.model_config, torch_dtype=torch.bfloat16,
            trust_remote_code=self.config.trust_remote_code
        )
    model.to_empty(device="cpu")
    model = self.patch_model_generation_config(model)
    lora_path = self.save_lora_adapter(state_dict)
    model.save_pretrained(self.config.target_dir, state_dict=state_dict)

这个方法的流程是: 1. 用 init_empty_weights() 创建空模型(不分配实际内存) 2. 将模型移到 CPU 3. 修补 generation_config 4. 如果有 LoRA 参数则单独保存 5. 使用 save_pretrained 保存模型和 tokenizer

7. HuggingFace Hub 上传 upload_to_huggingface()

def upload_to_huggingface(self):
    api = HfApi()
    api.create_repo(repo_id=self.config.hf_upload_path,
                    private=self.config.private, exist_ok=True)
    api.upload_folder(folder_path=self.config.target_dir,
                      repo_id=self.config.hf_upload_path, repo_type="model")

支持将合并后的模型直接上传到 HuggingFace Hub,包含完善的错误处理(认证失败、网络中断等)。

核心类/函数列表

名称 类型 说明
parse_args() 函数 命令行参数解析
ModelMergerConfig 数据类 合并器配置
generate_config_from_args() 函数 从参数生成配置对象
BaseModelMerger 抽象类 合并器基类
BaseModelMerger.get_transformers_auto_model_class() 方法 自动选择模型类
BaseModelMerger.patch_model_generation_config() 方法 修补生成配置
BaseModelMerger.save_lora_adapter() 方法 保存 LoRA 适配器
BaseModelMerger.save_hf_model_and_tokenizer() 方法 保存 HF 模型和 tokenizer
BaseModelMerger.upload_to_huggingface() 方法 上传到 HF Hub
BaseModelMerger.merge_and_save() 抽象方法 子类必须实现的合并逻辑
BaseModelMerger.cleanup() 抽象方法 子类必须实现的清理逻辑

与其他模块的关系

  • verl/utils/hf_tokenizer, hf_processor - 加载 tokenizer 和 processor
  • transformers - 使用 AutoConfig、AutoModel 等类
  • accelerate - 使用 init_empty_weights 节省内存
  • peft - 处理 LoRA 适配器保存
  • huggingface_hub - 上传模型到 HF Hub
  • 被 FSDPModelMerger 和 MegatronModelMerger 继承

小结

base_model_merger.py 是整个 model_merger 模块的基础设施层。它定义了命令行接口、配置管理和公共的模型保存逻辑。两个子类(FSDP / Megatron)只需专注于各自的 checkpoint 加载和合并逻辑,而保存为 HuggingFace 格式的通用流程则由基类统一处理。这是一个典型的模板方法模式应用。