Prhub

#7324 [ckpt, fsdp] feat: sharded delta weight sync for the TorchTitan engine

原始 PR 作者 attack204 合并时间 2026-08-21 10:15 文件变更 6 提交数 17 评论 4 代码增减 +300 / -76

执行摘要

TorchTitan 引擎接入分片 delta 权重同步,最高提速 17.2×

delta_sharded checkpoint 引擎只同步发生变化的权重元素,要求训练引擎按 rank 导出本地分片与放置元数据;TorchTitan 引擎此前只有全量导出,选择 delta_sharded 会直接抛 NotImplementedError。PR body 用基准数据说明收益:Qwen3-8B 在 32 卡 FSDP2 下 delta_sharded 为 3.34 s 而 nccl 全广播为 34.26 s(10.3×),Qwen3-30B-A3B MoE 在 FSDP2×EP8 下为 8.07 s vs 43.43 s(5.4×)。此外,不修复 EFSDP>1 丢专家的问题,delta 的 seed(全量导出)本身不可信;不修复 HSDP 下的 mesh bug,任何 HSDP 运行都会在 worker 初始化阶段崩溃。

值得精读。三个设计决策尤其值得学习:(1) 用数学分析而非特判接纳 _StridedShard——右到左切割只置换 block 归属而非交错 block,compute_local_shape_and_global_offset 已返回置换后偏移;(2) 用 slot table 而非 probe-driven 转换器解决 lockstep gather 下的命名一致性问题;(3) gather 组的「变化 Shard 维、固定其他维」构造避免 HSDP 副本浪费,并按 rank 集合缓存保证 dist.new_group 的建组顺序一致。若要在 verl 中为新的 trainer 引擎接入 delta 同步,本 PR 是很好的实现模板。

讨论亮点

外部 review 只有 wuxibin89 的 APPROVED,两条 review 评论均为作者 attack204 的自评说明。最有价值的反馈藏在 commit message 里:ec0f35b「drop the torchtitan delta export test」明确写道 “Review feedback. It needed 8 GPUs and a torchtitan install, so no CI job could have picked it up”——8 GPU 加 torchtitan 依赖使引擎级导出测试无法被任何 CI 任务拾取,最终删除该测试,保留直接覆盖 placement 数学的 test_sharded_delta_gather。作者在 spec.py 第 211 行的自评指出 “FSDP2 + TP is a _StridedShard, so we should modify the limitation”,即旧代码把 _StridedShard 当作「非单 block」拒绝,而实际上它只是置换 rank 归属的 block,由此驱动了 _shard_dim() 的引入与文档支持矩阵的同步更新。

实现拆解

1. 新增分片导出入口(transformer_impl.py)

  • TorchTitanEngine.get_per_tensor_param_shard() 遍历 self.modulestate_dict(),通过内层 _shard() 完成「搬到 device → 转 bf16 → to_local() → 展平」,产出 (hf_name, local_flat_bf16, ShardSpec) 三元组;ShardSpec.from_param() 记录 mesh 与 placements,接收端据此推导全局 flat 坐标。
  • get_per_tensor_param_delta_shard() 借助 verl.workers.engine.utils.hf_delta_export 把本次分片与 _delta_shard_snap 快照做逐字节差分,产出最终 HF 坐标的稀疏 delta 条目;快照每次导出时刷新,任何 rank 都不持有全量模型快照。
  • 默认实现上移到 BaseEngine(只需分片导出加 per-parameter entry builder),删除 FSDPEngine 中两处完全相同的拷贝;Megatron 保留自己的 override(差分对象不是固定快照),veomni 保留 _hf_delta_entry,TorchTitanEngine 自带实现以处理专家堆栈的 slot 拆分。

