执行摘要
- 一句话:修复 FSDP 保存检查点时 tied weights 重复导致的加载错误
- 推荐动作:值得精读,尤其是理解 FSDP 与 Hugging Face transformers 检查点互操作的深层坑点。建议阅读
drop_tied_target_keys 实现和 review 讨论中关于 identity 检测与不变式维护的权衡。考虑为核心逻辑添加单元测试。
功能与动机
PR body 中指出:'FSDP gather un-aliases tied params (one independent CPU tensor per state_dict key), so save_pretrained's storage-based dedup never fires for them. Both keys end up in safetensors with distinct values. On reload, transformers>=5 detects the duplicate and silently refuses to tie, defeating tie_word_embeddings=true.'
实现拆解
- 在
verl/utils/transformers_compat.py 中新增 drop_tied_target_keys 核心函数。函数在 tie_word_embeddings=True 时,先调用 model.tie_weights() 确保参数绑定,再通过 named_parameters(remove_duplicate=False) 迭代,利用 id(param) 检测同一参数的不同名称(别名),保留第一个遇到的名称为规范键,删除后续别名键;若规范键不在 state_dict 中,则提升别名键为规范键。
- 在
verl/utils/checkpoint/fsdp_checkpoint_manager.py 的 save_checkpoint 方法中,于 save_pretrained 前调用 _drop_tied_target_keys(state_dict, save_model, model_config)。
- 在
verl/model_merger/base_model_merger.py 的 save_hf_model_and_tokenizer 方法中,于 save_pretrained 前调用同样的函数。
- 两个调用者通过
from verl.utils.transformers_compat import drop_tied_target_keys as _drop_tied_target_keys 导入,以保持向后兼容。
关键文件:
verl/utils/transformers_compat.py(模块 工具函数;类别 source;类型 core-logic;符号 drop_tied_target_keys): 核心逻辑所在:新增 drop_tied_target_keys 函数解决 tied weights 保存问题。
verl/utils/checkpoint/fsdp_checkpoint_manager.py(模块 检查点管理;类别 source;类型 dependency-wiring): FSDP 检查点保存路径中集成调用,确保每次 HF 保存前清理。
verl/model_merger/base_model_merger.py(模块 模型合并;类别 source;类型 data-contract): 模型合并器保存路径中集成调用,保持与 FSDP 路径一致性。
关键符号:drop_tied_target_keys
关键源码片段
verl/utils/transformers_compat.py
核心逻辑所在:新增 drop_tied_target_keys 函数解决 tied weights 保存问题。
# verl/utils/transformers_compat.py
def drop_tied_target_keys(state_dict, model, model_config) -> None:
"""删除 `state_dict` 中绑定的别名键(如 ``lm_head.weight``)。
FSDP gather 和 model merger 会为每个 state_dict 键物化一个独立的 CPU 张量,
使得 ``save_pretrained`` 基于存储指针的去重失效。当两个键都存入 safetensors 后,
``transformers>=5`` 在加载时会静默拒绝重新绑定。本函数通过 ``Parameter`` 身份
在 ``tie_weights()`` 后检测别名。
仅当规范键已确认存在于 ``state_dict`` 时才删除别名;否则将别名提升为规范键,
避免意外删除某个绑定参数的所有条目。
"""
if not getattr(model_config, "tie_word_embeddings", False):
return
model.tie_weights()
canonical_by_id: dict[int, str] = {}
for name, param in model.named_parameters(remove_duplicate=False):
pid = id(param)
if pid not in canonical_by_id:
canonical_by_id[pid] = name
continue
# 遇到别名键
if canonical_by_id[pid] in state_dict:
# 规范键存在,安全删除别名
state_dict.pop(name, None)
else:
# 规范键缺失:提升别名键为规范键,避免数据丢失
canonical_by_id[pid] = name
评论区精华
- gemini-code-assist 指出两处代码重复,要求提取到公共模块,并修复可能完全删除参数的风险。作者在第二次提交中将函数移到
transformers_compat.py,并增强了别名提升逻辑。
-
wuxibin89 问为何不使用直接导入(from ... import get_auto_model_for_vision2seq, drop_tied_target_keys),作者解释为了保持向后兼容,并计划后续 PR 清理(#6356)。
-
重构重复代码到公共模块 (design): 作者在第二次提交中提取到 transformers_compat.py 并共享。
- 至少保留一个键的不变式 (correctness): 作者增加了别名提升逻辑:若规范键缺失,则将别名设为规范键。
- 导入方式选择 (question): 作者解释是为保持向后兼容,并计划在 #6356 中清理。
风险与影响
- 风险:
- 影响所有使用 FSDP 且
tie_word_embeddings=True 的训练检查点保存,但修复的是正确性问题,已通过 end-to-end 验证。
- 新增函数无单元测试,可能漏掉边界情况(如模型自定义
_tied_weights_keys、tie_weights() 抛出异常等)。
- 函数依赖
model.tie_weights(),若模型不支持或返回不一致,可能导致保存失败。
- 当
state_dict 中只出现别名键时,提升逻辑虽能保留参数,但可能改变预期规范键名称,影响下游加载。
- 影响:
- 用户:使用 tied embeddings 的 FSDP 训练将不再出现检查点加载后的参数分化,
tie_word_embeddings 语义得以保证。
- 系统:检查点保存路径增加轻量级预处理,性能影响可忽略。
- 团队:需关注后续 #6356 以完成导入清理,并考虑添加测试覆盖。
- 风险标记:核心路径变更, 缺少测试覆盖, 数据模型兼容性
关联脉络
- PR #6356 #6356 (清理遗留别名导入): PR 作者在讨论中提及的 follow-up PR,用于清理本 PR 遗留的向后兼容导入。
参与讨论