Prhub

#6334 [fsdp, ckpt] fix: drop tied target keys before HF save_pretrained

原始 PR 作者 ChangyiYang 合并时间 2026-05-14 08:29 文件变更 3 提交数 2 评论 10 代码增减 +43 / -0

执行摘要

修复 FSDP 保存检查点时 tied weights 重复导致的加载错误

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.'

值得精读,尤其是理解 FSDP 与 Hugging Face transformers 检查点互操作的深层坑点。建议阅读 drop_tied_target_keys 实现和 review 讨论中关于 identity 检测与不变式维护的权衡。考虑为核心逻辑添加单元测试。

讨论亮点
  • gemini-code-assist 指出两处代码重复,要求提取到公共模块,并修复可能完全删除参数的风险。作者在第二次提交中将函数移到 transformers_compat.py,并增强了别名提升逻辑。
  • wuxibin89 问为何不使用直接导入(from ... import get_auto_model_for_vision2seq, drop_tied_target_keys),作者解释为了保持向后兼容,并计划后续 PR 清理(#6356)。

实现拆解

  1. 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 中,则提升别名键为规范键。
  2. verl/utils/checkpoint/fsdp_checkpoint_manager.pysave_checkpoint 方法中,于 save_pretrained 前调用 _drop_tied_target_keys(state_dict, save_model, model_config)
  3. verl/model_merger/base_model_merger.pysave_hf_model_and_tokenizer 方法中,于 save_pretrained 前调用同样的函数。
  4. 两个调用者通过 from verl.utils.transformers_compat import drop_tied_target_keys as _drop_tied_target_keys 导入,以保持向后兼容。
文件 模块 状态 重要度
verl/utils/transformers_compat.py 工具函数 modified 7.23
verl/utils/checkpoint/fsdp_checkpoint_manager.py 检查点管理 modified 5.63
verl/model_merger/base_model_merger.py 模型合并 modified 5.44

关键符号

drop_tied_target_keys

关键源码片段

verl/utils/transformers_compat.py core-logic

核心逻辑所在:新增 drop_tied_target_keys 函数解决 tied weights 保存问题。

# verl/utils/transformers_compat.pydef 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 指出 `_drop_tied_target_keys` 在 `base_model_merger.py` 和 `fsdp_checkpoint_manager.py` 中重复,建议移到 `transformers_compat.py`。

结论:作者在第二次提交中提取到 `transformers_compat.py` 并共享。 · 已解决

至少保留一个键的不变式 正确性

gemini-code-assist 指出原逻辑可能将参数完全删除(当规范键不在 state_dict 而别名在时),建议确保至少保留一个键。

结论:作者增加了别名提升逻辑:若规范键缺失,则将别名设为规范键。 · 已解决

导入方式选择 question

wuxibin89 问为什么不用 `from ... import get_auto_model_for_vision2seq, drop_tied_target_keys` 而使用别名导入。

结论:作者解释是为保持向后兼容,并计划在 #6356 中清理。 · 已解决

风险与影响

  • 影响所有使用 FSDP 且 tie_word_embeddings=True 的训练检查点保存,但修复的是正确性问题,已通过 end-to-end 验证。
  • 新增函数无单元测试,可能漏掉边界情况(如模型自定义 _tied_weights_keystie_weights() 抛出异常等)。
  • 函数依赖 model.tie_weights(),若模型不支持或返回不一致,可能导致保存失败。
  • state_dict 中只出现别名键时,提升逻辑虽能保留参数,但可能改变预期规范键名称,影响下游加载。
  • 用户:使用 tied embeddings 的 FSDP 训练将不再出现检查点加载后的参数分化,tie_word_embeddings 语义得以保证。
  • 系统:检查点保存路径增加轻量级预处理,性能影响可忽略。
  • 团队:需关注后续 #6356 以完成导入清理,并考虑添加测试覆盖。
核心路径变更 缺少测试覆盖 数据模型兼容性

关联 Issue

#77724 FSDP: enhanced shared parameter support
#85949 [Distributed] Loading distributed checkpoint with FSDP fails with varying key errors (pos.embedding, shared.weight)

完整报告

参与讨论