ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Diffusers 潜空间一致性蒸馏(Latent Consistency Distillation)完整训练指南:从 Stable Diffusion 教师模型到少步数 LCM

Diffusers 潜空间一致性蒸馏(Latent Consistency Distillation)完整训练指南:从 Stable Diffusion 教师模型到少步数 LCM 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 ModelsLCM将传统扩散模型的数十步去噪压缩到 48 步即可生成高质量图像。本文基于 Diffusers 仓库中的官方训练指南 lcm_distill.md 及示例脚本 train_lcm_distill_sd_wds.py系统讲解如何对 Stable Diffusion 教师模型执行潜空间一致性蒸馏覆盖原理、环境搭建、参数解析、训练循环逐步拆解、完整启动命令、推理以及 LCM-LoRA 与 SDXL 变体帮助读者独立训练出属于自己的少步数推理模型。LCM 蒸馏的核心原理Latent Consistency ModelsLCM之所以能在极少步数内生成高质量图像是因为其训练方法——潜空间一致性蒸馏Latent Consistency DistillationLCD——直接作用于扩散模型的潜空间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 中声明的核心依赖包括依赖版本要求用途accelerate0.16.0多卡 / 混合精度训练调度transformers4.25.1CLIP 文本编码器与分词器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_pathNone替代 VAE 路径。SDXL 自带 VAE 存在数值不稳定问题可用 madebyollin 的 fp16 修复版 VAE 替代使其在 fp16 下稳定工作--w_min/--w_max5.0/15.0引导尺度guidance scale采样的最小 / 最大值训练时从U[w_min, w_max]均匀采样。注意脚本采用 Imagen CFG 公式所有引导尺度相比原论文都加了 1--num_ddim_timesteps50DDIM 采样使用的时间步数量决定 ODE 求解器轨迹的离散粒度--loss_typel2蒸馏损失类型可选l2或huber。Huber 损失对离群点更鲁棒实践上更受推荐--huber_c0.001Huber 损失参数仅当--loss_typehuber时生效--unet_time_cond_proj_dim256学生 U-Net 中引导尺度嵌入time_cond_proj的维度当教师 U-Net 未配置time_cond_proj_dim时使用--timestep_scaling_factor10.0计算 LCM 边界缩放时的乘法时间步缩放因子。取值越大近似误差越小默认 10.0 通常足够--vae_encode_batch_size32VAE 编码 / 解码图像的批大小。一次性编码整个 batch 可能 OOM拆小批处理更稳妥--ema_decay0.95目标学生模型target student U-Net的指数移动平均衰减率--cast_teacher_unetFalse是否将教师 U-Net 转换为--mixed_precision指定的精度--teacher_revisionNone教师模型的 revision用于从 Hub 拉取指定版本--proportion_empty_prompts0将图像提示替换为空字符串的比例01配合 CFG 无条件分支使用--allow_tf32False是否在 Ampere GPU 上启用 TF32 以加速训练此外训练常规参数还包括--output_dir默认lcm-xl-distilled、--checkpointing_steps默认 500、--checkpoints_total_limit、--resume_from_checkpoint支持latest自动选择最新检查点、--report_totensorboard/wandb/comet_ml、--validation_steps默认 200、--push_to_hub、--hub_model_id、--seed等。例如仅需在启动命令中加入以下参数即可开启 fp16 混合精度加速训练accelerate launch train_lcm_distill_sd_wds.py \ --mixed_precisionfp16训练脚本逐步拆解数据集类与 WebDataset 预处理流水线脚本首先定义数据集类SDText2ImageDataset源码 train_lcm_distill_sd_wds.py负责图像预处理与训练数据集构建。核心的图像变换逻辑如下def transform(example): image example[image] image TF.resize(image, resolution, interpolationinterpolation_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, handlerwds.ignore_and_continue), wds.rename(imagejpg;png;jpeg;webp, texttext;txt;caption, handlerwds.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.shuffle1000 样本洗牌缓冲与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与 VAEAutoencoderKL加载教师 U-Net并冻结 VAE、文本编码器与教师 U-Netrequires_grad_(False)创建在线学生 U-Net由优化器更新若教师 U-Net 没有time_cond_proj_dim配置则按--unet_time_cond_proj_dim添加引导尺度嵌入投影层再从教师权重初始化teacher_unet UNet2DConditionModel.from_pretrained( args.pretrained_teacher_model, subfolderunet, revisionargs.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_dimtime_cond_proj_dim) unet.load_state_dict(teacher_unet.state_dict(), strictFalse) unet.train()创建目标学生 U-Nettarget student由在线学生网络初始化之后只通过 EMAPolyak 平均更新、不参与梯度计算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-L1070optimizer optimizer_class( unet.parameters(), lrargs.learning_rate, betas(args.adam_beta1, args.adam_beta2), weight_decayargs.adam_weight_decay, epsargs.adam_epsilon, )其中optimizer_class在启用--use_8bit_adam时为bnb.optim.AdamW8bit否则为torch.optim.AdamW。数据集创建源码 L1079-L1091dataset SDText2ImageDataset( train_shards_path_or_urlargs.train_shards_path_or_url, num_train_examplesargs.max_train_samples, per_gpu_batch_sizeargs.train_batch_size, global_batch_sizeargs.train_batch_size * accelerator.num_processes, num_workersargs.dataloader_num_workers, resolutionargs.resolution, interpolation_typeargs.interpolation_type, shuffle_buffer_size1000, pin_memoryTrue, persistent_workersTrue, ) train_dataloader dataset.train_dataloader注意Accelerator构造时设置了split_batchesTrue——这对 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_outdef scalings_for_boundary_conditions(timestep, sigma_data0.5, timestep_scaling10.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_{nk}}上输出噪声预测再结合预测类型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-L1358if args.loss_type l2: loss F.mse_loss(model_pred.float(), target.float(), reductionmean) 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_typehuber的原因。随后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_checkpointlatest恢复每--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 12MCC12M数据集的 webdataset 分片为例数据通过--train_shards_path_or_url以pipe:前缀流式拉取教师模型选用 stable-diffusion-v1-5。使用环境变量管理模型与输出路径export MODEL_DIRstable-diffusion-v1-5/stable-diffusion-v1-5 export OUTPUT_DIRpath/to/saved/model accelerate launch train_lcm_distill_sd_wds.py \ --pretrained_teacher_model$MODEL_DIR \ --output_dir$OUTPUT_DIR \ --mixed_precisionfp16 \ --resolution512 \ --learning_rate1e-6 --loss_typehuber --ema_decay0.95 --adam_weight_decay0.0 \ --max_train_steps1000 \ --max_train_samples4000000 \ --dataloader_num_workers8 \ --train_shards_path_or_urlpipe:curl -L -s https://huggingface.co/datasets/laion/conceptual-captions-12m-webdataset/resolve/main/data/{00000..01099}.tar?downloadtrue \ --validation_steps200 \ --checkpointing_steps200 --checkpoints_total_limit10 \ --train_batch_size12 \ --gradient_checkpointing --enable_xformers_memory_efficient_attention \ --gradient_accumulation_steps1 \ --use_8bit_adam \ --resume_from_checkpointlatest \ --report_towandb \ --seed453645634 \ --push_to_hub关键参数速览--mixed_precisionfp16开启混合精度--learning_rate1e-6使用较低学习率--loss_typehuber采用更稳健的损失--train_shards_path_or_url的花括号{00000..01099}会被braceexpand展开为 1100 个 tar 分片地址--push_to_hub会在训练结束后把产物上传到 Hub需提前通过hf auth login完成认证且注意脚本禁止同时使用--report_towandb与--hub_token以免令牌泄露风险。训练完成后unet与unet_targetEMA 版本都会以 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, dtypetorch.float16, variantfp16) pipeline DiffusionPipeline.from_pretrained(stable-diffusion-v1-5/stable-diffusion-v1-5, unetunet, dtypetorch.float16, variantfp16) 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_steps4, guidance_scale1.0).images[0]由于 LCM 已把引导尺度信息蒸馏进模型推理时guidance_scale只需设为 1.0。LCMScheduler的实现位于 scheduling_lcm.py其step方法同样基于c_skip/c_out边界缩放完成单步去噪与训练目标网络的计算方式一致配套的完整管线 LatentConsistencyModelPipeline 默认num_inference_steps4也支持 img2img 与 LoRA 检查点组合使用。轻量变体LCM-LoRA 与 SDXLLCM-LoRALoRA 技术可以显著减少可训练参数量训练更快、产物更小约 100MB 量级且可注入任意同架构模型。仓库提供了两个 LoRA 变体脚本train_lcm_distill_lora_sd_wds.py面向 Stable Diffusion 1.xtrain_lcm_distill_lora_sdxl_wds.py面向 SDXL。LoRA 脚本基于peft的LoraConfig/get_peft_model实现如--lora_rank控制秩示例中为 64训练命令与全量蒸馏几乎一致仅将脚本替换为 LoRA 版本并加入--lora_rank64等参数。其完整说明见 LoRA training 指南。Stable Diffusion XLSDXL 是强大的高分辨率文生图模型架构上增加了一个文本编码器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),仅供参考
返回列表