Diffusers 潜空间一致性蒸馏(Latent Consistency Distillation)完整训练指南:从 Stable Diffusion 教师模型到少步数 LCM
【免费下载链接】diffusers🤗 Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers
潜空间一致性模型(Latent Consistency Models,LCM)将传统扩散模型的数十步去噪压缩到 4~8 步即可生成高质量图像。本文基于 🤗 Diffusers 仓库中的官方训练指南 lcm_distill.md 及示例脚本 train_lcm_distill_sd_wds.py,系统讲解如何对 Stable Diffusion 教师模型执行潜空间一致性蒸馏,覆盖原理、环境搭建、参数解析、训练循环逐步拆解、完整启动命令、推理,以及 LCM-LoRA 与 SDXL 变体,帮助读者独立训练出属于自己的少步数推理模型。
LCM 蒸馏的核心原理
Latent Consistency Models(LCM)之所以能在极少步数内生成高质量图像,是因为其训练方法——潜空间一致性蒸馏(Latent Consistency Distillation,LCD)——直接作用于扩散模型的潜空间(latent space)。传统扩散管线通常需要 25 步以上去噪,而 LCM 显著改变了这一局面。
蒸馏过程包含两个关键技术手段(对应论文 4.1、4.2、4.3 节):
- 单阶段引导蒸馏(one-stage guided distillation):让学生模型直接学习"从任意噪声点一步预测干净样本"的一致性映射,而不是像普通蒸馏那样逐步模仿教师轨迹;
- 跳步方法(skipping-step):在蒸馏过程中刻意跳过部分时间步,让一致性训练更高效地覆盖整个采样轨迹。
在仓库的示例脚本中,这两点分别体现在"边界缩放系数(boundary scalings)"计算与"DDIM ODE 求解器"的构造上,下文训练循环部分会详细展开。
环境准备与依赖安装
从源码安装 Diffusers
示例脚本随仓库持续更新,官方推荐从源码安装以保证脚本与库版本匹配。在当前仓库环境下,执行:
git clone https://github.com/huggingface/diffusers cd diffusers pip install .安装训练依赖
进入示例目录并安装该脚本所需的依赖:
cd examples/consistency_distillation pip install -r requirements.txtrequirements.txt 中声明的核心依赖包括:
| 依赖 | 版本要求 | 用途 |
|---|---|---|
accelerate | >=0.16.0 | 多卡 / 混合精度训练调度 |
transformers | >=4.25.1 | CLIP 文本编码器与分词器 |
webdataset | 无 | WebDataset 流式数据读取 |
torchvision | 无 | 图像预处理变换 |
ftfy/Jinja2 | 无 | 文本清洗与模板 |
tensorboard | 无 | 训练日志可视化 |
若使用 LoRA 脚本,还需另行安装peft;使用 8-bit Adam 优化器时需要bitsandbytes。
配置 Accelerate 环境
🤗 Accelerate 负责根据硬件自动配置多 GPU / TPU 训练与混合精度。有几种初始化方式:
交互式配置(推荐,可启用torch.compile显著加速训练):
accelerate config使用默认配置(不回答任何提问):
accelerate config default在笔记本等不支持交互式 Shell 的环境中,用 Python 方式写入基础配置:
from accelerate.utils import write_basic_config write_basic_config()[!TIP] 若你的 GPU 显存有限,可开启
--gradient_checkpointing(梯度检查点)、--gradient_accumulation_steps(梯度累积)与--mixed_precision(混合精度)来降低显存占用并加速训练;进一步可通过 xFormers 内存高效注意力 与 bitsandbytes 8-bit 优化器(--use_8bit_adam)继续压减显存。
训练脚本参数全解
全部参数定义集中在 parse_args() 函数中,每个参数都有默认值。其中绝大多数参数与 Text-to-image 训练指南 中一致(如--train_batch_size默认 16、--learning_rate默认 1e-4、--resolution默认 512、--lr_scheduler默认constant、--adam_beta1默认 0.9 等),本文聚焦于潜空间一致性蒸馏特有的参数:
| 参数 | 默认值 | 说明 |
|---|---|---|
--pretrained_teacher_model | 无(必填) | 教师模型路径,即待蒸馏的预训练潜在扩散模型(如 stable-diffusion-v1-5) |
--pretrained_vae_model_name_or_path | None | 替代 VAE 路径。SDXL 自带 VAE 存在数值不稳定问题,可用 madebyollin 的 fp16 修复版 VAE 替代,使其在 fp16 下稳定工作 |
--w_min/--w_max | 5.0/15.0 | 引导尺度(guidance scale)采样的最小 / 最大值,训练时从U[w_min, w_max]均匀采样。注意脚本采用 Imagen CFG 公式,所有引导尺度相比原论文都加了 1 |
--num_ddim_timesteps | 50 | DDIM 采样使用的时间步数量,决定 ODE 求解器轨迹的离散粒度 |
--loss_type | l2 | 蒸馏损失类型,可选l2或huber。Huber 损失对离群点更鲁棒,实践上更受推荐 |
--huber_c | 0.001 | Huber 损失参数,仅当--loss_type=huber时生效 |
--unet_time_cond_proj_dim | 256 | 学生 U-Net 中引导尺度嵌入(time_cond_proj)的维度;当教师 U-Net 未配置time_cond_proj_dim时使用 |
--timestep_scaling_factor | 10.0 | 计算 LCM 边界缩放时的乘法时间步缩放因子。取值越大近似误差越小,默认 10.0 通常足够 |
--vae_encode_batch_size | 32 | VAE 编码 / 解码图像的批大小。一次性编码整个 batch 可能 OOM,拆小批处理更稳妥 |
--ema_decay | 0.95 | 目标学生模型(target student U-Net)的指数移动平均衰减率 |
--cast_teacher_unet | False | 是否将教师 U-Net 转换为--mixed_precision指定的精度 |
--teacher_revision | None | 教师模型的 revision(用于从 Hub 拉取指定版本) |
--proportion_empty_prompts | 0 | 将图像提示替换为空字符串的比例(0~1),配合 CFG 无条件分支使用 |
--allow_tf32 | False | 是否在 Ampere GPU 上启用 TF32 以加速训练 |
此外,训练常规参数还包括:--output_dir(默认lcm-xl-distilled)、--checkpointing_steps(默认 500)、--checkpoints_total_limit、--resume_from_checkpoint(支持latest自动选择最新检查点)、--report_to(tensorboard/wandb/comet_ml)、--validation_steps(默认 200)、--push_to_hub、--hub_model_id、--seed等。
例如,仅需在启动命令中加入以下参数即可开启 fp16 混合精度加速训练:
accelerate launch train_lcm_distill_sd_wds.py \ --mixed_precision="fp16"训练脚本逐步拆解
数据集类与 WebDataset 预处理流水线
脚本首先定义数据集类SDText2ImageDataset(源码 train_lcm_distill_sd_wds.py),负责图像预处理与训练数据集构建。核心的图像变换逻辑如下:
def transform(example): image = example["image"] image = TF.resize(image, resolution, interpolation=interpolation_mode) c_top, c_left, _, _ = transforms.RandomCrop.get_params(image, output_size=(resolution, resolution)) image = TF.crop(image, c_top, c_left, resolution, resolution) image = TF.to_tensor(image) image = TF.normalize(image, [0.5], [0.5]) example["image"] = image return example即:先缩放到目标分辨率,再做随机裁剪,最后归一化到[-1, 1]。插值方式由--interpolation_type控制,可选bilinear、bicubic、box、nearest、nearest_exact、hamming、lanczos。
针对云端大规模数据集,脚本采用WebDataset 格式构建流式预处理流水线——图像按需解码、处理后直接进入训练循环,无需预先下载整个数据集:
processing_pipeline = [ wds.decode("pil", handler=wds.ignore_and_continue), wds.rename(image="jpg;png;jpeg;webp", text="text;txt;caption", handler=wds.warn_and_continue), wds.map(filter_keys({"image", "text"})), wds.map(transform), wds.to_tuple("image", "text"), ]该流水线依次完成:PIL 解码(忽略坏样本)、按扩展名重命名字段、过滤仅保留image/text、应用变换、输出元组。数据管线再经由wds.ResampledShards(无限重采样分片)、tarfile_to_samples_nothrow(容错解包 tar 分片)、wds.shuffle(1000 样本洗牌缓冲)与wds.batched组装成wds.WebLoader。值得注意,脚本对 webdataset 默认的group_by_keys做了不抛异常的重新实现(group_by_keys_nothrow),避免个别损坏样本中断整个训练。
组件加载与学生网络创建
在 main() 函数中依次完成组件装配:
- 从教师模型加载
DDPMScheduler,并由其alphas_cumprod推导出alpha_schedule = sqrt(alphas_cumprod)与sigma_schedule = sqrt(1 - alphas_cumprod); - 实例化
DDIMSolver(源码 L394-L418),它基于离散化的 DDIM 时间步预先计算alpha_cumprod与前一时刻的alpha_cumprod_prev,供训练中单步 ODE 求解使用; - 加载分词器(
AutoTokenizer)、文本编码器(CLIPTextModel)与 VAE(AutoencoderKL); - 加载教师 U-Net,并冻结 VAE、文本编码器与教师 U-Net(
requires_grad_(False)); - 创建在线学生 U-Net(由优化器更新):若教师 U-Net 没有
time_cond_proj_dim配置,则按--unet_time_cond_proj_dim添加引导尺度嵌入投影层,再从教师权重初始化:
teacher_unet = UNet2DConditionModel.from_pretrained( args.pretrained_teacher_model, subfolder="unet", revision=args.teacher_revision ) time_cond_proj_dim = ( teacher_unet.config.time_cond_proj_dim if teacher_unet.config.time_cond_proj_dim is not None else args.unet_time_cond_proj_dim ) unet = UNet2DConditionModel.from_config(teacher_unet.config, time_cond_proj_dim=time_cond_proj_dim) unet.load_state_dict(teacher_unet.state_dict(), strict=False) unet.train()- 创建目标学生 U-Net(target student),由在线学生网络初始化,之后只通过 EMA(Polyak 平均)更新、不参与梯度计算:
target_unet = UNet2DConditionModel.from_config(unet.config) target_unet.load_state_dict(unet.state_dict()) target_unet.train() target_unet.requires_grad_(False)EMA 更新逻辑位于 update_ema():每个同步梯度步后,target = rate * target + (1 - rate) * online,衰减率由--ema_decay控制(默认 0.95)。
优化器与数据集装配
优化器只作用于在线学生 U-Net 参数(源码 L1063-L1070):
optimizer = optimizer_class( unet.parameters(), lr=args.learning_rate, betas=(args.adam_beta1, args.adam_beta2), weight_decay=args.adam_weight_decay, eps=args.adam_epsilon, )其中optimizer_class在启用--use_8bit_adam时为bnb.optim.AdamW8bit,否则为torch.optim.AdamW。
数据集创建(源码 L1079-L1091):
dataset = SDText2ImageDataset( train_shards_path_or_url=args.train_shards_path_or_url, num_train_examples=args.max_train_samples, per_gpu_batch_size=args.train_batch_size, global_batch_size=args.train_batch_size * accelerator.num_processes, num_workers=args.dataloader_num_workers, resolution=args.resolution, interpolation_type=args.interpolation_type, shuffle_buffer_size=1000, pin_memory=True, persistent_workers=True, ) train_dataloader = dataset.train_dataloader注意Accelerator构造时设置了split_batches=True——这对 webdataset 至关重要,否则学习率调度的步数计算会因批次被多进程拆分而出错。
训练循环中的一致性蒸馏实现
训练循环(源码 L1185 起)对应论文 Algorithm 1 的完整流程,每一步骤如下:
① 潜变量编码。图像像素值以不超过--vae_encode_batch_size的批大小送入 VAE 编码器,采样潜变量后乘以vae.config.scaling_factor。
② 时间步采样与跳步。先按topk = num_train_timesteps // num_ddim_timesteps计算跳步间隔,再从num_ddim_timesteps个离散 ODE 步中均匀随机采样起点start_timesteps,目标时间步为timesteps = start_timesteps - topk(小于 0 时截断为 0)——这就是"跳步"加速蒸馏的体现。
③ 边界缩放。调用 scalings_for_boundary_conditions()(与LCMScheduler.get_scalings_for_boundary_condition_discrete同源)计算起点与终点的c_skip、c_out:
def scalings_for_boundary_conditions(timestep, sigma_data=0.5, timestep_scaling=10.0): scaled_timestep = timestep_scaling * timestep c_skip = sigma_data**2 / (scaled_timestep**2 + sigma_data**2) c_out = scaled_timestep / (scaled_timestep**2 + sigma_data**2) ** 0.5 return c_skip, c_out④ 加噪。采样高斯噪声并执行前向扩散noisy_model_input = noise_scheduler.add_noise(latents, noise, start_timesteps)。
⑤ 引导尺度采样与嵌入。从U[w_min, w_max]均匀采样引导尺度w,再由 guidance_scale_embedding()(源自LatentConsistencyModel.get_guidance_scale_embedding,与 VDM 同源的正余弦位置编码)生成维度为time_cond_proj_dim的引导尺度嵌入,作为timestep_cond输入 U-Net。
⑥ 在线学生预测。学生 U-Net 在加噪潜变量z_{t_{n+k}}上输出噪声预测,再结合预测类型(epsilon/sample/v_prediction)通过get_predicted_original_sample还原原始样本预测,最终合成一致性模型输出:
pred_x_0 = get_predicted_original_sample( noise_pred, start_timesteps, noisy_model_input, noise_scheduler.config.prediction_type, alpha_schedule, sigma_schedule, ) model_pred = c_skip_start * noisy_model_input + c_out_start * pred_x_0⑦ 教师 CFG 预测与 ODE 求解。在torch.no_grad()下,教师 U-Net 分别对条件嵌入与无条件嵌入做预测,得到各自的原样本预测与噪声预测,再按 LCM 论文的 CFG 公式合成:
pred_x0 = cond_pred_x0 + w * (cond_pred_x0 - uncond_pred_x0) pred_noise = cond_pred_noise + w * (cond_pred_noise - uncond_pred_noise) x_prev = solver.ddim_step(pred_x0, pred_noise, index)ddim_step依据 DDIM 反演公式x_prev = sqrt(alpha_prev) * pred_x0 + sqrt(1 - alpha_prev) * pred_noise前进一步,得到增强 PF-ODE 轨迹上的下一点x_prev。
⑧ 目标学生预测。目标学生 U-Net 在x_prev、时间步t_n与同一引导尺度嵌入下再次输出,得到一致性回归目标:
target = c_skip * x_prev + c_out * pred_x_0⑨ 损失计算与反向传播。对model_pred与target计算蒸馏损失(源码 L1352-L1358):
if args.loss_type == "l2": loss = F.mse_loss(model_pred.float(), target.float(), reduction="mean") elif args.loss_type == "huber": loss = torch.mean( torch.sqrt((model_pred.float() - target.float()) ** 2 + args.huber_c**2) - args.huber_c )Huber 损失对离群点更稳健,这也是官方示例命令默认选用--loss_type="huber"的原因。随后accelerator.backward(loss)反向传播、按--max_grad_norm裁剪梯度、优化器步进,并在sync_gradients时对目标学生网络执行 EMA 更新。
检查点与验证
- 脚本通过
accelerate的register_save_state_pre_hook/register_load_state_pre_hook自定义序列化格式,将unet与unet_target分开保存为 diffusers 原生格式,训练中断后可用--resume_from_checkpoint=latest恢复; - 每
--checkpointing_steps步保存检查点,--checkpoints_total_limit控制保留数量(超出时自动删除最旧检查点); - 每
--validation_steps步调用 log_validation(),用LCMScheduler以 4 步采样对一组验证提示词(如 "Astronaut in a jungle...")生成图像,同时记录在线网络与目标网络(EMA)两套结果,可上报到 TensorBoard 或 wandb。
若想深入理解去噪循环的基本范式,可参考 Understanding pipelines, models and schedulers tutorial。
启动训练
下面的命令以Conceptual Captions 12M(CC12M)数据集的 webdataset 分片为例(数据通过--train_shards_path_or_url以pipe:前缀流式拉取),教师模型选用 stable-diffusion-v1-5。使用环境变量管理模型与输出路径:
export MODEL_DIR="stable-diffusion-v1-5/stable-diffusion-v1-5" export OUTPUT_DIR="path/to/saved/model" accelerate launch train_lcm_distill_sd_wds.py \ --pretrained_teacher_model=$MODEL_DIR \ --output_dir=$OUTPUT_DIR \ --mixed_precision=fp16 \ --resolution=512 \ --learning_rate=1e-6 --loss_type="huber" --ema_decay=0.95 --adam_weight_decay=0.0 \ --max_train_steps=1000 \ --max_train_samples=4000000 \ --dataloader_num_workers=8 \ --train_shards_path_or_url="pipe:curl -L -s https://huggingface.co/datasets/laion/conceptual-captions-12m-webdataset/resolve/main/data/{00000..01099}.tar?download=true" \ --validation_steps=200 \ --checkpointing_steps=200 --checkpoints_total_limit=10 \ --train_batch_size=12 \ --gradient_checkpointing --enable_xformers_memory_efficient_attention \ --gradient_accumulation_steps=1 \ --use_8bit_adam \ --resume_from_checkpoint=latest \ --report_to=wandb \ --seed=453645634 \ --push_to_hub关键参数速览:--mixed_precision=fp16开启混合精度;--learning_rate=1e-6使用较低学习率;--loss_type="huber"采用更稳健的损失;--train_shards_path_or_url的花括号{00000..01099}会被braceexpand展开为 1100 个 tar 分片地址;--push_to_hub会在训练结束后把产物上传到 Hub(需提前通过hf auth login完成认证,且注意脚本禁止同时使用--report_to=wandb与--hub_token,以免令牌泄露风险)。训练完成后,unet与unet_target(EMA 版本)都会以 diffusers 原生格式保存到OUTPUT_DIR。
若需要准备自己的训练数据,可参考 Create a dataset for training 指南,构建与脚本兼容的 webdataset 格式数据集。
用蒸馏产物进行推理
训练完成后,用训练好的学生 U-Net 替换 Stable Diffusion 管线中的 U-Net,并将调度器切换为LCMScheduler,即可用 4 步完成采样:
from diffusers import UNet2DConditionModel, DiffusionPipeline, LCMScheduler import torch unet = UNet2DConditionModel.from_pretrained("your-username/your-model", dtype=torch.float16, variant="fp16") pipeline = DiffusionPipeline.from_pretrained("stable-diffusion-v1-5/stable-diffusion-v1-5", unet=unet, dtype=torch.float16, variant="fp16") pipeline.scheduler = LCMScheduler.from_config(pipe.scheduler.config) pipeline.to("cuda") # or "mps", "xpu", "cpu" prompt = "sushi rolls in the form of panda heads, sushi platter" image = pipeline(prompt, num_inference_steps=4, guidance_scale=1.0).images[0]由于 LCM 已把引导尺度信息蒸馏进模型,推理时guidance_scale只需设为 1.0。LCMScheduler的实现位于 scheduling_lcm.py,其step方法同样基于c_skip/c_out边界缩放完成单步去噪,与训练目标网络的计算方式一致;配套的完整管线 LatentConsistencyModelPipeline 默认num_inference_steps=4,也支持 img2img 与 LoRA 检查点组合使用。
轻量变体:LCM-LoRA 与 SDXL
LCM-LoRA
LoRA 技术可以显著减少可训练参数量,训练更快、产物更小(约 100MB 量级),且可注入任意同架构模型。仓库提供了两个 LoRA 变体脚本:
- train_lcm_distill_lora_sd_wds.py:面向 Stable Diffusion 1.x;
- train_lcm_distill_lora_sdxl_wds.py:面向 SDXL。
LoRA 脚本基于peft的LoraConfig/get_peft_model实现(如--lora_rank控制秩,示例中为 64),训练命令与全量蒸馏几乎一致,仅将脚本替换为 LoRA 版本并加入--lora_rank=64等参数。其完整说明见 LoRA training 指南。
Stable Diffusion XL
SDXL 是强大的高分辨率文生图模型,架构上增加了一个文本编码器(CLIP 双塔)。使用 train_lcm_distill_sdxl_wds.py 即可对 SDXL 执行蒸馏(该脚本在计算嵌入时额外生成 SDXL U-Net 所需的added_cond_kwargs)。由于 SDXL 自带 VAE 存在数值不稳定性,强烈建议通过--pretrained_vae_model_name_or_path指定数值更稳定的替代 VAE(如 madebyollin 的 fp16 修复版)。详细说明见 SDXL training 指南。
下一步:进阶学习路径
- 阅读 Latent Consistency Models 管线文档,掌握 LCM 在文生图、图生图及 LoRA 检查点场景下的完整 API 用法;
- 若对 LCM 论文细节感兴趣,可研读原论文中关于单阶段引导蒸馏与跳步方法的设计动机与消融实验;
- 对照 一致性蒸馏示例目录 与 示例脚本,在理解训练循环的每一步后,尝试按自己的数据集与超参数组合复现实验。
【免费下载链接】diffusers🤗 Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考