ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

【Bug已解决】FSDP MoE PEFT hangs during forward pass 解决方案

【Bug已解决】FSDP MoE PEFT hangs during forward pass 解决方案 【Bug已解决】FSDP MoE PEFT hangs during forward pass 解决方案一、现象长什么样把 MoE混合专家模型接上 PEFT如 LoRA再用 FSDPfully_shard做切分在多卡上跑前向时最诡异的现象是程序毫无报错就静静地卡住GPU 利用率掉到 0nvidia-smi显示进程还在但日志不再前进Ctrl-C 后才看到一堆 NCCL 相关的 stackWatchdog caught collective operation timeout: ... at fsdp/_fully_shard/_fsdp_collectives.py:... RuntimeError: Collective operations must be called from all ranks或者更沉默一些训练直接 hang 在第一个loss model(batch)上永远不返回。这种挂起和上一类报错退出不同——它不崩溃只是停住常常让人误以为是数据加载慢、或者以为还在算白白等上几十分钟。它的触发条件很明确MoE路由导致各 rank token 分布不均 PEFT额外参数层 FSDP集合通信三者叠加且通常只在某个 batch 恰好把某些 expert 的 token 路由成 0 时才出现因此还带偶发特征——前几个 batch 正常某个 batch 突然 hang。二、背景要理解 hang先要理解 FSDP 前向里的集合通信。FSDP2fully_shard在每层前向时会按 shard 把参数all-gather到本卡算完再释放反向时做reduce-scatter。这些集合通信是群体操作collective所有参与 rank 必须成对、同序地调用少一个 rank 调用或调用次数对不上其他 rank 就会无限等待。问题出在 MoE PEFT 上MoE 的 expert 路由token 经 router 后分散到不同 expert每个 rank 上各 expert 拿到的 token 数不同。极端情况某个 rank 当前 batch 没有任何 token 被路由到expert_k。如果某层只在有 token 时才进入 expert 计算而 expert 计算内部又嵌了 FSDP 的集合通信因为 expert 也被fully_shard切分了那么没有 token的 rank 就整层跳过于是它这次 all-gather 少调用一次。其他 rank 调用了它没调用 → NCCL 死锁 → hang。PEFT 的 LoRA 层又添一层复杂LoRA 给原线性层加了lora_A/lora_B两个小矩阵。如果fully_shard只 shard 了原权重没覆盖 LoRA 矩阵比如 PEFT 包装在fully_shard之后才加LoRA 矩阵就会走一条非 FSDP的路径在某些 rank 上不触发通信进一步打乱 collective 的对称性。下面用可运行代码复现集合通信调用次数在 rank 间不对称 → 死锁的核心机制。三、根因根因一句话FSDP 的集合通信要求所有 rank 同序同次调用而 MoE 路由 PEFT 包装让某些 rank 跳过了部分层的通信触发 NCCL 集体操作死锁表现为前向 hang。三个具体失配MoE 路由导致层被条件性跳过某 rank 因 0 token 进入某 expert整层含其 FSDP all-gather不执行通信计数与其他 rank 不对齐。PEFT 层未被fully_shard覆盖LoRA 矩阵走独立路径在部分 rank 上不参与 FSDP 通信破坏对称性。fully_shard调用顺序与 PEFT 包装顺序错配先fully_shard再get_peft_modelLoRA 参数游离在 FSDP 管理之外或反之导致某些参数既被 shard 又被 PEFT 重包装通信路径双份。四、最小可运行复现用一个 2 进程的 Barrier 模拟集合通信rank1 因0 token跳过一次通信复现死锁为不真卡死这里用带超时的 barrier超时即证明不对称import multiprocessing as mp import time # 用清晰版本复现rank1 因 0 token 跳过 expert 通信导致最后集体操作不对称 def _clean_rank(rank, gate, result): try: gate.wait(timeout2) # 第一层通信对齐 tokens 3 if rank 0 else 0 if tokens 0: gate.wait(timeout2) # 仅 rank0 进 expert 通信 gate.wait(timeout2) # 最后层通信rank1 永远等不到 rank0 result[rank] done except Exception as e: result[rank] fHANG/timeout: {type(e).__name__} if __name__ __main__: import threading # noqa mgr mp.Manager() res mgr.dict() # Barrier parties2两边必须都 wait 同一道门相同次数 g mp.Barrier(2) ps [mp.Process(target_clean_rank, args(i, g, res)) for i in range(2)] for p in ps: p.start() for p in ps: p.join(timeout5) print(dict(res)) # 输出会显示 rank1 在最后一道门超时 - 这就是 forward hang 的本质运行后rank1会在最后一道 barrier 超时等价于 FSDP 前向里某 rank 因跳过 expert 通信而永远等不到其他 rank——即 hang。五、解决方案第一层最小直接修复最立竿见影的修复保证每个 rank 无论有没有 token都走完全相同的通信路径。对 MoE把所有 expert 的 all-gather 提前到路由判断之前统一做对 PEFT确保 LoRA 矩阵也纳入 FSDP 管理。import torch import torch.nn as nn import torch.nn.functional as F class SafeMoELayer(nn.Module): 修复版先统一 all-gather 所有 expert 参数再做路由避免某 rank 跳过通信。 def __init__(self, num_experts, hidden): super().__init__() # 每个 expert 是一个独立线性层即便本 rank 没 token也先 materialize/gather self.experts nn.ModuleList( [nn.Linear(hidden, hidden) for _ in range(num_experts)] ) def forward(self, x, router_logits): # x: [n_tokens, hidden]; router_logits: [n_tokens, num_experts] # 关键修复不论本 rank 有无 token都对每个 expert 做一次占位前向 # 这里用 dummy token 保证所有 expert 的 FSDP all-gather 都被触发 present x.shape[0] if present 0: x torch.zeros(1, x.shape[1], devicex.device) # 占位触发通信 out torch.zeros_like(x) for i, expert in enumerate(self.experts): mask router_logits[:, i].sigmoid() 0.5 if mask.any(): out[mask] expert(x[mask]) * router_logits[mask, i].sigmoid().unsqueeze(-1) # 占位 token 的结果丢弃不影响真实梯度 return out[:present] if present 0 else out第一层修复直接消除了0 token rank 跳过 expert 通信的死锁。六、解决方案第二层结构性改进把FSDP 必须覆盖所有可训练参数含 PEFT和MoE 通信路径必须对称收口成一个ShardPlan 包装顺序约定避免以后再出现 PEFT 层游离。import torch import torch.nn as nn from dataclasses import dataclass, field from typing import List dataclass class ShardPolicy: 声明哪些模块需要 fully_shard保证 MoE 与 PEFT 都在内。 shard_modules: List[str] field(default_factorylist) def all_covered(self, model: nn.Module) - bool: names {n.split(.)[0] for n, _ in model.named_modules()} return all(m in names for m in self.shard_modules) def apply_fsdp_then_peft(model, lora_targets, policy: ShardPolicy): 正确顺序先 fully_shard 所有目标模块再叠加 PEFT 且 PEFT 的 LoRA 矩阵也要被同一套 FSDP 管理。 # 1) 对所有 MoE expert 主干做 fully_shard示意真实用 torch.distributed.fsdp for name, mod in model.named_modules(): if any(name.startswith(s) for s in policy.shard_modules): if isinstance(mod, nn.Linear): pass # 真实场景: fully_shard(mod, mesh) # 2) PEFT 包装必须在 fully_shard 之后确保 LoRA 矩阵也被后续 shard 覆盖 # get_peft_model(...) 在此调用 assert policy.all_covered(model), 仍有模块未被 FSDP 覆盖会破坏通信对称 return model def main(): model nn.Sequential( nn.Linear(8, 8), nn.ReLU(), nn.Linear(8, 8), # 第二个线性模拟 expert 层 ) policy ShardPolicy(shard_modules[0, 2]) model apply_fsdp_then_peft(model, [0, 2], policy) print(FSDP PEFT 覆盖校验通过通信路径对称) if __name__ __main__: main()第二层的关键在于顺序契约与覆盖断言PEFT 在 FSDP 之后、且 LoRA 矩阵也被 shard所有 rank 的通信调用次数天然对齐。七、解决方案第三层断言 / CI 守护加 pytest 守护每一层 forward 在所有 rank 上都被调用的不变量并验证占位 token 不污染梯度。用单进程模拟 rank 计数import torch import torch.nn as nn import pytest class CountingLayer(nn.Module): def __init__(self): super().__init__() self.linear nn.Linear(4, 4) self.calls 0 def forward(self, x, active): # 修复后即便 activeFalse 也走一次占位保证计数对齐 if x.shape[0] 0: x torch.zeros(1, 4) out self.linear(x) self.calls 1 return out[:active] if active else out def test_forward_called_even_with_zero_tokens(): layer CountingLayer() # rank0 有 token layer(torch.randn(2, 4), active2) # rank1 模拟 0 token但占位后仍调用一次修复后行为 layer(torch.zeros(0, 4), active0) assert layer.calls 2, 某 rank 跳过了 forward通信将不对称 def test_placeholder_does_not_pollute_grad(): layer CountingLayer() x torch.randn(2, 4, requires_gradTrue) out layer(x, active2) out.sum().backward() assert x.grad is not None assert torch.isfinite(x.grad).all() if __name__ __main__: pytest.main([__file__, -q])CI 里test_forward_called_even_with_zero_tokens通过就能保证0 token rank 不再跳过通信从根上防住 MoEFSDP 的 hang。八、排查清单FSDP MoE PEFT 前向 hang按此顺序排查先确认是集合通信死锁加NCCL_TIMEOUT60环境变量或用TORCH_DISTRIBUTED_DEBUGDETAIL超时后报错会指向fsdp_collectives或watchdog caught collective即为死锁。看是否偶发如果前几个 batch 正常、某 batch 突然 hang高度怀疑 MoE 路由把某些 expert 的 token 路由成 0。检查 MoE expert 是否在所有 rank 上都被调用在每层 expert 前后打印该 rank 的 token 数找某 rank 为 0 且跳过了通信的证据。检查 PEFT 与 FSDP 的包装顺序确认是先fully_shard再get_peft_model且 LoRA 矩阵也被 FSDP 覆盖顺序反了 LoRA 会游离。用占位 token 兜底对 0 token 的 rank喂一个 dummy token 触发 all-gather结果丢弃——低成本消除不对称。统一通信路径把所有 expert 的 all-gather 提到路由判断之前统一做不要有 token 才 gather。降规模复现先把 expert 数、rank 数降到 2用确定性路由复现 hang再放大避免在大集群上盲目等超时。九、小结FSDP MoE PEFT 前向 hang根因不是 GPU 故障而是FSDP 的集合通信要求所有 rank 同序同次调用而 MoE 路由让0 token 的 rank跳过了某层 expert 的 all-gatherPEFT 层若未被 FSDP 覆盖又进一步破坏对称性最终触发 NCCL 集体操作死锁。它不崩溃、只静默卡死且常在某个 batch 路由不均时偶发最易误判为数据慢。修复三层第一层用占位 token 保证即使 0 token 的 rank 也走相同通信路径第二层用ShardPolicy收口FSDP 必须覆盖 MoE 与 PEFT 全部模块、且 PEFT 在 FSDP 之后的顺序契约第三层用 pytest 断言每层 forward 在所有 rank 都被调用来守护对称性。记住集合通信最怕不对称MoE 路由再不均通信路径也要对齐。
返回列表