CoDi 条件扩散蒸馏(Conditional Diffusion Distillation)实战指南:从预训练文本到图像模型蒸馏出 1–4 步条件生成模型
2026/9/20 2:10:57 网站建设 项目流程
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/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.13jaxlib==0.4.13flax==0.7.2optax==0.1.7
  • 模型与扩散组件diffusers==0.24.0transformers==4.36.0huggingface-hub==0.19.4
  • 数据与工具datasets==2.15.0torch==2.1.1torchvision==0.16.1tensorstore==0.1.45orbax-checkpoint==0.2.3等。

需要注意:脚本以FlaxStableDiffusionControlNetPipelineFlaxControlNetModelFlaxDDPMScheduler等 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_x

3.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,其流程为:

  1. 从 COYO-700M 数据集中挑选 100 万对图像-文本样本;
  2. 下载每张图像,并用 Canny 边缘检测器生成条件图像(conditioning image);
  3. 生成一份meta.jsonl元数据文件,将原始图像、处理后图像与文本标题关联起来。

运行命令如下(若已将数据盘挂载到 TPU,建议把train_data_dircache_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.jsonl

4.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_steps50蒸馏模型学习到的采样步数。训练时把完整去噪轨迹划分为该数量的步长,见源码中skipped_schedule = num_train_timesteps // distill_learning_steps的计算(train_codi_flax.py),即每步跳过的时间步跨度
--onestepodecontrol在预测z_t时使用哪种模式:control表示用条件模型做单步 ODE 采样,uncontrol表示无控制信号
--onestepode_control_paramstarget单步 ODE 采样所用的 ControlNet 参数来源:target(使用 EMA 参数)或online(使用在线训练参数)(train_codi_flax.py)
--onestepode_sample_epsv_prediction单步 ODE 采样时 epsilon 的预测模式:v_predictionx_predictionepsilon(train_codi_flax.py)
--distill_lossconsistency_x蒸馏损失形式:consistency_x(对预测的x0做一致性约束)或consistency_epsilon(对预测的噪声/速度场做一致性约束)(train_codi_flax.py)
--ema_decay0.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):支持linearcosinecosine_with_restartspolynomialconstantconstant_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_stepjax.lax.fori_loop(train_codi_flax.py)。
  • --mixed_precisionno/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_epssampler_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_conditionsc_skipc_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_xconsistency_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.pyargs.py,可用于在 TPU/GPU 上复现本文所述的全部训练流程。

  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询