尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

Flash Diffusion开发者指南:手把手带你蒸馏自己的条件扩散模型

Flash Diffusion开发者指南:手把手带你蒸馏自己的条件扩散模型 Flash Diffusion开发者指南手把手带你蒸馏自己的条件扩散模型【免费下载链接】flash-diffusionFlash Diffusion — accelerating conditional diffusion models (AAAI 2025 Oral)项目地址: https://gitcode.com/gh_mirrors/fl/flash-diffusionFlash 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.pySD1.5configs/flash_sd.yamltrain_flash_sdxl.pySDXLconfigs/flash_sdxl.yamltrain_flash_pixart.pyPixArt-α (DiT)configs/flash_pixart.yamltrain_flash_canny_adapter.pyCanny 适配器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_NPROCS1 export SLURM_NNODES1 # 蒸馏 SD1.5SDXL / PixArt / Canny 同理 python3.10 examples/train_flash_sd.py训练全程有 WB 日志、每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-L177GUIDANCE_MIN/MAX教师 CFG 引导系数的下/上限LORA_RANK学生 LoRA 秩SD1.5/SDXL 用 128PixArt 用 64 即可三种损失蒸馏损失默认 LPIPS 感知损失 DMD 损失 对抗损失三者权重按阶段递增保证少步生成既有分布对又有细节真进阶蒸馏你自己的条件扩散模型项目天然支持自定义模型蒸馏——只要你能组装出三件套VAEAutoencoderKLDiffusers图像与潜空间互转条件编码器ClipEmbedder 等文本编码器或多个编码器组合ConditionerWrapper 支持任意条件拼接比如文本 低清图教师去噪器DiffusersUNet2DCondWrapperUNet或 DiffusersTransformer2DWrapperDiT一个典型的组装流程摘自 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( versionstabilityai/stable-diffusion-xl-base-1.0, text_embedder_subfoldertext_encoder_2, tokenizer_subfoldertokenizer_2, input_keytext, always_return_pooledTrue, )) conditioner ConditionerWrapper(conditioners[embedder]) # 教师去噪器 → 加载教师权重后学生 教师深拷贝 LoRA unet DiffusersUNet2DCondWrapper( in_channels4, out_channels4, cross_attention_dim1280, projection_class_embeddings_input_dim1280, class_embed_typeprojection, ) 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, subfoldertransformer ) transformer PeftModel.from_pretrained(transformer, jasperai/flash-pixart) pipe PixArtAlphaPipeline.from_pretrained( PixArt-alpha/PixArt-XL-2-1024-MS, transformertransformer ) pipe.scheduler LCMScheduler.from_pretrained( PixArt-alpha/PixArt-XL-2-1024-MS, subfolderscheduler, timestep_spacingtrailing, ) image pipe(A raccoon reading a book in a lush forest., num_inference_steps4, guidance_scale0).images[0]常见问题与避坑清单问题建议显存不够降低BATCH_SIZEPixArt/SDXL 官方仅用 2训练已启用bf16-mixed混合精度出图偏模糊适当调高后期阶段的ADVERSARIAL_LOSS_SCALE与DMD_LOSS_SCALE少步出图分布偏移保持TIMESTEP_DISTRIBUTION: mixture与官方MODE_PROBS的阶段调度数据格式报错确认每个 tar 样本同时含jpg与jsoncaptionaesthetic_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),仅供参考
返回列表