【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,按分片参数 × 2);
- all-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(全分片)显存低,FSDP2+LoRA(基座不分片)反而高。
一句话: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_base=True, shard_lora=True) lora_full_base = peak_mem(total, n, shard_base=False, shard_lora=True) print("纯 FSDP2:", pure) print("FSDP2+LoRA(基座不分片):", lora_full_base) assert lora_full_base > pure, "复现:基座不分片导致 LoRA 更费显存" if __name__ == "__main__": main()运行输出:
纯 FSDP2: 750.0 FSDP2+LoRA(基座不分片): 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×总参/N; - LoRA 的 A/B 随所在模块一起被 shard,不额外占完整副本;
- 仅当基座是 4-bit(QLoRA,见第 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_quantized=True) plan_float = LoraMemPlan(1000, 4, base_quantized=False) 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_base=False) sharded = peak(1000, 4, shard_base=True) 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_base=True) 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_shard,m/v 回到 /N; - 第二层:用
LoraMemPlan预算器区分量化/浮点基座,选最优分片方案; - 第三层:pytest 验证浮点基座分片、QLoRA 基座本地完整且省,锁进 CI。
核心教训:LoRA 省的是"参数量",但显存账本里优化器状态(尤其 m/v)才是大头。让基座(无论是否 LoRA)在浮点下不分片,等于把最大的那块 m/v 从 /N 变完整——这是"加 LoRA 反而更费"的最常见根因。量化基座例外,因其体积小且无 m/v。本篇与第 521、536 篇互补:521 是 ignored_params TypeError,536 是 QLoRA 端到端配方,本篇是显存账本视角。