DiffSynth Studio 模型压缩实战:3 种蒸馏方案 + 8 步出图快速上手
【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio
30 步采样出一张 1024 图像要等 40 秒,实时交互场景等不起。DiffSynth Studio 内置的蒸馏训练让模型用更少步数学会多步生成的效果,Diffusion 推理加速到 8~15 步,观感接近原模型。本文跑通 LoRA 蒸馏全流程,并按数据流讲清原理。
方案总览
| 方案 | 原理一句话 | 常用步数 | 适合谁 |
|---|---|---|---|
| 全量直接蒸馏 | 优化全部参数,让少步输出对齐多步输出 | 8~15 | 能全量微调、要部署成整包权重的团队 |
| LoRA 直接蒸馏 | 冻结底模,只训低秩适配器 | 8~15 | 预算有限、要兼容开源 LoRA 生态 |
| 轨迹模仿蒸馏 | 学生逐步对齐教师的采样轨迹 + 感知正则 | 8 | 把 Turbo 类快速模型压到极限步数 |
三种方案共用 examples/qwen_image/model_training/train.py 这套训练入口,切换方式只是换一个--task参数。下面挑最常用的 LoRA 直接蒸馏,三步跑通。
三步跑通
Step 1:环境就绪
克隆仓库并安装依赖,训练命令基于accelerate,装好后即可直接跑:
git clone https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio cd DiffSynth-Studio pip install -e .Step 2:跑蒸馏训练
官方示例脚本是 Qwen-Image-Distill-LoRA.sh,关键参数如下,可以直接复用:
accelerate launch examples/qwen_image/model_training/train.py \ --dataset_base_path data/diffsynth_example_dataset/qwen_image/Qwen-Image-Distill-LoRA \ --dataset_metadata_path data/diffsynth_example_dataset/qwen_image/Qwen-Image-Distill-LoRA/metadata.csv \ --model_id_with_origin_paths "Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors,Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors" \ --lora_base_model dit \ --lora_rank 32 \ --num_epochs 5 \ --use_gradient_checkpointing \ --task direct_distill这条蒸馏训练命令和普通 LoRA SFT 脚本几乎一样,唯一的实质区别就是--task direct_distill。
Step 3:少步推理验证
训练完加载蒸馏后的 LoRA,把num_inference_steps改到 15、cfg_scale设为 1:
from diffsynth.pipelines.qwen_image import QwenImagePipeline, ModelConfig import torch pipe = QwenImagePipeline.from_pretrained(torch_dtype=torch.bfloat16, device="cuda", model_configs=[ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="transformer/diffusion_pytorch_model*.safetensors"), ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="text_encoder/model*.safetensors"), ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="vae/diffusion_pytorch_model.safetensors")], tokenizer_config=ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="tokenizer/")) pipe.load_lora(pipe.dit, "path/to/distilled_lora.safetensors") image = pipe("水下少女,梦幻唯美", seed=0, num_inference_steps=15, cfg_scale=1)完整示例在 examples/qwen_image/model_inference/Qwen-Image-Distill-LoRA.py,含模型下载逻辑,可以直接复用。
原理拆解:一次推理里到底变了什么
不绕术语,直接看数据流。
调度器决定"跑多少步"
正常推理时scheduler.set_timesteps(30)给 DiT 排 30 步去噪路径。直接蒸馏训练里做同样的事,只是步数换成小数字:set_timesteps(inputs["num_inference_steps"]),然后按这个短时间表逐步跑model_fn和pipe.step直到走完,代码在 diffsynth/diffusion/loss.py 的DirectDistillLoss。
损失函数只对齐"终点"
注意DirectDistillLoss不约束中间轨迹,只在循环结束后比较最终 latent 和数据集里的干净 latentinput_latents(mse_loss(inputs["latents"], inputs["input_latents"]))。这张图是底模自己用多步采样生成的,相当于"教师"就是同模型多步输出的自己,学生是同模型的少步路径加 LoRA。这也解释了为什么数据集 metadata 里要带num_inference_steps字段——每条样本指定目标步数。
轨迹模仿:把中间步也管住
TrajectoryImitationLoss多做两件事:先让教师按cfg_scale=2、50 步采样,用fetch_trajectory记下 8 个目标时刻的中间 latent;再让学生的 8 步跑,每步把输入 latent 拉回教师对应时刻,目标速度取(teacher_next - teacher_now) / Δσ;最后把双方结果都解码回像素,加一项 LPIPS 感知损失。效果是学生不仅终点到位,走的路径也相似。官方示例命令见 Z-Image-Turbo.sh。
实测对比与选型建议
| 模型系列 | 蒸馏方案 | 步数(原始→蒸馏后) | 大致加速 | 质量体感 |
|---|---|---|---|---|
| Qwen-Image | LoRA 直接蒸馏 | 30 → 15 | 3~5 倍 | 观感无明显差异 |
| Z-Image-Turbo | 轨迹模仿 | 28 → 8 | 4~6 倍 | 更稳定,伪影少 |
| FLUX | 全量直接蒸馏 | 30 → 8 | 3~5 倍 | 细节尚可 |
| Wan Video | 直接蒸馏 | 50 → 20 | 2~3 倍 | 视频场景仍偏多步 |
以上数据基于官方示例脚本的常见配置给出,实际耗时随分辨率和 GPU 变化,建议用同款硬件自测。
- 只能训 LoRA、希望产物是常规 LoRA 文件塞进现有推理流程的,选 LoRA 直接蒸馏
- 底模本身是 Turbo 快速模型、还要再压到 8 步的,用轨迹模仿,小步数下比直接蒸馏更稳
- 视频模型蒸馏收益有限,先配合 split training 和梯度检查点把显存压下去,再谈步数压缩
容易踩的坑
蒸馏后画质掉点先查哪?
按顺序查三处:num_inference_steps(官方推理脚本用 15,8 步是下限);cfg_scale(蒸馏 LoRA 按 cfg=1 训练,沿用原模型的高 CFG 会偏色);数据集规模——仓库自带的 sample 数据集只演示格式,真实质量上限取决于你自己用底模生成的图像规模。
lora_rank 选多大?
官方脚本统一用 32。蒸馏脚本的--lora_target_modules覆盖注意力投影加 MLP(如to_q,to_k,to_v,...,img_mlp.net.2),比普通 LoRA 宽,所以 32 是合理起点;明显欠拟合(loss 不降、细节糊)再加到 64。
为什么训练显存比 SFT 高?
DirectDistillLoss要对整条 N 步循环反向传播,计算量和激活内存随步数放大,官方脚本都带了--use_gradient_checkpointing,你跑的时候别删。Z-Image 的轨迹模仿脚本还会先用普通 LoRA 预热一版,再通过--lora_checkpoint接上去蒸馏,小显卡可以照抄这个顺序。
跑完 Step 3 后,拿原模型num_inference_steps=30和蒸馏 LoRA=15各出一张同 prompt 的图并排对比,再决定要不要切轨迹模仿压到 8 步。
【免费下载链接】DiffSynth-StudioEnjoy the magic of Diffusion models!项目地址: https://gitcode.com/GitHub_Trending/dif/DiffSynth-Studio
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考