Flash Diffusion训练实战:4阶段蒸馏SDXL,所有YAML配置参数逐项讲透
【免费下载链接】flash-diffusionFlash Diffusion — accelerating conditional diffusion models (AAAI 2025 Oral)项目地址: https://gitcode.com/gh_mirrors/fl/flash-diffusion
Flash Diffusion 训练实战指南来了!Flash Diffusion(AAAI 2025 Oral)是一种通过 4 阶段蒸馏把 SDXL 等扩散模型加速到 4 步出图的方法。本文将带你逐项读懂 SDXL 蒸馏的 YAML 配置参数,几小时 GPU 训练时间,即可得到一个少步高速出图的 LoRA。
1. 先懂方法:Flash Diffusion 蒸馏原理
Flash Diffusion 的核心思想很简单:让"学生模型"单步预测教师模型多步去噪的结果,并用一个随训练动态变化的时间步分布引导学生逐步学会更少的去噪步数。
整个流程由训练脚本examples/train_flash_sdxl.py驱动,配置则完全来自examples/configs/flash_sdxl.yaml。脚本会自动完成:加载 SDXL 教师模型 → 深拷贝出学生 UNet 并挂上 LoRA → 组装 CLIP/Timesteps 条件编码器 → 构建判别器 → 启动训练。
2. 快速安装:3 步跑通环境
- 创建并激活 Python 3.10 虚拟环境(
venv或conda均可) - 安装依赖:
pip install -r requirements.txt - 以可编辑模式安装项目:
pip install -e .
仓库的 requirements.txt 已整理好全部依赖;setup.py 负责把flash包注册为可安装模块。
3. 4 阶段蒸馏:K 与 NUM_ITERATIONS_PER_K 怎么读?
K: [32, 32, 32, 32] NUM_ITERATIONS_PER_K: [5000, 5000, 5000, 5000]这是整套配置的灵魂,定义了渐进式蒸馏的 4 个阶段:
- K:每个阶段学生模型允许的最大去噪步数。4 个阶段都从 32 步开始,逐步"压缩"步数空间,让学生从"32 步以内也能学好"过渡到"4 步也够用"。
- NUM_ITERATIONS_PER_K:每个阶段训练 5000 个 step,4 个阶段共 20000 步——这就是论文所说"只需几小时 GPU 时间"的原因。
训练中的"换挡"逻辑
训练时模型按累计步数判断当前阶段(见 flash_diffusion_model.py,内部用K_steps = np.cumsum(NUM_ITERATIONS_PER_K)计算切换点),阶段越往后:
| 参数 | 阶段1 → 阶段4 的变化趋势 | 作用 |
|---|---|---|
MODE_PROBS | [0,0,0.5,0.5]→[0.4,0.2,0.2,0.2] | 时间步混合分布的重心逐步向高噪声区域移动,逼迫学生适应少步生成 |
ADVERSARIAL_LOSS_SCALE | 0→0.3 | GAN 对抗损失逐渐加入,后期提升图像真实感 |
DMD_LOSS_SCALE | 0→0.7 | DMD 分布匹配损失逐步增强,让学生拟合教师多步输出分布 |
DISTILL_LOSS_SCALE | 恒定1.0 | 基础蒸馏损失始终开启 |
4. YAML 参数逐项讲透
4.1 数据集部分:webdataset 分片流
SHARDS_PATH_OR_URLS: - pipe:cat /path/to/tar/files/{000000..000010}.tar VALIDATION_PROMPTS: - A beautiful red car on the beach at sunset, 4k, photorealistic, awesome.SHARDS_PATH_OR_URLS:训练数据以webdataset 的 tar 分片喂入,只需把路径换成你自己的数据。每个样本需包含jpg图片和带caption、aesthetic_score两个键的json文件(数据管线在 train_flash_sdxl.py 中由KeysFromJSONMapper等映射器完成解析,并自动过滤美学分低于 6.0 的样本)。VALIDATION_PROMPTS:训练中周期性采样的验证提示词,决定你在日志里看到哪些对比图。
4.2 模型部分:损失函数与调度器
LORA: True LORA_RANK: 64 DISTILL_LOSS_TYPE: lpips UCG_KEYS: [text] TIMESTEP_DISTRIBUTION: mixture MIXTURE_NUM_COMPONENTS: 4 MIXTURE_VAR: 0.5 GAN_LOSS_TYPE: lsgan TEACHER_SCHEDULER: DPMSolverMultistepScheduler SAMPLING_SCHEDULER: LCMScheduler TEACHER_SAMPLING_SCHEDULER: EulerDiscreteScheduler USE_TEACHER_AS_REAL: False USE_EMPTY_PROMPT: False逐项拆解:
LORA: True+LORA_RANK: 64:只训练 UNet 注意力层的 LoRA(to_q/to_k/to_v),可训练参数极少,是方法高效的关键。DISTILL_LOSS_TYPE: lpips:蒸馏损失用 LPIPS 感知损失(VGG 网络),比 L1/L2 更能保住纹理细节,可选值l2 / l1 / lpips(定义见 flash_diffusion_config.py)。UCG_KEYS: [text]:教师模型做分类器引导(UCG)时随机置空的维度,只针对文本条件。TIMESTEP_DISTRIBUTION: mixture+MIXTURE_NUM_COMPONENTS: 4+MIXTURE_VAR: 0.5:时间步从 4 分量的高斯混合分布中采样,MODE_PROBS控制各分量权重并逐阶段漂移——这是"动态时间步分布"的实现核心。GAN_LOSS_TYPE: lsgan:判别器损失类型(可选hinge / vanilla / non-saturating / wgan / lsgan)。判别器是挂在教师 UNet 特征上的 4 层卷积网络,直接写在训练脚本里。- 三个调度器分工:
TEACHER_SCHEDULER负责教师加噪/多步去噪参考;SAMPLING_SCHEDULER: LCMScheduler是学生推理时的调度器;TEACHER_SAMPLING_SCHEDULER仅用于日志中教师对照图的采样。 USE_TEACHER_AS_REAL: False:对抗损失的"真实图"用数据集原图而非教师生成图(避免学生模仿教师的偏差)。USE_EMPTY_PROMPT: False:SDXL 引导用完整空文本嵌入,故关闭;SD1.5 与 Pixart 配置中为True。
4.3 训练部分:学习率与批大小
LR: 0.00001 LR_DISCRIMINATOR: 0.00001 MAX_EPOCHS: 100 BATCH_SIZE: 2学生 LoRA 与判别器各用一个 AdamW 优化器(配置见 training_config.py),学习率均为 1e-5。BATCH_SIZE是每 GPU 的批大小——SDXL 显存占用大,默认只开 2;多卡用环境变量SLURM_NPROCS/SLURM_NNODES控制。
4.4 日志部分:每 N 步看效果
LOG_EVERY_N_BATCHES: 200 NUM_STEPS: [1, 2, 4] LOG_TEACHER_SAMPLES: True CKPT_EVERY_N_STEPS: 5000 TEACHER_SAMPLING_GUIDANCE_SCALE: 7.5每 200 个 batch 用验证提示词分别以 1、2、4 步采样学生图,并附教师对照图写入 WandB 日志;每 5000 步保存一次检查点。训练日志器实现位于 loggers.py。
5. 训练效果:4 步出图照样能打
训练完成后,得到的 LoRA 配合 LCMScheduler 只需4 NFE即可出图,质量对标几十步的原始模型:
同样的方法在 SD1.5、Pixart-α(DiT) 和 Canny 适配器上都有对应脚本与配置:train_flash_sd.py、train_flash_pixart.py、train_flash_canny_adapter.py。
6. 四套官方配置的差异速查
| 参数 | flash_sd.yaml | flash_sdxl.yaml | flash_pixart.yaml | flash_canny_adapter.yaml |
|---|---|---|---|---|
| LORA_RANK | 128 | 64 | 64 | 128 |
| K | [32,32,32,32] | [32,32,32,32] | [16,16,16,16] | [16,16,16,16] |
| 每阶段迭代数 | 5000 | 5000 | 10000 | 5000 |
| GUIDANCE_MIN / MAX | 3 / 13 | 3 / 13 | 2 / 9 | 3 / 13 |
| ADVERSARIAL_LOSS_SCALE | 0→0.3 | 0→0.3 | 0→0.2 | 0→0.3 |
| BATCH_SIZE | 4 | 2 | 2 | 2 |
| USE_EMPTY_PROMPT | True | False | True | False |
规律一目了然:SDXL 这类 1024 分辨率大模型 batch 只能开 2;DiT 与适配器从 16 步起步;Pixart 需要 4 倍迭代量补偿。改一个配置就能复刻论文任意实验。
7. 一键启动训练
# 设置 GPU 数与节点数 export SLURM_NPROCS=1 export SLURM_NNODES=1 # 蒸馏 SDXL(配置在 examples/configs/flash_sdxl.yaml) python3.10 examples/train_flash_sdxl.py训练日志与检查点默认输出到logs/时间戳-FlashSDXL/。每隔CKPT_EVERY_N_STEPS步落盘一份检查点,取最后一份即可作为最终的 Flash LoRA 使用。
小结:Flash Diffusion 的 YAML 看似参数众多,实则围绕一条主线——4 阶段渐进蒸馏:K定阶段、NUM_ITERATIONS_PER_K定时长、MODE_PROBS让时间步分布逐步漂移、三类损失(蒸馏/DMD/GAN)按阶段加权接力。读懂这张表,你就掌握了 AAAI 2025 Oral 的完整训练配方,几小时 GPU 时间,让 SDXL 实现 4 步闪电出图 ⚡
【免费下载链接】flash-diffusionFlash Diffusion — accelerating conditional diffusion models (AAAI 2025 Oral)项目地址: https://gitcode.com/gh_mirrors/fl/flash-diffusion
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考