- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
CoDi(Conditional Diffusion Distillation)是 Google Research 提出的一种单阶段条件扩散蒸馏方法,本仓库(CoDi/)提供了其官方 Flax 训练实现。本文以 CoDi/README.md 为主线,完整梳理方法核心思想、HuggingFace 数据集与自有数据集两套训练流程、全部命令行参数的含义(结合 training_scripts/args.py 源码),并深入 training_scripts/train_codi_flax.py 的train_step实现,讲清“单阶段蒸馏”在代码层面是如何落地的一致性损失、单步 ODE 采样与 EMA 参数更新。读完本文,你将能够基于 Stable Diffusion v1.5 复现 CoDi 训练,并在 Inpainting、InstructPix2Pix、超分辨率等条件下用 1–4 步采样快速生成高保真图像。
一、CoDi 是什么:单阶段条件扩散蒸馏
CoDi 旨在把一个无条件的扩散模型(如 Stable Diffusion)高效地蒸馏为条件扩散模型,使得模型在 Inpainting、InstructPix2Pix、深度图生成、超分辨率等条件设定下,仅用 1–4 步采样即可生成高质量图像。
与之对比,以往的条件蒸馏方法多为两阶段管线:
- 先蒸馏后微调(distillation-first):先对无条件模型做一致性蒸馏,再针对新条件微调;
- 先微调后蒸馏(fine-tuning-first):先把模型微调为条件模型,再做一致性蒸馏。
CoDi 是论文提出的首个单阶段蒸馏策略,直接从文本到图像的预训练模型出发,在加入新条件的同时完成蒸馏,一步到位地得到一个完整蒸馏的条件扩散模型。仓库 README 的架构图清晰对比了这两种范式:
在标准真实世界图像超分辨率基准上,README 报告:CoDi 仅用4 步采样即可达到原模型50 步采样的 FID 与 LPIPS 水平,明显优于此前的 guided-distillation 与 consistency model(一致性模型)方法;而在文本引导 Inpainting 这类相对简单的任务上,CoDi 论文提出的一种参数高效蒸馏(parameter-efficient distillation)方案,甚至能在 FID 与 LPIPS 指标上超越原始 50 步采样结果。
论文出处与引用信息:
CoDi: Conditional Diffusion Distillation for Higher-Fidelity and Faster Image Generation — Kangfu Mei, Mauricio Delbracio, Hossein Talebi, Zhengzhong Tu, Vishal M. Patel, Peyman Milanfar(Johns Hopkins University 与 Google Research 合作)。
二、环境准备与依赖
本仓库的训练实现基于 HuggingFace Diffusers 的 Flax(JAX)生态。核心依赖见 CoDi/requirements.txt,主要包括:
- 深度学习框架:
jax==0.4.13、jaxlib==0.4.13、flax==0.7.2、optax==0.1.7; - 模型与扩散组件:
diffusers==0.24.0、transformers==4.36.0、huggingface-hub==0.19.4; - 数据与工具:
datasets==2.15.0、torch==2.1.1、torchvision==0.16.1、tensorstore==0.1.45、orbax-checkpoint==0.2.3等。
需要注意:脚本以FlaxStableDiffusionControlNetPipeline、FlaxControlNetModel、FlaxDDPMScheduler等 Flax 组件为核心(见 train_codi_flax.py),并调用check_min_version("0.16.0.dev0")做 diffusers 最低版本校验,请按 requirements 安装匹配版本。脚本通过--from_pt从 PyTorch 检查点加载权重,因此环境中同时需要 PyTorch。
三、使用 HuggingFace 数据集训练 CoDi
这是上手最快的方式:只需把DATASET_NAME换成 HuggingFace Hub 上的数据集(可以是私有数据集)即可开始训练。README 建议优先参考jax-diffusers-event组织下的数据集(例如jax-diffusers-event/canny_diffusiondb,即基于 Canny 边缘图的 ControlNet 风格条件数据)。
3.1 完整训练命令
export HF_HOME="/data/huggingface/" export DISK_DIR="/data/huggingface/cache" export MODEL_DIR="runwayml/stable-diffusion-v1-5" export OUTPUT_DIR="/data/canny_model" export DATASET_NAME="jax-diffusers-event/canny_diffusiondb" python3 training_scripts/train_codi_flax.py \ --pretrained_model_name_or_path=$MODEL_DIR \ --output_dir=$OUTPUT_DIR \ --dataset_name=$DATASET_NAME \ --load_from_disk \ --cache_dir=$DISK_DIR \ --resolution=512 \ --learning_rate=1e-5 \ --train_batch_size=2 \ --revision="non-ema" \ --from_pt \ --max_train_steps=500000 \ --checkpointing_steps=10000 \ --dataloader_num_workers=16 \ --distill_learning_steps 50 \ --onestepode control \ --onestepode_control_params target \ --onestepode_sample_eps v_prediction \ --distill_loss consistency_x3.2 根据数据集调整列名
不同数据集的字段命名不同,需要按数据实际情况指定三组列名。例如jax-diffusers-event/canny_diffusiondb需追加:
--image_column original_image --caption_column prompt --conditioning_image transformed_image对应到 args.py 中的三个参数:--image_column(目标图像列,默认image)、--conditioning_image_column(ControlNet 条件图像列,默认conditioning_image)、--caption_column(文本提示列,默认text)。README 示例中使用的--conditioning_image对应脚本中的--conditioning_image_column,请以所下载数据集的字段名为准。
四、使用自有数据训练 CoDi
4.1 数据预处理
README 以“训练一个基于 Canny 边缘条件的 ControlNet 模型”为例,演示如何从大规模图文数据构建条件训练集。预处理脚本参考 HuggingFace community-events 的coyo_1m_dataset_preprocess.py,其流程为:
- 从 COYO-700M 数据集中挑选 100 万对图像-文本样本;
- 下载每张图像,并用 Canny 边缘检测器生成条件图像(conditioning image);
- 生成一份
meta.jsonl元数据文件,将原始图像、处理后图像与文本标题关联起来。
运行命令如下(若已将数据盘挂载到 TPU,建议把train_data_dir与cache_dir都放在挂载盘上):
python3 coyo_1m_dataset_preprocess.py \ --train_data_dir="/data/dataset" \ --cache_dir="/data" \ --max_train_samples=1000000 \ --num_proc=32预处理完成后,train_data_dir下应生成如下目录结构:
data ├── images │ ├── image_1.png │ ├── ....... │ └── image_1000000.jpeg ├── processed_images │ ├── image_1.png │ ├── ....... │ └── image_1000000.jpeg └── meta.jsonl4.2 从本地目录加载数据并训练
训练时只需把DATASET_NAME换成DATASET_DIR(指向上述数据文件夹):
export HF_HOME="/data/huggingface/" export DISK_DIR="/data/huggingface/cache" export MODEL_DIR="runwayml/stable-diffusion-v1-5" export OUTPUT_DIR="/data/canny_model" export DATASET_DIR="/data/dataset" python3 training_scripts/train_codi_flax.py \ --pretrained_model_name_or_path=$MODEL_DIR \ --output_dir=$OUTPUT_DIR \ --train_data_dir=$DATASET_DIR \ --load_from_disk \ --cache_dir=$DISK_DIR \ --resolution=512 \ --learning_rate=1e-5 \ --train_batch_size=2 \ --revision="non-ema" \ --from_pt \ --max_train_steps=500000 \ --checkpointing_steps=10000 \ --dataloader_num_workers=16 \ --distill_learning_steps 50 \ --onestepode control \ --onestepode_control_params target \ --onestepode_sample_eps v_prediction \ --distill_loss consistency_x--load_from_disk指示脚本使用datasets.load_from_disk从--train_data_dir加载此前用save_to_disk保存的数据集;--dataset_name与--train_data_dir二者只能指定其一(args.py 中有对应的 sanity check,两者同时设置或都未设置都会抛出ValueError)。
五、核心参数全解:从 args.py 看每个开关的真实含义
训练脚本的入口为 CoDi/training_scripts/train_codi_flax.py,参数解析位于 CoDi/training_scripts/args.py。除上述命令用到的参数外,以下几个参数对蒸馏结果起着决定性作用:
5.1 蒸馏专属参数(CoDi 的核心开关)
| 参数 | 默认值 | 含义与取值 |
|---|---|---|
--distill_learning_steps | 50 | 蒸馏模型学习到的采样步数。训练时把完整去噪轨迹划分为该数量的步长,见源码中skipped_schedule = num_train_timesteps // distill_learning_steps的计算(train_codi_flax.py),即每步跳过的时间步跨度 |
--onestepode | control | 在预测z_t时使用哪种模式:control表示用条件模型做单步 ODE 采样,uncontrol表示无控制信号 |
--onestepode_control_params | target | 单步 ODE 采样所用的 ControlNet 参数来源:target(使用 EMA 参数)或online(使用在线训练参数)(train_codi_flax.py) |
--onestepode_sample_eps | v_prediction | 单步 ODE 采样时 epsilon 的预测模式:v_prediction、x_prediction或epsilon(train_codi_flax.py) |
--distill_loss | consistency_x | 蒸馏损失形式:consistency_x(对预测的x0做一致性约束)或consistency_epsilon(对预测的噪声/速度场做一致性约束)(train_codi_flax.py) |
--ema_decay | 0.999 | 蒸馏过程中 EMA 参数的衰减系数(train_codi_flax.py) |
5.2 通用训练参数
--pretrained_model_name_or_path(必填):预训练模型路径或 HuggingFace Hub 模型标识,例如runwayml/stable-diffusion-v1-5。--controlnet_model_name_or_path:预训练 ControlNet 路径;不指定时 ControlNet 权重由 UNet 初始化(见 args.py)。--revision/--from_pt/--controlnet_revision/--controlnet_from_pt:模型版本分支以及是否从 PyTorch 检查点加载(Flax 加载 PyTorch 权重时需要--from_pt)。--resolution(默认512):输入图像统一缩放的分辨率。--train_batch_size(默认1):每个设备的训练批大小。--learning_rate(默认1e-4)、--scale_lr:初始学习率及其按 GPU 数/梯度累积步数/批大小的缩放开关。--lr_scheduler(默认constant):支持linear、cosine、cosine_with_restarts、polynomial、constant、constant_with_warmup。--snr_gamma:SNR 加权 gamma(建议5.0,对应论文 arXiv:2303.09556),用于重平衡损失;在 train_codi_flax.py 中按 SNR 对一致性损失加权。--max_train_steps/--num_train_epochs:总训练步数或轮数;两者会互相换算(train_codi_flax.py)。--checkpointing_steps(默认5000):每多少步保存一次检查点。--dataloader_num_workers(默认0):数据加载子进程数。--gradient_accumulation_steps(默认1):梯度累积步数,实现见cumul_grad_step的jax.lax.fori_loop(train_codi_flax.py)。--mixed_precision:no/fp16/bf16(bf16 需要 PyTorch ≥ 1.10 与 NVIDIA Ampere 架构 GPU)。--validation_prompt/--validation_image/--validation_steps:验证用的提示词与条件图像集合及验证频率;二者必须同时设置,且数量需匹配(args.py)。--report_to:目前仅支持wandb,配合--wandb_entity、--tracker_project_name使用。--streaming/--max_train_samples:流式加载大型 Hub 数据集(流式模式必须显式指定max_train_samples)。--debug:调试模式,跳过jax.pmap设备并行与梯度 all-reduce。--output_dir:输出目录,支持{timestamp}占位符,解析时替换为%Y%m%d_%H%M%S时间戳(args.py)。
六、源码级原理:train_step 中的单阶段条件蒸馏
README 的结论在 train_codi_flax.py 的train_step中有完整的代码对应,整个蒸馏过程可以拆解为三个关键步骤(源码注释中标注为 step1/step2,对应论文 Algorithm 11/12):
Step 1:用“教师路径”构造目标(论文 Algorithm 12)
- 对潜变量加噪得到
noisy_latents,并按distill_learning_steps计算“下一个时间步”next_timesteps; - 用 ControlNet(默认取EMA 参数,即
--onestepode_control_params target)+ UNet 在时间步t上预测,把预测结果转换为单步 ODE 的估计sampler_eps与sampler_x(对应论文公式 7); - 通过
hat_noisy_latents_s = alpha_s * sampler_x + sigma_s * sampler_eps完成一次跳跃式反推,得到s时刻的噪声潜变量; - 再次用 EMA ControlNet + UNet 在
s时刻预测,得到目标预测target_model_pred_x/target_model_pred_epsilon,并用scalings_for_boundary_conditions(c_skip、c_out,源码中timestep_scaling=10)做边界条件缩放。
Step 2:学生路径预测(论文 Algorithm 11)
- 对原始
noisy_latents,用在线训练中的 ControlNet 参数(params,即正在被优化的参数)与 UNet 预测,得到online_model_pred_x/online_model_pred_epsilon。
Step 3:一致性损失与 EMA 更新
- 损失为在线预测与冻结梯度(
jax.lax.stop_gradient)的目标预测之间的 MSE,可选择consistency_x或consistency_epsilon两种形式; - 此外损失中还包含一个边界回归项
beta_reg = (online_model_pred_x - stop_gradient(latents))^2,促使学生模型直接逼近真实潜变量(train_codi_flax.py); - 梯度通过
jax.value_and_grad计算,支持梯度累积、跨设备pmean求平均;训练状态TrainState额外维护ema_params,每个 step 结束后按ema_decay更新 EMA 参数(train_codi_flax.py 与 train_codi_flax.py)。
从实现看,CoDi 的“单阶段”体现在:全程只有一次针对 ControlNet 参数的梯度更新(UNet 与 VAE 保持冻结),教师目标(EMA 参数)与学生模型(在线参数)同步演进,无需先做一致性蒸馏再微调,这正是它与两阶段方法的本质区别。
七、引用与致谢
若你的工作使用了 CoDi,请按如下格式引用:
@article{mei2023conditional, title={CoDi: Conditional Diffusion Distillation for Higher-Fidelity and Faster Image Generation}, author={Mei, Kangfu and Delbracio, Mauricio and Talebi, Hossein and Tu, Zhengzhong and Patel, Vishal M and Milanfar, Peyman}, journal={arXiv preprint arXiv:2310.01407}, year={2023} }该实现基于 HuggingFace Diffusers 与 HuggingFace community-events 的 jax-controlnet-sprint 代码构建,README 已明确提示使用时应同时遵守上述项目的开源许可。仓库于 2023-12-02 发布了 CoDi 的训练脚本(README News 条目),即本仓库 training_scripts/ 下的train_codi_flax.py与args.py,可用于在 TPU/GPU 上复现本文所述的全部训练流程。
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
Ultralytics YOLO 知识蒸馏(Knowledge Distillation)实战指南:从教师模型蒸馏到轻量学生模型
Ultralytics YOLO 知识蒸馏(Knowledge Distillation)实战指南:从教师模型蒸馏到轻量学生模型 知识蒸馏(Knowledge
人工智能计算机视觉深度学习机器学习预训练DiffSynth-Studio 直接蒸馏(Direct Distill):端到端的扩散模型蒸馏加速训练指南
DiffSynth Studio 直接蒸馏(Direct Distill):端到端的扩散模型蒸馏加速训练指南 本篇技术指南围绕 DiffSynth Studio
人工智能大模型媒体生成深度学习微调4步生成高质量图像:Google扩散模型蒸馏技术全解析
4步生成高质量图像:Google扩散模型蒸馏技术全解析 你还在为扩散模型 DM 训练耗时长、采样步骤多而烦恼?本文将带你深入Google Research的扩散
人工智能深度学习NLP计算机视觉强化学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考