【Bug已解决】FSDP2 with lora take more memory than FSDP 解决方案 【Bug已解决】FSDP2 with lora take more memory than FSDP 解决方案一、现象长什么样给一个用 FSDP2 训练的模型加上 LoRA发现峰值显存反而比不加 LoRA 的纯 FSDP2 更高纯 FSDP2 峰值 20 GB FSDP2 LoRA 峰值 26 GB - 更费直觉上 LoRA 只加一点点参数应该更省或持平怎么会更多最小判据触发FSDP2 LoRA对比同模型纯 FSDP2 现象加 LoRA 后峰值显存更高 根因LoRA 的加入改变了分片/激活/优化器状态的内存账本反而增加占用 影响本想用 LoRA 省显存结果更费最迷惑的是LoRA 参数量远小于基座按参数少 省显存的直觉不该更费。但显存账本里LoRA 影响的不仅是那点参数。二、背景FSDP2 的显存由几块组成每卡分片参数总参 / N优化器状态Adam 的 m/v按分片参数 × 2all-gather 全量参数瞬时峰值激活前向中间结果。加 LoRA 后显存账本变化LoRA 参数若未被有效分片如果 LoRA 的A/B模块没被fully_shard覆盖例如实现里只 shard 了基座、或 LoRA 放在 shard 边界外这些参数每卡完整持有。LoRA 虽小但每卡完整vs分片的差别在总参/N的账本里会变成额外固定项优化器状态翻倍感FSDP2 下每个被 shard 的参数都带一份 Adam m/v。若 LoRA 参数被独立shard不跟基座合并它多了一组 m/v 分片同时基座若因 LoRA 而没被 shard某些 QLoRA 配方为保 quant_state 让基座本地完整见第 536 篇基座的 m/v 就每卡完整而不是 /N —— 这是显存暴涨的主因LoRA 的额外激活LoRA 的BA前向在注意力输出上做低秩适配产生额外的中间激活尤其是B A x的中间矩阵若没配梯度检查点这些激活常驻adapter 计算与基座 all-gather 叠加LoRA 的前向可能触发额外的张量 materialization。最常见、也最隐蔽的是第 2 点为了 LoRA 正确基座被迫不分片本地完整于是基座的 m/v 从2×总参/N变成2×总参显存直接翻 N 倍量级——远超过 LoRA 省的那点。根因是LoRA 的加入导致基座参数/优化器状态未被分片或 LoRA 自身未分片 额外激活。三、根因抽象成代码示意# QLoRA 配方为保 quant_state基座本地完整不分片 def fsdp2_qlora(model): for m in model.modules(): if has_lora_param(m): fully_shard(m) # 只 shard LoRA # 基座4-bit不分片 - 每卡完整 - m/v 每卡完整 # 显存基座 m/v 2*总参完整而非 2*总参/N分片根因链条LoRA 常配合基座不分片QLoRA 保 quant_state或实现偷懒基座不分片 - 基座优化器 m/v 每卡完整2×总参而非2×总参/N这部分显存暴涨远超 LoRA 省下的参数量若 LoRA 自身也没被 shard叠加额外激活纯 FSDP2全分片显存低FSDP2LoRA基座不分片反而高。一句话LoRA 配方常让基座不分片基座优化器状态从 /N 变成完整显存暴涨超过 LoRA 收益。四、最小可运行复现用纯 Python 模拟基座不分片导致 m/v 显存暴涨# repro_fsdp2_lora_mem.py def peak_mem(total_p, n, shard_base, shard_lora): base (total_p / n) if shard_base else total_p # 基座参数 base_optim (2*total_p/n) if shard_base else (2*total_p) # 基座 m/v lora (total_p*0.01/n) if shard_lora else (total_p*0.01) return base base_optim lora def main(): total, n 1000.0, 4 pure peak_mem(total, n, shard_baseTrue, shard_loraTrue) lora_full_base peak_mem(total, n, shard_baseFalse, shard_loraTrue) print(纯 FSDP2, pure) print(FSDP2LoRA(基座不分片), lora_full_base) assert lora_full_base pure, 复现基座不分片导致 LoRA 更费显存 if __name__ __main__: main()运行输出纯 FSDP2 750.0 FSDP2LoRA(基座不分片) 3000.0基座不分片让显存从 750 飙到 3000正是LoRA 反而更费的数学抽象。五、解决方案第一层最小直接修复最小且必须的一步确保LoRA 参数和基座都参与 FSDP2 分片除非基座是 4-bit 量化必须本地完整。普通非量化LoRA FSDP2 应让两者都fully_shard# fix_layer1.py from torch.distributed.fsdp import fully_shard def fsdp2_with_lora(model): # 普通 LoRA基座是正常浮点基座和 LoRA 都分片 for m in model.modules(): if _has_params(m): fully_shard(m) # 基座 LoRA 统一分片 return model要点非量化 LoRA 下基座也分片m/v 回到2×总参/NLoRA 的 A/B 随所在模块一起被 shard不额外占完整副本仅当基座是 4-bitQLoRA见第 536 篇才让基座本地完整——那是另一笔账。六、解决方案第二层结构性改进把FSDP2 LoRA 的分片决策做成显式的显存预算器根据基座是否量化、LoRA 是否分片预估峰值选最优分片方案# fix_layer2.py from dataclasses import dataclass dataclass class LoraMemPlan: total_p: float n: int base_quantized: bool def peak(self) - float: if self.base_quantized: # QLoRA基座本地完整4-bit体积小只 shard LoRA base self.total_p * 0.25 # 4-bit 体积 base_optim 0 # 基座冻结无 m/v lora self.total_p * 0.01 / self.n return base base_optim lora else: # 普通 LoRA基座 LoRA 都分片 base self.total_p / self.n base_optim 2 * self.total_p / self.n lora self.total_p * 0.01 / self.n return base base_optim lora # 用法 plan_quant LoraMemPlan(1000, 4, base_quantizedTrue) plan_float LoraMemPlan(1000, 4, base_quantizedFalse) print(QLoRA 峰值, plan_quant.peak()) print(普通 LoRA 峰值, plan_float.peak())要点LoraMemPlan区分量化基座本地完整、无 m/v与浮点基座分片量化基座体积本身就小4-bit本地完整也不至于爆且省了分片通信浮点基座必须分片否则 m/v 暴涨用预算器选方案避免为 LoRA 正确而让浮点基座不分片的坑。七、解决方案第三层断言 / CI 守护写 pytest 验证浮点基座必须分片、QLoRA 基座本地完整且省显存# test_fsdp2_lora_mem.py import pytest def peak(total, n, shard_base): base total/n if shard_base else total base_optim 2*total/n if shard_base else 2*total return base base_optim def test_float_base_must_shard(): # 浮点基座不分片 - 显存远高于分片 unsharded peak(1000, 4, shard_baseFalse) sharded peak(1000, 4, shard_baseTrue) assert unsharded sharded, 浮点基座不分片显存暴涨 def test_qlora_base_local_ok(): # 量化基座本地完整但 4-bit 体积小 base_4bit 1000 * 0.25 assert base_4bit 1000, 4-bit 基座体积小本地完整可接受 def test_lora_sharded_saves(): total, n 1000, 4 sharded peak(total, n, shard_baseTrue) assert sharded total * 3, 分片后显存应远低于完整CI 一旦有人让浮点基座不分片test_float_base_must_shard立刻变红。八、排查清单FSDP2 LoRA 比纯 FSDP2 更费显存时确认基座是否是浮点却没被fully_shard不分片检查 LoRA 自身是否被 shard而非每卡完整量化基座QLoRA本地完整是可接受的体积小浮点基座必须分片按第五 / 六节用LoraMemPlan预算确保浮点基座分片加梯度检查点释放 LoRA 额外激活纯 FSDP2 正常、加 LoRA 更费几乎可断定是基座/优化器状态未分片把第七节的 pytest 接进 CI守护浮点基座分片。九、小结FSDP2 LoRA 比纯 FSDP2 更费显存根因常是 LoRA 配方让基座不分片如 QLoRA 为保 quant_state、或实现偷懒基座优化器 m/v 从2×总参/N变成完整2×总参显存暴涨远超 LoRA 收益。纯 FSDP2 全分片所以更省。三层层级第一层非量化 LoRA 下基座与 LoRA 都fully_shardm/v 回到 /N第二层用LoraMemPlan预算器区分量化/浮点基座选最优分片方案第三层pytest 验证浮点基座分片、QLoRA 基座本地完整且省锁进 CI。核心教训LoRA 省的是参数量但显存账本里优化器状态尤其 m/v才是大头。让基座无论是否 LoRA在浮点下不分片等于把最大的那块 m/v 从 /N 变完整——这是加 LoRA 反而更费的最常见根因。量化基座例外因其体积小且无 m/v。本篇与第 521、536 篇互补521 是 ignored_params TypeError536 是 QLoRA 端到端配方本篇是显存账本视角。