Flash Diffusion训练实战:4阶段蒸馏SDXL,所有YAML配置参数逐项讲透
Flash Diffusion训练实战4阶段蒸馏SDXL所有YAML配置参数逐项讲透【免费下载链接】flash-diffusionFlash Diffusion — accelerating conditional diffusion models (AAAI 2025 Oral)项目地址: https://gitcode.com/gh_mirrors/fl/flash-diffusionFlash Diffusion 训练实战指南来了Flash DiffusionAAAI 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 个 step4 个阶段共 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_SCALE0→0.3GAN 对抗损失逐渐加入后期提升图像真实感DMD_LOSS_SCALE0→0.7DMD 分布匹配损失逐步增强让学生拟合教师多步输出分布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: TrueLORA_RANK: 64只训练 UNet 注意力层的 LoRAto_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: mixtureMIXTURE_NUM_COMPONENTS: 4MIXTURE_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: FalseSDXL 引导用完整空文本嵌入故关闭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.yamlflash_sdxl.yamlflash_pixart.yamlflash_canny_adapter.yamlLORA_RANK1286464128K[32,32,32,32][32,32,32,32][16,16,16,16][16,16,16,16]每阶段迭代数50005000100005000GUIDANCE_MIN / MAX3 / 133 / 132 / 93 / 13ADVERSARIAL_LOSS_SCALE0→0.30→0.30→0.20→0.3BATCH_SIZE4222USE_EMPTY_PROMPTTrueFalseTrueFalse规律一目了然SDXL 这类 1024 分辨率大模型 batch 只能开 2DiT 与适配器从 16 步起步Pixart 需要 4 倍迭代量补偿。改一个配置就能复刻论文任意实验。7. 一键启动训练# 设置 GPU 数与节点数 export SLURM_NPROCS1 export SLURM_NNODES1 # 蒸馏 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),仅供参考