Flash Diffusion开发者指南:手把手带你蒸馏自己的条件扩散模型
【免费下载链接】flash-diffusionFlash Diffusion — accelerating conditional diffusion models (AAAI 2025 Oral)项目地址: https://gitcode.com/gh_mirrors/fl/flash-diffusion
Flash Diffusion 是 AAAI 2025 Oral 论文《Flash Diffusion: Accelerating Any Conditional Diffusion Model for Few Steps Image Generation》的官方开源实现,它能用几个 GPU 小时的训练,把 SD1.5、SDXL、PixArt-α 等任意条件扩散模型蒸馏成只需4 步甚至 1 步出图的快速模型,且质量几乎无损。本文是一份面向新手的手把手指南,带你从零完成:环境安装、复现官方蒸馏实验、以及蒸馏你自己的条件扩散模型。
一分钟认识 Flash Diffusion:为什么它快?
传统扩散模型(如 Stable Diffusion)需要 20~50 次去噪迭代才能生成一张图,速度慢是落地的最大痛点。Flash Diffusion 的做法可以概括为一句话:训练一个"学生"网络去单步预测"教师"网络多步去噪的最终结果,配合会随训练进程动态移动的时间步采样分布,逐步把学生的能力从"粗去噪"压到"精去噪"。
整个方法只需训练少量 LoRA 参数(rank 64~128),远低于全参数蒸馏的开销,因此几块消费级 GPU 就能跑完。下图就是仅用 4 个 NFE(函数调用次数)生成的图片效果,细节依然非常扎实:
三种骨干网络,同一套蒸馏配方
Flash Diffusion 的"通用性"体现在:无论你的去噪器是 UNet 还是 DiT,是文生图、修复还是人脸交换,都适用同一套蒸馏管线:
- Flash SD:由 SD1.5 教师蒸馏,UNet 骨干
- Flash SDXL:由 SDXL 教师蒸馏,UNet 骨干
- Flash PixArt:由 PixArt-α 教师蒸馏,DiT 骨干(证明方法不依赖 UNet)
环境安装:三步跑起来
项目要求Python 3.10+,核心依赖在 requirements.txt 与 setup.py 中(lightning 2.2.5、diffusers 生态、peft、webdataset、wandb 等)。
# 1. 克隆代码 git clone https://gitcode.com/gh_mirrors/fl/flash-diffusion cd flash-diffusion # 2. 创建并激活虚拟环境 python3.10 -m venv envs/flash_diffusion source envs/flash_diffusion/bin/activate # 3. 安装依赖并以可编辑模式安装 pip install --upgrade pip pip install -r requirements.txt pip install -e .💡 多卡训练通过环境变量
SLURM_NPROCS/SLURM_NNODES控制,单机单卡直接设为 1 即可。
最快上手:复现一个现成的蒸馏实验
examples/目录提供了 4 个开箱即用的训练脚本,每个脚本对应一个官方配置:
| 训练脚本 | 蒸馏对象 | 配置文件 |
|---|---|---|
| train_flash_sd.py | SD1.5 | configs/flash_sd.yaml |
| train_flash_sdxl.py | SDXL | configs/flash_sdxl.yaml |
| train_flash_pixart.py | PixArt-α (DiT) | configs/flash_pixart.yaml |
| train_flash_canny_adapter.py | Canny 适配器 | configs/flash_canny_adapter.yaml |
第一步:准备 WebDataset 格式的训练数据
数据流由webdataset驱动(封装在 src/flash/data/datasets/dataset.py)。把数据打包成.tar,每条样本包含一张jpg图片和一个json文件:
sample = { "jpg": dummy_image, "json": { "caption": "dummy caption", "aesthetic_score": 6.0 } }然后把 yaml 中的SHARDS_PATH_OR_URLS改成你的 tar 路径(支持{000000..000010}这种花括号批量写法),例如 flash_sd.yaml 里的写法:
SHARDS_PATH_OR_URLS: - pipe:cat /path/to/tar/files/{000000..000010}.tar第二步:一条命令启动蒸馏
export SLURM_NPROCS=1 export SLURM_NNODES=1 # 蒸馏 SD1.5(SDXL / PixArt / Canny 同理) python3.10 examples/train_flash_sd.py训练全程有 W&B 日志、每LOG_EVERY_N_BATCHES步自动生成 1/2/4 步的对比样图,让你直观看到学生模型"越蒸越快、越蒸越稳"的过程。
核心机制拆解:4 个阶段 + 3 种损失
蒸馏训练的"灵魂"都在配置里(默认值见 flash_diffusion_config.py),训练逻辑在 flash_diffusion_model.py。以 flash_sd.yaml 为例:
K: [32, 32, 32, 32] # 每个阶段教师的时间步数 NUM_ITERATIONS_PER_K: [5000, ...] # 每个阶段的迭代步数 TIMESTEP_DISTRIBUTION: mixture # 高斯混合分布采样时间步 USE_DMD_LOSS: True # 启用 DMD 分布匹配损失 ADVERSARIAL_loss_SCALE: [0, 0.1, 0.2, 0.3] # GAN 损失逐渐增强🔑关键参数速查:
K:教师一次采样的步数,训练分 4 个阶段推进,学生逐步学会更少步数出图Timestep Distribution:支持uniform/gaussian/mixture三种时间步采样分布,mixture的概率质量会随阶段向高噪声区间移动(实现见 flash_diffusion_model.py#L135-L177)GUIDANCE_MIN/MAX:教师 CFG 引导系数的下/上限LORA_RANK:学生 LoRA 秩,SD1.5/SDXL 用 128,PixArt 用 64 即可- 三种损失:蒸馏损失(默认 LPIPS 感知损失)+ DMD 损失 + 对抗损失,三者权重按阶段递增,保证少步生成既有"分布对"又有"细节真"
进阶:蒸馏你自己的条件扩散模型
项目天然支持自定义模型蒸馏——只要你能组装出三件套:
- VAE:AutoencoderKLDiffusers,图像与潜空间互转
- 条件编码器:ClipEmbedder 等文本编码器,或多个编码器组合(ConditionerWrapper 支持任意条件拼接,比如文本 + 低清图)
- 教师去噪器:DiffusersUNet2DCondWrapper(UNet)或 DiffusersTransformer2DWrapper(DiT)
一个典型的组装流程(摘自 README 示例):
from copy import deepcopy from flash.models.unets import DiffusersUNet2DCondWrapper from flash.models.vae import AutoencoderKLDiffusers, AutoencoderKLDiffusersConfig from flash.models.embedders import ( ClipEmbedder, ClipEmbedderConfig, ConditionerWrapper, ) # VAE(从 HF Hub 加载) vae = AutoencoderKLDiffusers( AutoencoderKLDiffusersConfig("stabilityai/sdxl-vae") ) # 文本条件编码器(冻结) embedder = ClipEmbedder(ClipEmbedderConfig( version="stabilityai/stable-diffusion-xl-base-1.0", text_embedder_subfolder="text_encoder_2", tokenizer_subfolder="tokenizer_2", input_key="text", always_return_pooled=True, )) conditioner = ConditionerWrapper(conditioners=[embedder]) # 教师去噪器 → 加载教师权重后,学生 = 教师深拷贝 + LoRA unet = DiffusersUNet2DCondWrapper( in_channels=4, out_channels=4, cross_attention_dim=1280, projection_class_embeddings_input_dim=1280, class_embed_type="projection", ) student_denoiser = deepcopy(unet)之后只需把这三个对象连同FlashDiffusionConfig传入 FlashDiffusion 模型,再用 TrainingPipeline 启动训练——整个训练框架(src/flash/trainer/)会自动处理优化器、日志、断点保存。
一次蒸馏,处处可用:修复、放大、换脸
蒸馏后的模型不止能做文生图。官方实验证明同一方法可无缝迁移到多种下游任务:
甚至包括条件适配器——用 Canny 边缘图或深度图引导的文生图适配器蒸馏后,4 步即可出图(对应 train_flash_canny_adapter.py 与 DiffusersT2IAdapterWrapper):
推理:4 步出图的 Hugging Face 管线
蒸馏产出的 LoRA 可以直接挂回原版管线,配合 LCMScheduler 就能少步推理,无需任何魔改:
from diffusers import PixArtAlphaPipeline from peft import PeftModel transformer = Transformer2DModel.from_pretrained( "PixArt-alpha/PixArt-XL-2-1024-MS", subfolder="transformer" ) transformer = PeftModel.from_pretrained(transformer, "jasperai/flash-pixart") pipe = PixArtAlphaPipeline.from_pretrained( "PixArt-alpha/PixArt-XL-2-1024-MS", transformer=transformer ) pipe.scheduler = LCMScheduler.from_pretrained( "PixArt-alpha/PixArt-XL-2-1024-MS", subfolder="scheduler", timestep_spacing="trailing", ) image = pipe("A raccoon reading a book in a lush forest.", num_inference_steps=4, guidance_scale=0).images[0]常见问题与避坑清单
| 问题 | 建议 |
|---|---|
| 显存不够 | 降低BATCH_SIZE(PixArt/SDXL 官方仅用 2),训练已启用bf16-mixed混合精度 |
| 出图偏模糊 | 适当调高后期阶段的ADVERSARIAL_LOSS_SCALE与DMD_LOSS_SCALE |
| 少步出图分布偏移 | 保持TIMESTEP_DISTRIBUTION: mixture与官方MODE_PROBS的阶段调度 |
| 数据格式报错 | 确认每个 tar 样本同时含jpg与json(caption+aesthetic_score字段),且aesthetic_score >= 6.0才会被采样 |
| 想训全参数 | 将 yaml 中LORA设为False(注意显存与训练时长会显著上升) |
总结:你的首个 4 步扩散模型,只需 4 步走
- ✅
pip install -e .装好环境 - ✅ 把数据打成 webdataset 的 tar,改 yaml 中的路径
- ✅ 运行
examples/train_flash_sd.py复现官方蒸馏 - ✅ 换成自己的 VAE + 编码器 + 去噪器,蒸馏你自己的条件扩散模型
项目采用 CC BY-NC 4.0 许可证发布,研究引用可参考 README.md 末尾的 BibTeX。如果这篇 Flash Diffusion 蒸馏指南帮你跑通了第一个少步扩散模型,不妨把它的 4 步出图能力接入你的 AIGC 产品——快的同时,画质不打折。⚡
【免费下载链接】flash-diffusionFlash Diffusion — accelerating conditional diffusion models (AAAI 2025 Oral)项目地址: https://gitcode.com/gh_mirrors/fl/flash-diffusion
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考