fsdp 默认 use_orig_params=true 要求参数为可训练 nn.parameter,手动冻结后仍被分片致显存激增或报错;须在 fsdp 包装前用 model.train() 和显式 requires_grad 设置,并禁用 auto_wrap_policy 初期调试。

为什么直接用 FSDP 会 OOM 或卡死?
PyTorch 2.4 的 FSDP(Fully Sharded Data Parallel)默认启用 use_orig_params=True,这要求模型参数必须是可训练的、非冻结的 nn.Parameter 对象。如果你沿用旧写法(比如手动 model.requires_grad_(False) 再只对部分层设为 True),FSDP 会在初始化时尝试分片所有参数——包括那些被设为 requires_grad=False 的——导致显存占用反升、甚至触发 RuntimeError: cannot assign a Tensor to parameter 'xxx' before FSDP initialization。
实操建议:
- 微调前先用
model.train()+ 显式for name, p in model.named_parameters(): p.requires_grad = True if "lora" in name or "classifier" in name else False,再传入FSDP - 务必在
FSDP包装前完成所有requires_grad设置,包装后修改会失效 - 禁用
auto_wrap_policy初期调试:直接用fsdp_config = {"sharding_strategy": ShardingStrategy.FULL_SHARD},避免策略误判子模块
torch.compile 和 FSDP 能不能一起用?怎么配才不崩?
可以,但 PyTorch 2.4 中二者协同有严格顺序:必须先 FSDP 包装模型,再对其 forward 方法调用 torch.compile。反序(先 compile 再 FSDP)会导致 FSDP 无法识别参数分片结构,报错 ValueError: compiled module has no _fsdp_wrapped_module。
实操建议:
- 写法必须是:
model = FSDP(model, ...); model.forward = torch.compile(model.forward, mode="reduce-overhead") - 避免对整个
model调用torch.compile(即不用model = torch.compile(model)),否则会绕过FSDP的梯度同步逻辑 - 若用
mode="max-autotune",首次迭代可能卡住 2–3 分钟——这是正常编译开销,不是死锁
LoRA + FSDP 微调时,ignored_modules 怎么设才不漏参?
LoRA 层(如 lora_A/lora_B)通常插入在原始线性层内部,属于子模块。如果只把 LoRA 模块本身传给 ignored_modules,FSDP 仍会对父线性层做分片,导致 LoRA 参数被错误地跨 GPU 拆分或同步异常。
实操建议:
- 把包含 LoRA 的完整父模块(如
model.layers[0].self_attn.q_proj)加入ignored_modules,而不是只加q_proj.lora_A - 用
print(list(model.named_modules()))确认 LoRA 所在层级路径,避免字符串匹配遗漏(例如"q_proj"不等于"self_attn.q_proj") - 若用
peft库,推荐直接传get_peft_model(...).base_model.model给FSDP,并设置ignored_modules=list(model.modules())(仅限 LoRA 全量启用场景)
验证 FSDP 是否真生效:看什么指标最靠谱?
别只盯着 nvidia-smi 显存——FSDP 的核心收益是降低单卡显存峰值和提升吞吐,但显存读数受缓存、CUDA graph 等干扰。真正有效的验证点只有两个:
- 启动后立即检查
torch.cuda.memory_allocated():多卡下每卡应接近均等,且总和显著低于 DDP 或无并行时的单卡峰值 - 训练 step time:开启
torch.profiler,重点看cudaLaunchKernel和ncclGroupEnd占比;FSDP 正常时通信占比应 30% 说明sharding_strategy或backward_prefetch配置不当
容易被忽略的是梯度归约时机:PyTorch 2.4 默认启用 use_orig_params=True 后,optimizer.step() 前必须确保所有分片梯度已同步,否则会静默跳过更新——建议在 loss.backward() 后加一行 torch.cuda.synchronize() 做简单确认。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











