源码深度解析:Erasing Concepts from Diffusion Models共享训练核心与Adapter设计模式

📅 2026/8/20 20:42:20
源码深度解析:Erasing Concepts from Diffusion Models共享训练核心与Adapter设计模式
源码深度解析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计算损失反向传播保存 checkpointsave_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且训练时vaeNone不加载 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_gradFalsestudent_params学生参数requires_gradTrue。训练时通过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),仅供参考