源码深度解析:Erasing Concepts from Diffusion Models共享训练核心与Adapter设计模式
【免费下载链接】erasingErasing Concepts from Diffusion Models项目地址: https://gitcode.com/gh_mirrors/er/erasing
在 AI 绘图安全治理领域,Erasing Concepts from Diffusion Models(简称 ESD,中文常称"概念擦除")是一套通过微调权重、让 Stable Diffusion / SDXL / FLUX 等扩散模型"忘记"指定概念的开源方案。无论是清除版权风格、物体还是不安全内容,它都只需一段文本描述即可定向擦除。本文将深入解析该项目最新版代码:一套共享训练核心如何用Adapter 设计模式优雅地统一驱动 SD、SDXL、FLUX、FLUX.2 Klein 四种模型家族的训练,同时节省近一半显存并提速 5-8 倍。
一、先看懂整体架构:四入口一核心
整个仓库采用"薄入口 + 厚核心"的经典布局。四个训练入口脚本只负责解析命令行参数、组装配置:
- esd_sd.py —— Stable Diffusion V1.x 家族入口
- esd_sdxl.py —— SDXL 家族入口
- esd_flux.py —— FLUX.1 家族入口
- esd_flux2_klein.py —— FLUX.2 Klein 家族入口
仔细观察会发现,这四个脚本的main()结构几乎完全一致:构建 argparse 解析器 → 把参数填进统一的ESDConfig数据类 → 调用共享的run_esd_training(config)→ 打印 checkpoint 保存路径。真正的训练逻辑全部沉淀在 utils/esd_trainer.py 这一个文件里,这就是项目的共享训练核心。
以 SD 为例,一条命令即可启动概念擦除训练:
python esd_sd.py --erase_concept 'Van Gogh' --train_method 'esd-x'二、共享训练核心:ESDConfig 与统一训练循环
1. 配置数据类 ESDConfig:一份配置走天下
utils/esd_trainer.py 中的ESDConfig是一个 dataclass,把四种模型家族的训练参数统一收纳:family(模型家族标识)、base_model_id(基础模型)、erase_concept(要擦除的概念)、train_method(训练方法)、iterations、lr、negative_guidance等。它还提供了一个巧妙的属性erase_from_effective:当用户不指定erase_from时,默认把"要从谁身上擦除"当成"擦除概念本身"。
2. 统一训练循环 run_esd_training:真正的共享核心
utils/esd_trainer.py 的run_esd_training是整套代码的心脏,无论训练哪个模型家族都走同一条流水线:
- 获取适配器:
get_adapter(config.family)按家族名取出对应 Adapter; - 加载管线:
adapter.load_pipeline(config)加载对应 diffusers 管线并冻结非训练组件; - 准备可训练参数:
adapter.create_prepared_component(...)选出要训练的参数并构建学生/教师双份快照; - 缓存文本嵌入:
adapter.prepare_context(...)一次性算出"擦除概念 / 空概念 / 来源概念"三组 embedding; - 迭代优化:循环里
adapter.training_step(...)产出预测与目标,用F.mse_loss计算损失反向传播; - 保存 checkpoint:
save_esd_checkpoint(...)写出带元数据的安全张量文件。
这份循环对四种模型家族完全透明——差异全部被 Adapter 吸收掉了。
三、Adapter 设计模式:抽象基类如何屏蔽模型差异
1. BaseESDAdapter:定义统一契约
utils/esd_trainer.py 中的抽象基类BaseESDAdapter定义了每个模型家族必须实现的方法契约:
load_pipeline—— 加载本家族的 diffusers 管线normalize_train_method—— 把历史别名(如xattn)归一化为esd-xselect_parameter_names—— 按训练方法筛选可训练参数名prepare_context—— 预处理提示词嵌入等上下文training_step—— 执行单步 ESD 训练
基类还提供了一组"模板方法"式的默认实现:resolve_learning_rate(未指定学习率时按家族取默认值)、resolve_resolution(未指定分辨率时取模型原生尺寸)、build_metadata、build_checkpoint_path等,子类只需按需覆写。
2. 四大 Adapter 实现:各自精彩
仓库末尾的 utils/esd_trainer.py 用注册表模式把它们组织起来:
ADAPTERS = { "sd": StableDiffusionESDAdapter(), "sdxl": StableDiffusionXLESDAdapter(), "flux": FluxESDAdapter(), "flux2_klein": Flux2KleinESDAdapter(), }- StableDiffusionESDAdapter:组件指向
unet,支持esd-x(只训交叉注意力 attn2)、esd-u、esd-all、esd-x-strict、selfattn等多种训练方法; - StableDiffusionXLESDAdapter:同样训练
unet,但默认学习率区分方法(esd-x用 2e-4,其余用 1e-5),并会对大范围更新的esd-u/esd-all发出质量警告; - FluxESDAdapter:组件切换到
transformer,且训练时vae=None不加载 VAE,采样潜变量完全由 transformer 前向产生; - Flux2KleinESDAdapter:基于更新的
Flux2KleinPipeline,参数筛选精确到注意力投影后缀(to_q/to_k/to_v等)。
得益于这个设计,新增一个模型家族只需"实现接口 + 注册",共享训练核心一行都不用改,这就是 Adapter 设计模式带来的扩展性红利。
四、训练方法选择:esd-x 与参数筛选的精妙
不同训练方法决定了"擦除"的颗粒度。utils/esd_trainer.py 中 SD 适配器的select_parameter_names展示了筛选逻辑:
esd-x:只选择包含attn2的模块(交叉注意力),这是概念知识的主要"存储仓库";esd-u:选择不含attn2的模块(UNet 其余部分);esd-all:全部参数;esd-x-strict:进一步收窄到attn2.to_k与attn2.to_v,用最少的参数实现擦除。
通用筛选器 select_parameter_names 只接受TARGET_MODULE_TYPES中的模块类型(Linear、Conv2d、LoRA 兼容层等),并用named_modules遍历 + 去重,保证选出的参数名干净可靠。对普通用户而言,日常使用esd-x或esd-x-strict即可获得擦除效果与生成质量的较好平衡。
五、学生-教师双快照:PreparedComponent 的内存妙招
ESD 训练最核心的机制是"用冻结的原模型(教师)指导可微的学生模型"。utils/esd_trainer.py 的PreparedComponent为此设计了同一份权重上的双参数快照:
base_params:教师参数副本(requires_grad=False);student_params:学生参数(requires_grad=True)。
训练时通过use_base()/use_student()原地替换模块参数,配合set_module递归设置,实现了"同一模型、两套状态"的零拷贝切换。这避免了维护两份完整 UNet/Transformer 的大内存开销——正是官方宣称"几乎省一半显存"的关键之一。单步训练的逻辑(以 SD 为例,见 utils/esd_trainer.py)分两大阶段:
- 采样阶段(教师):随机选取一个去噪步长
run_till_timestep,用基础模型采样中间潜变量xt,并前向算出三组噪声预测:noise_pred_erase(擦除概念)、noise_pred_null(空概念)、noise_pred_erase_from(来源概念); - 优化阶段(学生):切回学生参数,对
xt再次前向得到model_pred,构造 ESD 专属目标:
target = noise_pred_from - negative_guidance * (noise_pred_erase - noise_pred_null)这个目标公式的直觉是:既要让输出"靠近保留内容",又要沿"擦除概念与空概念之差"的方向反向推开,从而在保持其余生成能力的同时定向抹除目标概念。最后用 MSE 损失F.mse_loss(model_pred, target)更新学生权重。
六、工程化细节:内存卸载与元数据感知的 Checkpoint
共享训练核心还内置了两项值得借鉴的工程优化:
1. 训练时的显存卸载
prepare_context 在缓存完文本嵌入后立即调用offload_modules_to_cpu把 VAE、文本编码器等非训练组件搬到 CPU,并执行torch.cuda.empty_cache()+gc.collect(),让宝贵的 GPU 显存全部服务于可训练组件。
2. 元数据感知的 Checkpoint 体系
utils/esd_checkpoint.py 实现了完整的保存/加载协议:
save_esd_checkpoint写入format: erasing-esd-v2格式标记,并附带 family、component、train_method、erase_concept 等完整元数据;infer_checkpoint_component在无元数据时也能通过参数名重叠度自动推断 checkpoint 属于unet还是transformer,因此 evalscripts/generate-images.py 等评估脚本可以无差别处理 SD/SDXL/FLUX 的模型;- checkpoint 文件名规则清晰:
esd-{概念}-from-{来源}-{方法}.safetensors,见 build_checkpoint_path。
七、应用效果:从裸体内容到艺术风格与物体
共享训练核心最终服务的是一系列真实应用场景。README.md 中的案例图直观展示了擦除效果:同一模型在擦除前能生成目标概念,擦除后则完全规避。
在艺术风格擦除上,ESD 相比 Safe Concept Deletion、SLD 等基线能更彻底地抹除特定画家风格,同时保留提示词其他语义:
在NSFW 内容治理上,NudeNet 定量评估显示,ESD 对女性胸部、生殖器等暴露部位的消除比例远超 SLD Medium 与 SD 2.0/2.1 等方案:
总结:一套核心,四种模型,一个模式
回看整个代码库,共享训练核心 + Adapter 设计模式的收益非常清晰:训练循环、损失构造、checkpoint 协议、显存优化全部收敛在一处;模型差异被四个 Adapter 封装成统一接口;新增模型家族只需"实现契约 + 注册表登记"。如果你也在设计多后端训练框架,这套分层思路非常值得借鉴——让复杂的扩散模型概念擦除,变得像插拔适配器一样简单。
【免费下载链接】erasingErasing Concepts from Diffusion Models项目地址: https://gitcode.com/gh_mirrors/er/erasing
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考