2. HF 命名对齐与专家堆栈

  • _to_hf_named_params() 把原 get_per_tensor_param 里的 sd_adapter.to_hf()lm_head.weight 补回逻辑抽成共享方法,保证全量与分片导出按名字精确配对——名字只出现在一边会静默把权重冻结在 seed 值。
  • _expert_stack_slots() 利用 sd_adapter.from_hf_map 结构性识别 fused expert stack(一个 torchtitan key 对应带两个 {} 占位的 HF 模板),整体导出堆栈并在 spec.hf_slots 中枚举每个专家的 HF 张量,避免 to_hf() 只产出本地专家导致的 lockstep 命名不一致。
  • _hf_entry_row_slots() 处理 dim-0 identity 拆分:行 e 就是 full[e],用一次 searchsorted 从全局索引恢复每个专家的计数,规避 probe-driven 转换器的复杂度。

3. 放置推导扩展(spec.py)

  • 新增 _shard_dim():torch 2.13 中 _StridedShard 迁到 C++ 且不再继承 Shardp.is_shard() 会漏判;该函数同时对两种写法给出答案。
  • 新增 _assert_even_strided():拒绝 strided 偏移公式无法表达的 uneven cut。
  • 新增 _shard_dims_group() 取代 _flattened_mesh_group():变化 Shard 维、固定其他维构造 gather 组,按 rank 集合缓存(dist.new_group 是集合通信,所有 rank 以相同顺序建组),避免 HSDP 下 replicate 度被计入 wire 字节与 rank 0 的 staging 缓冲。

4. 边界防护与两个预存 bug 修复

  • _assert_shard_export_supported() 在导出边界以明确信息拒绝 PP:各 stage 持有不相交切片,导出顺序在 rank 间不一致,无法满足 lockstep gather。
  • _get_data_parallel_mesh() 改用 parallel_dims.get_optional_mesh("loss") 获取 1-D 扁平 loss mesh,修复 HSDP 下 get_data_parallel_rank()/get_data_parallel_group() 遇到 2-D mesh 崩溃的问题。
  • get_per_tensor_param()(全量导出)改为先对整个 DTensor 专家堆栈做 full_tensor() 再拆分,修复 EFSDP>1 时先 to_hf() 拆分再在 ep 组内 all-gather 导致静默丢失 50% 专家的 bug。

5. 测试、兼容与文档配套

  • tests/special_distributed/test_sharded_delta_gather.py 扩展 2D/3D mesh 用例:FSDP×TP(_StridedShard)、EFSDP×EP、HSDP×TP/EP、replicate 度 4 的 sharp case,并新增断言「gather 组大小等于 Shard 维大小的乘积」。
  • sglang_rollout.pyLocalSerializedTensor 的 import 增加 SGLang 0.5.16+ 的兼容分支。
  • docs/advance/delta_weight_sync.mddocs/workers/torchtitan_workers.rst 更新支持矩阵、设计说明与实测性能表。
文件 模块 状态 重要度
verl/workers/engine/torchtitan/transformer_impl.py 训练引擎 modified 8.95
verl/workers/engine/spec.py 放置推导 modified 8.51
verl/workers/rollout/sglang_rollout/sglang_rollout.py 回放端 modified 5.52
tests/special_distributed/test_sharded_delta_gather.py 测试覆盖 modified 5.79
docs/advance/delta_weight_sync.md 文档 modified 3.01
docs/workers/torchtitan_workers.rst 文档 modified 2.02

关键符号

get_per_tensor_param_shard get_per_tensor_param_delta_shard _to_hf_named_params _expert_stack_slots _hf_entry_row_slots _assert_shard_export_supported _shard_dim _assert_even_strided _shard_dims_group derive_dtensor_placement

关键源码片段

verl/workers/engine/torchtitan/transformer_impl.py core-logic

核心实现文件:新增分片导出与 delta 差分入口,处理 HF 命名对齐与专家堆栈 slot table,同时修复 EFSDP 丢专家与 HSDP mesh 两个预存 bug

def get_per_tensor_param_shard(self, **kwargs):
    """Yield this rank's local shard ``(hf_name, local_flat_bf16, ShardSpec)`` instead of the full tensor."""
    self._assert_shard_export_supported()
    raw = {}
    for module in self.module:
        raw.update(module.state_dict())
