FastGen与DMD2蒸馏:Model Optimizer让文生图模型4步出高清图
【免费下载链接】Model-OptimizerA unified library of SOTA model optimization techniques like quantization, distillation, pruning, neural architecture search, speculative decoding, etc. It compresses deep learning models for downstream deployment frameworks like TensorRT-LLM, TensorRT, vLLM, etc. to optimize inference speed.项目地址: https://gitcode.com/GitHub_Trending/te/Model-Optimizer
文生图模型出图快不快,瓶颈往往不在显卡,而在采样步数——传统扩散模型要跑几十上百步去噪才能生成一张图。而 Model Optimizer 提供的FastGen 加速库,通过DMD2 蒸馏把 Qwen-Image 等文生图模型的采样步数直接压到4 步甚至 1 步,出图速度提升数十倍且画质几乎无损。本文带你快速上手这套方案。
为什么文生图模型这么慢?🐢
扩散模型(Diffusion)的生成过程是"从纯噪声一步步去噪成图像"。为保证细节质量,推理时通常需要 30~50 步(甚至上百步),每一步都要完整前向一次大 Transformer,这就是出图慢的根源。
DMD2(Distribution Matching Distillation,分布匹配蒸馏)的思路是:不去逐步模仿教师的每一步去噪,而是让学生网络一次性学会"教师最终输出的整体分布"。这样学生推理时只需要 1~4 步就能直接落到高质量图像上。
DMD2 蒸馏原理:三网协同训练
DMD2 同时训练三个网络,各司其职:
| 网络 | 角色 | 一句话解释 |
|---|---|---|
| Student(学生) | 你最终保留的少步生成器 | 学会用 1~4 步直接出图 |
| Fake-score(假评分网络) | 跟踪学生当前输出分布 | 负责"反向推力" |
| Teacher(教师) | 冻结的原始 Qwen-Image 模型 | 代表目标分布 |
训练时按student_update_freq交替两个阶段:每隔 5 步 fake-score 更新,才更新 1 次学生(分布匹配损失 + 可选 GAN 损失);其余步骤则让 fake-score 去追平学生的当前分布。分布匹配梯度会把学生推向教师、拉离假评分网络,最终学生的输出分布逼近教师。
规范配置还会叠加两个"画质增强器":
- CFG(无分类器引导):给教师加 guidance_scale(默认 4.0),让学生学到"被引导后"的更优分布;
- GAN 分支:在教师第 30 个 Transformer Block 上挂判别器头,配 R1 梯度惩罚,让图像细节更锐利。
FastGen 库:源码结构一览
DMD2 的全部数学逻辑都实现在modelopt/torch/fastgen/中,结构非常清晰:
| 模块 | 作用 |
|---|---|
| config.py | DMDConfig/DistillationConfig等 Pydantic 配置类,支持从 YAML 加载 |
| methods/dmd.py | DMD2 核心:学生损失、fake-score 损失、判别器损失 |
| pipeline.py | DistillationPipeline基类,持有 student/teacher 引用并冻结教师 |
| plugins/qwen_image.py | Qwen-Image 专用插件(2×2 patch packing、img_shapes) |
| ema.py / discriminators.py | 学生 EMA 权重跟踪 / GAN 判别器实现 |
配套资源:
- 内置蒸馏配方:modelopt_recipes/general/distillation/dmd2_qwen_image.yaml
- 完整示例与文档:examples/diffusers/fastgen/README.md
- 蒸馏 API 指南:docs/source/guides/4_distillation.rst
三步训练你的 4 步学生模型 🏋️
整个训练流程收敛在 examples/diffusers/fastgen/ 目录下,分三步:
1️⃣ 构建训练数据缓存
用官方预处理脚本把原始图像转成Qwen-Image VAE 潜变量 + 文本嵌入的缓存(避免训练时重复编码):
python examples/diffusers/fastgen/preprocess_qwen_image.py image \ --image_dir <原始图像目录> --output_dir <缓存目录> --processor qwen_image同时生成 CFG 所需的负提示嵌入(一次性):
python examples/diffusers/fastgen/make_negative_prompt_embedding.py \ --output <缓存目录>/negative_prompt_embedding.pt2️⃣ 启动多卡训练
使用规范配置 configs/dmd2_qwen_image.yaml(4 步学生 + CFG + GAN 分支),以 8 卡为例:
torchrun --nproc-per-node=8 \ examples/diffusers/fastgen/dmd2_finetune.py \ --config examples/diffusers/fastgen/configs/dmd2_qwen_image.yaml \ --step_scheduler.max_steps=5000所有 DMD2 参数都支持命令行覆盖,例如--dmd2.guidance_scale=3.5、--fsdp.dp_size=16(显存不够时加卡或开启--fsdp.activation_checkpointing=true)。
3️⃣ 断点续训
检查点会完整保存学生、fake-score、EMA 影子权重和迭代计数器,设置restore_from: LATEST后重启即自动从最新检查点恢复,无需手动处理。
关键配置速查表 ⚙️
| 配置项 | 默认值 | 说明 |
|---|---|---|
student_sample_steps | 4 | 学生推理步数,改成 1 即单步学生 |
t_list | [0.999, 0.74925, 0.4995, 0.24975, 0.0] | 4 步的时间步调度,与推理采样点严格对齐 |
guidance_scale | 4.0 | 教师 CFG 强度,设为 null 可关闭 |
student_update_freq | 5 | 1 次学生更新 : N 次 fake-score/判别器更新 |
gan_loss_weight_gen | 0.03 | GAN 生成器权重,设 0 关闭 GAN 分支 |
ema.decay | 0.9999 | 学生 EMA 衰减系数 |
4 步推理:训练完如何出图 🎨
训练结束后,用 inference_dmd2_qwen_image.py 加载"整合版学生 Transformer + 基础 VAE/文本编码器",接口与 diffusers 风格一致:
pipe = QwenImageDMDInferencePipeline.from_pretrained( student_path="/path/to/checkpoint/.../model/consolidated", base_pipeline_path="Qwen/Qwen-Image", ) image = pipe(prompt="a small red cube on a white table", num_inference_steps=4).images[0] # 与训练步数一致几个实用技巧:
num_inference_steps必须等于训练时的student_sample_steps(4 步学生就填 4);- 想验证 EMA 权重效果,把
ema_path指向检查点里的ema_shadow.pt; - 也可以直接跑命令行 CLI:
python examples/diffusers/fastgen/inference_dmd2_qwen_image.py --student_path ... --prompt "..."。
常见问题 FAQ 💡
| 问题 | 原因与解决 |
|---|---|
| CUDA 显存不足 | 训练要同时持有学生+教师+fake-score 三个 Transformer,增加 GPU 数(--fsdp.dp_size)或开启激活检查点 |
| 第 0 步 loss 变 NaN | 几乎总是时间步越界:别把dmd2.pred_type改成flow以外(Qwen-Image 是 rectified-flow 模型),也不要改动时间步调度 |
| 报"CFG 缺少负提示嵌入" | 设置negative_prompt_embedding_path,或把guidance_scale设为 null 关闭 CFG |
| Dataloader 空批次 | 缓存样本数需 ≥local_batch_size × dp_size,分布式采样器会丢弃不完整批次 |
还能叠加哪些加速?🚀
DMD2 蒸馏解决的是步数问题,Model Optimizer 还能同时解决精度与显存问题:
- 量化:FP8 / INT8 / NVFP4 量化推理,见 examples/diffusers/quantization/;
- 缓存扩散(Cache Diffusion):利用相邻步特征相似性跳过重复计算,见 examples/diffusers/cache_diffusion/;
- 蒸馏后的学生模型同样可以量化部署到 TensorRT-LLM、TensorRT 等推理框架。
少步蒸馏 × 低比特量化组合拳,正是生产级文生图推理加速的完整答案。
总结
- FastGen(modelopt/torch/fastgen/)把 DMD2 蒸馏的完整实现打包成开箱即用的库 + 配方;
- 训练只需"预处理数据 → 一条 torchrun 命令",examples/diffusers/fastgen/ 提供了端到端示例;
- 最终产物是一个4 步(甚至 1 步)出图的文生图学生模型,配合量化与缓存加速,推理成本可下降一个数量级。
上手路径建议:先通读 examples/diffusers/fastgen/README.md,跑通推理 CLI 感受 4 步出图效果,再按上文配置启动自己的蒸馏训练。
【免费下载链接】Model-OptimizerA unified library of SOTA model optimization techniques like quantization, distillation, pruning, neural architecture search, speculative decoding, etc. It compresses deep learning models for downstream deployment frameworks like TensorRT-LLM, TensorRT, vLLM, etc. to optimize inference speed.项目地址: https://gitcode.com/GitHub_Trending/te/Model-Optimizer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考