【Bug已解决】[Bug]: [FSDP2] auto-exclude incompatible Params4bit from fully_shard to prevent silent QLoRA

发布时间:2026/8/1 17:51:01
【Bug已解决】[Bug]: [FSDP2] auto-exclude incompatible Params4bit from fully_shard to prevent silent QLoRA 【Bug已解决】[Bug] [FSDP2] auto-exclude incompatible Params4bit from fully_shard to prevent silent QLoRA corruption 解决方案一、现象长什么样用FSDP2torch.distributed.fsdp.fully_shard做QLoRA4-bit 量化基座 LoRA训练多卡下出现诡异结果LoRA 权重训完合并回基座后输出乱码 / 精度远低于单卡 QLoRA。只在多卡 FSDP2fully_shard时炸单卡 QLoRA 正常或不 FSDP仅 DDP也正常。没有报错是「4-bit 基座的梯度被悄悄算错」导致训练污染。有时伴随RuntimeError: ... quantization scale mismatch但更多是静默错。本质FSDP2 的fully_shard(model)默认把模型所有参数**含 QLoRA 的 4-bit 量化基座Params4bit都分片。但 4-bit 量化参数有特殊的反量化dequant路径和量化 scale不能被普通的分片 all-reduce 梯度处理——分片后每个 rank 只拿到 4-bit 权重的一个碎片反量化所需的全局 scale 上下文被破坏梯度在 4-bit 上做 all-reduce 毫无意义于是基座本应冻结被悄悄改坏、LoRA 训练被污染。**二、背景QLoRA 的做法基座用 4-bit 量化bitsandbytes的Params4bit冻结只在上面挂可训练的 LoRAfloat。训练时只有 LoRA 的梯度流动基座不动。FSDP2fully_shard的做法把模型参数按张量分片到各 rank反向时各 rank 的梯度做 all-reduce 聚合。它假设「每个参数都是普通 float 张量可切分、可聚合梯度」。冲突点4-bitParams4bit不是普通 float——它是「量化值 scale」的打包表示且基座被冻结requires_gradFalse。FSDP2 却把它当成普通参数去fully_shardfully_shard会对它做分片把一个 4-bit 权重张量沿某维切成 N 份每份丢掉全局 scale 的上下文 → 反量化出错。即便基座requires_gradFalseFSDP2 理论上不该聚合其梯度分片本身已经破坏了 4-bit 权的数据布局若基座在训练中被任何路径触碰如zero_grad误清、或混合精度 cast就静默损坏。更糟的是有些实现里fully_shard对requires_gradFalse的参数仍做分片为了内存于是 4-bit 基座被分片且无法正确反量化 →静默 QLoRA 损坏。一句话FSDP2 把 QLoRA 的 4-bit 基座当成普通参数分片破坏了其量化反量化上下文导致基座被静默改坏、LoRA 训练污染。三、根因根因是FSDP2fully_shard未识别并排除 QLoRA 的 4-bit 量化参数把它们当普通参数分片破坏量化语义三层第一层主因4-bitParams4bit被fully_shard分片。fully_shard(model)遍历所有参数做分片没判断「这是不是量化参数」。4-bit 参数一被切分反量化所需的 scale/零点的全局性被破坏权重值错。第二层4-bit 梯度聚合无意义且危险。即便基座冻结FSDP2 的分片/通信逻辑可能仍对 4-bit 参数做 shape 相关的处理若基座在混合精度下被 cast 或误参与梯度4-bit 上的 all-reduce 既无意义又可能写坏量化缓冲区。第三层无「量化参数自动排除」机制。FSDP2 没有「遇到Params4bit等量化参数自动不 shard、保留在单一 rank / 或整体复制」的策略也没有报错提示于是静默损坏。一句话fully_shard 不分青红皂白分片所有参数、含 4-bit 量化参数破坏其语义且无自动排除导致静默 QLoRA 损坏。四、最小可运行复现下面用纯 Python 模拟「4-bit 量化参数被分片后反量化失败、导致值错」的控制流不需要 GPUfrom dataclasses import dataclass from typing import List dataclass class Param: name: str is_4bit: bool value: float scale: float 1.0 def fully_shard_buggy(params: List[Param], world: int): 有 bug对所有参数含 4bit都分片。 for p in params: if p.is_4bit: # 错误4bit 被切分scale 上下文丢失 - 反量化值错 p.value p.value / world 0.5 # 模拟分片破坏 return params def dequant(p: Param) - float: # 4bit 反量化依赖全局 scale return p.value * p.scale def main(): params [ Param(base.4bit, is_4bitTrue, value2.0, scale0.25), Param(lora.A, is_4bitFalse, value1.0), ] # 单卡4bit 不分片反量化正确 base_single dequant(params[0]) # 2.0 * 0.25 0.5 # 多卡 fully_shard错误分片 4bit fully_shard_buggy(params, world4) base_sharded dequant(params[0]) # 被破坏后 ! 0.5 print(单卡反量化值:, base_single) print(多卡分片后反量化值:, base_sharded, (应相同实际错 - 静默损坏)) if __name__ __main__: main()跑出来单卡 0.5、多卡分片后值变损坏和线上「4-bit 基座被分片悄悄改坏」一致。五、解决方案第一层最小直接修复最省事的救火只对 LoRAfloat参数做fully_shard把 4-bit 基座排除在外。FSDP2 支持「只 shard 模型的一部分」from torch.distributed.fsdp import fully_shard import torch # QLoRA基座是 4bitParams4bitLoRA 是 float # 只对 LoRA 参数所在模块做 fully_shard基座保持完整不分片 # 方法找出所有非 4bit 的子模块逐模块 fully_shard for module in model.modules(): # 跳过包含 Params4bit 的基座层 has_4bit any(getattr(p, quant_state, None) is not None for p in module.parameters(recurseFalse)) if not has_4bit and any(True for _ in module.parameters(recurseFalse)): fully_shard(module, meshmesh)或者更直接的把 4-bit 基座requires_gradFalse且整体放在 rank0不分片只 shard LoRA# 仅对 LoRA 参数进行 shard lora_params [p for n, p in model.named_parameters() if lora in n] for m in lora_modules: fully_shard(m, meshmesh) # 只 shard LoRA 模块这样 4-bit 基座不被分片量化语义完好。六、解决方案第二层结构性改进第一层是「手动挑 LoRA 模块」第二层是「实现 FSDP2 的自动排除检测Params4bit量化参数并跳过其分片」从设计上消灭误分片from typing import List def is_quantized_param(p) - bool: 检测是否为 4bit 量化参数Params4bit。 # bitsandbytes Params4bit 带 quant_state 属性 return hasattr(p, quant_state) and p.quant_state is not None def collect_shardable_modules(model, mesh): 自动排除含量化参数的模块只对纯 float 模块 fully_shard。 shard_targets [] for name, module in model.named_modules(): params list(module.parameters(recurseFalse)) if not params: continue # 含量化参数 - 排除不分片保留量化语义 if any(is_quantized_param(p) for p in params): continue # 纯 float 且有可训练参数 - 分片 if any(p.requires_grad for p in params): shard_targets.append(module) return shard_targets def fully_shard_qlora_safe(model, mesh): QLoRA FSDP2 安全分片自动排除 4bit 基座。 targets collect_shardable_modules(model, mesh) for m in targets: fully_shard(m, meshmesh) # 只 shard LoRA / float 部分 return model # 用法 model load_qlora_model(...) # 4bit 基座 LoRA fully_shard_qlora_safe(model, mesh) # 4bit 基座自动排除不静默损坏关键改动is_quantized_param识别Params4bit有quant_state。collect_shardable_modules跳过任何含量化参数的模块只 shard 纯 float 模块。fully_shard_qlora_safe把「排除量化参数」做成默认行为用户不必手动挑模块。七、解决方案第三层断言 / CI 守护把「4bit 自动排除」「不分片」「不静默损坏」固化成测试import pytest def test_detect_quantized_param(): p4 FakeParam(quant_statex) # 4bit pfp FakeParam(quant_stateNone) # float assert is_quantized_param(p4) is True assert is_quantized_param(pfp) is False def test_quantized_module_excluded(): model FakeQLoRAModel() # 基座 4bit LoRA float targets collect_shardable_modules(model, meshFakeMesh()) # 含 4bit 基座的模块不应在分片目标里 for m in targets: assert not any(is_quantized_param(p) for p in m.params()) # LoRA 模块应在 assert any(lora in m.name for m in targets) def test_4bit_not_sharded(): model FakeQLoRAModel() fully_shard_qlora_safe(model, FakeMesh()) # 4bit 基座应仍保持完整未被分片标记 assert model.base_4bit.sharded is False def test_no_silent_corruption(): # 排除后4bit 反量化值应与单卡一致 model FakeQLoRAModel() fully_shard_qlora_safe(model, FakeMesh()) assert abs(dequant(model.base_4bit) - 0.5) 1e-6 def test_lora_still_sharded(): model FakeQLoRAModel() fully_shard_qlora_safe(model, FakeMesh()) assert model.lora.sharded is True再加一个端到端回归QLoRA FSDP2 多卡训练不静默损坏基座def test_qlora_fsdp2_multi_gpu_no_corruption(): model load_qlora_model() fully_shard_qlora_safe(model, meshmake_mesh(4)) # 基座量化参数未被分片破坏 assert not is_sharded_4bit(model) # 训练若干步基座反量化值稳定 base_val dequant(model.base) train_steps(model, 5) assert abs(dequant(model.base) - base_val) 1e-3 # 基座未静默损坏八、排查清单看 QLoRA 多卡 FSDP2 训练结果乱码/精度差、无报错 → 是 4bit 被分片静默损坏。检查fully_shard(model)是否把含Params4bit的基座层也分片了。临时救火只对 LoRAfloat模块fully_shard排除 4bit 基座。确认基座requires_gradFalse仍可能被分片FSDP2 为省内存也会分片冻结参数。长期修复用fully_shard_qlora_safe自动检测并排除量化参数。升级 accelerate/pytorch 到合了量化参数自动排除的版本并跑上面的test_4bit_not_sharded。若用device_map卸载 FSDP24bit 基座同样不能 shard需一并排除。九、小结FSDP2 下 QLoRA 静默损坏不是 QLoRA 错了而是**fully_shard不分青红皂白地把 4-bit 量化基座当普通参数分片破坏了其量化反量化所需的全局 scale 上下文基座被悄悄改坏、LoRA 训练污染**。最小修复是只对 LoRAfloat模块fully_shard、排除 4bit 基座结构性修复是fully_shard_qlora_safe自动检测Params4bit并排除其分片最后用 pytest 把「量化参数排除」「不分片」「不静默损坏」锁死。抓住「量化参数不能被普通张量分片语义处理、FSDP 分片前必须识别并排除量化层」这条所有 QLoRA FSDP 的静默损坏都能照此化解。