​
    # 专家堆栈必须整体导出并携带 slot table:to_hf() 只会保留本地拥有的专家,
    # 而 lockstep gather 要求所有 rank 对“当前参数是哪个”给出同一答案
    stacks = {}
    for name, param in raw.items():
        slots = self._expert_stack_slots(name, param)
        if slots is not None:
            stacks[name] = slots
    params = self._to_hf_named_params({k: v for k, v in raw.items() if k not in stacks})
    device = get_device_id() # 开启 offload 策略时本地分片位于 CPU
​
    from ..spec import ShardSpec
​
    def _shard(param):
        # 先把参数搬到 device 再转 bf16:混合精度下 fp32 主权重保留在优化器,
        # 线路上只传输 bf16 分片,差分时才做类型对齐
        p = param.to(device, non_blocking=True)
        if p.is_floating_point():
            p = p.to(torch.bfloat16, non_blocking=True)
        local = p.to_local() if isinstance(p, DTensor) else p
        return local.reshape(-1)
​
    def _gen():
        # 稠密参数:按 HF 名字逐个导出,ShardSpec 携带 mesh 与 placements,
        # 供接收端 derive_dtensor_placement 推导 flat 坐标
        for name, param in params.items():
            yield name, _shard(param), ShardSpec.from_param(param)
        # 专家堆栈:以 torchtitan 原始名字为 key,spec.hf_slots 记录
        # 每个专家的 HF 张量名与形状,接收端按 slot 拆分为逐专家 delta
        for name, slots in stacks.items():
            spec = ShardSpec.from_param(raw[name])
            spec.hf_slots = slots
            yield name, _shard(raw[name]), spec
​
    return _gen(), None
verl/workers/engine/spec.py core-logic

共享放置推导逻辑:_shard_dim / _assert_even_strided / _shard_dims_group 使 _StridedShard 与 HSDP 多 Shard 维成为一等公民,是所有 delta 后端的共同基础设施

def _shard_dim(p) -> Optional[int]:
    """返回该 placement 切割的张量维,不切割则返回 None。    torch 2.13 中 _StridedShard 已迁到 C++ 且不再继承 Shard,p.is_shard() 会漏判;
    _StridedShard 与 Shard 一样是单 block 切割,只是右到左的切割顺序置换 rank 与 block 的归属,而不是交错 block 内部。
    """
    if p.is_shard():
        return int(p.dim)
    return int(p.dim) if type(p).__name__ == "_StridedShard" else None
​
​
def derive_dtensor_placement(spec: ShardSpec) -> tuple[int | BlockPlacement, bool, Optional[ProcessGroup]]:
    """推导本 rank 的 (place, contributes, gather_group) 三元组。    place 供 translate_flat_indices 做分片内 flat 坐标到全量坐标的映射;
    contributes 标记副本 rank(Replicate 维坐标非 0)以空 delta 参与 lockstep;
    gather_group 决定向哪些 rank 收集分片。
    """
    import torch.distributed as dist
​
    assert spec.place is None # 显式 place 的 spec 由调用方分发,不会走到这里
​
    if spec.mesh is None:
        # 未分片:本地张量就是完整参数,无需 gather
        return 0, (dist.get_rank() == 0 if dist.is_initialized() else True), None
​
    placements = spec.placements
    # 用 _shard_dim 而非 p.is_shard(),否则 torch 2.13 下 _StridedShard 会被静默漏进
    # “无 shard 维”分支,mesh 维从 gather 组中消失
    shard_dims = [d for d, p in enumerate(placements) if _shard_dim(p) is not None]
    _assert_even_strided(spec, placements)
​
    coord = spec.mesh.get_coordinate()
    contributes = True
    if coord is not None:
        # 任一 Replicate 维坐标非 0 的 rank 不贡献数据,只以空 delta 保持 lockstep
        for d, p in enumerate(placements):
            if p.is_replicate() and coord[d] != 0:
                contributes = False
                break
​
    if not shard_dims:
        return 0, contributes, None
​
    local_shape, global_offset = compute_local_shape_and_global_offset(spec.full_shape, spec.mesh, list(placements))
    place = BlockPlacement(tuple(local_shape), tuple(global_offset), tuple(spec.full_shape))
