GPU显存不够怎么办?Erasing Concepts from Diffusion Models训练常见问题与优化技巧
【免费下载链接】erasingErasing Concepts from Diffusion Models项目地址: https://gitcode.com/gh_mirrors/er/erasing
在本地微调扩散模型、做扩散模型概念擦除(Erasing Concepts from Diffusion Models)训练时,"CUDA out of memory"是最常见的劝退理由。这个开源项目通过只微调注意力层的小部分参数,就能把梵高风格、裸体内容、特定物体等概念从 Stable Diffusion / SDXL / FLUX 中"擦除"掉,而它的新版本代码相比旧版几乎省一半 GPU 显存、速度提升 5-8 倍。本文总结 ESD 训练中的显存优化技巧与常见问题排查方法,帮你用一块消费级显卡跑通概念擦除训练。
先看效果:概念擦除到底能做什么
项目效果一目了然:左边是原始模型的输出,右边是擦除后的输出。无论是不安全内容、知名艺术家风格,还是"汽车"这类具体物体,都能被精准移除,而画面其余部分保持合理。
为什么训练时会爆显存
ESD 训练本质上是一次轻量微调:加载一个完整的扩散模型(UNet 或 FLUX Transformer),在显存里同时保留冻结的原始权重和可训练的副本,因此显存占用通常比单纯推理高出一倍左右。显存主要消耗在三个地方:
- 模型权重本身:SD1.4 约 4GB、SDXL 约 7GB、FLUX.1-dev 更大;
- 前向采样过程:训练时每步都要先用原始模型采样噪声潜变量
x_t; - 反向传播的中间激活值:这是最容易爆显存的环节。
好消息是,这个项目在代码层面已经做了大量省显存设计,下面这些技巧按"见效速度"排序。
8 个立竿见影的GPU显存优化技巧
1. 使用新版代码,显存直接减半
项目官方明确说明:新版代码相比旧版几乎使用一半的 GPU 显存,且快 5-8 倍。如果你是从旧 commit 或第三方教程复制来的命令,请务必用最新代码重新安装:
git clone https://gitcode.com/gh_mirrors/er/erasing cd erasing conda create --name erasing python=3.14 conda activate erasing pip install -r requirements.txt新版的核心训练逻辑统一收敛在utils/esd_trainer.py中,SD、SDXL、FLUX 共用同一套流程,安装和排错都更简单。
2. 选对训练方法:esd-x 只动注意力层
这是该项目最核心的省显存设计。ESD 不需要全量微调,esd-x方法只更新交叉注意力层(attn2)的参数,其余全部冻结;更极致的esd-x-strict甚至只更新to_k和to_v两个投影矩阵。可训练参数量骤减,意味着优化器状态、梯度和激活值都大幅减少,显存压力直接下降。相关参数选择逻辑可以在utils/esd_trainer.py中看到,比如 SD 的esd-x只选择包含attn2的模块。
python esd_sd.py --erase_concept 'Van Gogh' --train_method 'esd-x'3. 开启梯度检查点:--gradient_checkpointing
如果仍然 OOM,加上梯度检查点参数,用少量计算换显存,通常能再省 30%-50%:
python esd_sdxl.py --erase_concept 'Van Gogh' --train_method 'esd-x-strict' --gradient_checkpointing梯度检查点会丢弃中间激活、在反向传播时重新计算,是低显存训练的"后悔药",SDXL 和 FLUX 的训练入口脚本都支持该参数。
4. 降低训练分辨率:--resolution 512
显存占用与分辨率的平方成正比。FLUX 训练脚本默认就把分辨率设为 512 以控制显存占用。如果你的显卡只有 8GB,建议显式指定:
python esd_flux.py --erase_concept 'monster' --train_method 'esd-x' --resolution 512SD 和 SDXL 脚本默认使用模型原生分辨率(512 或 1024),训练 SDXL 时如果爆显存,先降到 512 试一次。
5. 保持 batchsize=1
ESD 训练每一步只需要一张噪声潜变量图,所有脚本默认--batchsize 1,这也是官方推荐值。不要为了"加快训练"去调大 batch,除非你的显存非常充裕。
6. 利用默认的 bfloat16 混合精度
所有训练入口默认使用torch.bfloat16精度加载模型,显存占用约为 FP32 的一半。只需确认你的显卡支持 bfloat16(RTX 30/40 系列、A100 等都没问题),无需额外配置即可享受减半显存。
7. 允许 TF32 加速:--allow_tf32
在 Ampere 及以上架构的显卡上,可以开启 TF32 让矩阵乘法更快、更省显存:
python esd_flux2_klein.py --erase_concept 'monster' --train_method 'esd-x' --allow_tf32FLUX.2 Klein 等新模型的训练入口均支持该开关。
8. 减少迭代步数:--iterations
默认情况下 SD 系列训练 200 步、FLUX 训练 1400 步。如果只是快速验证效果,可以先用较小步数跑通流程:
python esd_sd.py --erase_concept 'Van Gogh' --train_method 'esd-x' --iterations 100训练完成后,检查点会以.safetensors格式保存到esd-models/目录,并自动带上训练元信息。
各模型的显存参考与推荐参数
不同基座模型的参数量差异巨大,这里给出一个直观参考:
| 基座模型 | 默认分辨率 | 建议最低显存 | 推荐训练方法 | 关键参数 |
|---|---|---|---|---|
| Stable Diffusion V1.4 | 512 | 8GB 可跑通 | esd-x / esd-x-strict | iterations 200 |
| SDXL | 1024 | 12GB(建议16GB) | esd-x-strict | 分辨率降 512 更稳 |
| FLUX.1-dev | 512 | 16GB(建议24GB) | esd-x / esd-x-strict | iterations 1400 |
| FLUX.2 Klein | 512 | 16GB(建议24GB) | esd-x | 需较新 diffusers |
如果显存低于上表建议值,优先叠加使用梯度检查点 + 降低分辨率两个技巧。
ESD 训练原理:为什么会这么省显存
这张流程图解释了 ESD 的优化目标:训练一个可学习的参数副本(Fine Tune ESD),让它输出的噪声预测逐渐远离"原始模型 + 被擦除概念提示词"的组合,逼近"原始模型 + 空提示词"的预测,从而实现概念擦除。全程只优化少量参数,所以显存和速度都很友好。
常见报错排查与解决
"CUDA out of memory"
按顺序尝试:加--gradient_checkpointing→ 降低--resolution→ 确认--batchsize 1→ 换成esd-x-strict方法。
"CUDA: no kernel image is available"
通常是 PyTorch 与显卡驱动、CUDA 版本不匹配。按项目requirements.txt安装匹配版本的 torch,或升级驱动后重装虚拟环境。
想用多块显卡训练
训练脚本通过--device指定设备,默认是cuda:0。ESD 训练本身单卡即可完成,如果显存不够更建议先尝试上面 8 个技巧,而不是贸然上多卡。
训练中断、想恢复
ESD 训练步数少、速度快,中断后直接重新运行即可。新版会以safetensors保存元数据感知的检查点,训练完成后也可直接用推理脚本加载,无需额外转换。
训练完成后如何验证效果
训练完成后,可以用notebooks/目录下的推理 notebook(如esd_inference_sd.ipynb、esd_inference_sdxl.ipynb、esd_inference_flux.ipynb)快速生成对比图;要批量评估擦除效果,则使用evalscripts/generate-images.py,它会自动识别检查点目标是unet还是transformer,SD、SDXL、FLUX 通用:
python evalscripts/generate-images.py --base_model 'stabilityai/stable-diffusion-xl-base-1.0' --esd_path 'esd-models/sdxl/esd-kelly-from-kelly.safetensors' --num_samples 1 --prompts_path 'data/kelly_prompts.csv' --num_inference_steps 20 --guidance_scale 7data/目录下还提供了art_prompts.csv、vangogh_prompts.csv、unsafe-prompts4703.csv等现成提示词集,省去自己整理数据的麻烦。
总结
GPU 显存不够不是放弃扩散模型概念擦除训练的理由。记住三个核心思路:只训练注意力层(esd-x 系列方法)、开梯度检查点、降低分辨率,再配合新版代码本身省一半显存的设计,多数 8-16GB 的消费级显卡都能顺利跑通。如果在安装或训练中遇到问题,优先检查utils/esd_trainer.py的参数配置和各训练入口脚本esd_sd.py、esd_sdxl.py、esd_flux.py、esd_flux2_klein.py的命令行选项,大部分显存问题都能在参数层面解决。
【免费下载链接】erasingErasing Concepts from Diffusion Models项目地址: https://gitcode.com/gh_mirrors/er/erasing
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考