尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

CoDi 条件扩散蒸馏(Conditional Diffusion Distillation)实战指南:从预训练文本到图像模型蒸馏出 1–4 步条件生成模型

CoDi 条件扩散蒸馏(Conditional Diffusion Distillation)实战指南:从预训练文本到图像模型蒸馏出 1–4 步条件生成模型 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载CoDiConditional 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 MilanfarJohns Hopkins University 与 Google Research 合作。二、环境准备与依赖本仓库的训练实现基于 HuggingFace Diffusers 的 FlaxJAX生态。核心依赖见 CoDi/requirements.txt主要包括深度学习框架jax0.4.13、jaxlib0.4.13、flax0.7.2、optax0.1.7模型与扩散组件diffusers0.24.0、transformers4.36.0、huggingface-hub0.19.4数据与工具datasets2.15.0、torch2.1.1、torchvision0.16.1、tensorstore0.1.45、orbax-checkpoint0.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_DIRrunwayml/stable-diffusion-v1-5 export OUTPUT_DIR/data/canny_model export DATASET_NAMEjax-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 \ --resolution512 \ --learning_rate1e-5 \ --train_batch_size2 \ --revisionnon-ema \ --from_pt \ --max_train_steps500000 \ --checkpointing_steps10000 \ --dataloader_num_workers16 \ --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_columnControlNet 条件图像列默认conditioning_image、--caption_column文本提示列默认text。README 示例中使用的--conditioning_image对应脚本中的--conditioning_image_column请以所下载数据集的字段名为准。四、使用自有数据训练 CoDi4.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_samples1000000 \ --num_proc32预处理完成后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_DIRrunwayml/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 \ --resolution512 \ --learning_rate1e-5 \ --train_batch_size2 \ --revisionnon-ema \ --from_pt \ --max_train_steps500000 \ --checkpointing_steps10000 \ --dataloader_num_workers16 \ --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_prediction、x_prediction或epsilontrain_codi_flax.py--distill_lossconsistency_x蒸馏损失形式consistency_x对预测的x0做一致性约束或consistency_epsilon对预测的噪声/速度场做一致性约束train_codi_flax.py--ema_decay0.999蒸馏过程中 EMA 参数的衰减系数train_codi_flax.py5.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_gammaSNR 加权 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_looptrain_codi_flax.py。--mixed_precisionno/fp16/bf16bf16 需要 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/12Step 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_conditionsc_skip、c_out源码中timestep_scaling10做边界条件缩放。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-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/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),仅供参考
返回列表