​
    if len(shard_dims) == 1:
        # 单 Shard 维:该维子组天然只含贡献 rank,副本被跳过
        return place, contributes, spec.mesh.get_group(mesh_dim=shard_dims[0])
    # 多 Shard 维:_shard_dims_group 变化 Shard 维、固定其他维,
    # 避免 HSDP 下 replicate 度被计入 wire 字节与 rank 0 的 staging 缓冲
    return place, contributes, _shard_dims_group(spec.mesh, shard_dims)

评论区精华

FSDP2+TP 的 _StridedShard 支持与 limitation 更新 设计

作者在 spec.py 第 211 行附近自评指出 “FSDP2 + TP is a _StridedShard, so we should modify the limitation”:旧代码把 _StridedShard 当作「非单 block」拒绝,而实际上右到左切割只置换 rank 与 block 的归属,compute_local_shape_and_global_offset 已返回置换后偏移

结论:引入 _shard_dim() 接纳 _StridedShard,并同步更新 delta_weight_sync.md 与 torchtitan_workers.rst 的支持矩阵 · 已解决

8 GPU 引擎级导出测试的可运行性 测试

commit ec0f35b 记录 “Review feedback. It needed 8 GPUs and a torchtitan install, so no CI job could have picked it up”——引擎级导出测试需要 8 GPU 且依赖 torchtitan 安装,任何 CI job 都无法拾取

结论:删除 torchtitan delta 导出测试,保留直接覆盖 placement 数学的 test_sharded_delta_gather;引擎级覆盖依赖后续 CI 基建 · 已解决

_hf_entry_row_slots 的专家堆栈命名转换语义 question

作者在 utils.py 自评解释 TrainSide fused stack(shape 为 (num_experts, *, *))如何映射为 rollout 侧的 model.layers.L.mlp.experts.0.w1 直至 experts.127.w1 逐专家张量

结论:纯解释性评论,无代码变更;转换语义由 dim-0 identity 的 searchsorted 拆分解法确定 · 已解决

风险与影响

  1. torch 版本敏感_shard_dim()_assert_even_strided() 的语义依赖 torch 2.13 中 _StridedShard 的行为;同时 torchtitan 引擎本身需要 torch≥2.12(DataParallelMeshDims),而 SGLang ≤0.5.17 锁定 torch 2.11,仅在 SGLang main(torch 2.13)上收敛,环境组合较脆弱。
  2. 引擎级 e2e 测试缺失:删除 8 GPU 导出测试后,torchtitan 引擎的分片导出路径没有 CI 覆盖,只有 placement 数学的单元测试;PR body 也注明 HSDP+EP“unit-tested; no end-to-end row”。
  3. 依赖上游修复:HSDP×EP 需要 torchtitan 侧 MoE state dict adapter 修复(to_hf() 在读 placement.dim 前未检查类型),修复前该组合在 dcp_load 阶段即崩溃。
  4. 共享逻辑影响面spec.pyderive_dtensor_placement() 是 FSDP/veomni/Megatron 共享的放置推导,本次改动为向后兼容的扩展(不再拒绝 _StridedShard),但任何回归都会波及其他 delta 后端。

对 torchtitan 引擎用户,delta_sharded 从不可用到可用,跨节点 32 卡场景权重同步从 34 s 级降到 3 s 级,MoE 模型从 43 s 级降到 8 s 级;HSDP 用户额外获得 _get_data_parallel_mesh() 修复(此前任何 HSDP 运行都会在 worker 初始化崩溃)。对 nccl 全量同步用户,EFSDP>1 时全量导出静默丢失专家的预存 bug 被修复,属于正确性提升。对团队而言,17 个 commit 完整展示了从「拒绝 TP/EP」到逐步放开的设计演进、负控制验证与大量实测,文档沉淀了支持矩阵与基准数据,为后续其他引擎接入 delta 同步提供了可复用的模板。

依赖 torch 2.13 行为 引擎级 e2e 测试缺失 HSDP×EP 依赖上游修复 共享放置推导逻辑改动

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论