完整訓練指南:從 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),僅供參考