base_model_merger.py — 抽象基类与配置¶
源码路径:
verl/model_merger/base_model_merger.py
文件概述¶
这个文件定义了模型合并器的基础设施,包括:
- 命令行参数解析 (
parse_args) - 配置数据类 (
ModelMergerConfig) - 配置生成函数 (
generate_config_from_args) - 抽象基类 (
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 格式的通用流程则由基类统一处理。这是一个典型的模板方法模式应用。