使用 Diffusers 训练文生图模型:train_text_to_image.py 全流程实战指南
【免费下载链接】diffusers🤗 Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers
本文以 docs/source/en/training/text2image.md 为骨架,结合仓库中的训练脚本 examples/text_to_image/train_text_to_image.py 源码,系统讲解基于 Stable Diffusion 的文生图(text-to-image)全量微调流程:从环境安装、数据集准备、脚本参数解析,到训练循环的底层原理与推理部署。读完本文,你将能够在单卡或分布式环境下,使用公开数据集或自建数据集完成一次完整的文生图模型微调,并掌握 checkpoint 恢复、Min-SNR 加权、EMA 等关键进阶技巧。
概述:文生图微调能做什么
文生图模型(如 Stable Diffusion)以文本提示(text prompt)为条件生成对应图像。其训练目标是让模型学会「给定文本 → 生成符合语义的图像」。本文介绍的train_text_to_image.py脚本对模型进行全量微调(fine-tune 整个 UNet 权重),与仅训练 LoRA 适配器的方式相比,可以更充分地适应目标数据分布,但也更容易过拟合。
⚠️ 注意:该脚本是实验性的。全量微调很容易过拟合,并可能引发灾难性遗忘(catastrophic forgetting)等问题。建议针对你的数据集尝试不同的超参数组合以获得最佳效果。
硬件要求与显存优化
训练模型对硬件有一定要求,但通过两个关键开关可以显著降低门槛:
gradient_checkpointing(梯度检查点):以更慢的反向传播换取显存,把激活值按需重算而非全部驻留显存;mixed_precision(混合精度):将计算精度降为 fp16/bf16,减少显存占用并加速。
在这两项开启的情况下,在单张 24GB 显存的 GPU 上即可完成训练。如果需要更大的 batch size 或更快的训练速度,建议使用 30GB 以上显存的 GPU。此外,还可以通过启用 xFormers 内存高效注意力(memory-efficient attention)进一步压缩显存足迹,详见 xFormers 优化指南。
环境准备与安装
从源码安装 Diffusers
由于训练脚本迭代频繁,官方强烈建议从源码安装并保持更新。在新建的虚拟环境中执行:
git clone https://github.com/huggingface/diffusers cd diffusers pip install .说明:以上命令中的仓库地址为公开的 diffusers 上游仓库;若在本地环境中已有本仓库副本,可直接在仓库根目录执行
pip install .。本文所引用的全部脚本路径均以当前仓库为准。
然后进入示例目录并安装该训练脚本所需的依赖:
cd examples/text_to_image pip install -r requirements.txtrequirements.txt 中的核心依赖包括:
accelerate>=0.16.0:负责多 GPU/TPU 分布式训练与混合精度配置;transformers>=4.25.1:提供 CLIP 文本编码器与分词器;datasets>=2.19.1:加载与预处理数据集;torchvision:图像变换(resize、crop、flip 等);tensorboard:默认的训练日志记录后端;peft>=0.17.0:作为 LoRA 训练的底层后端(本脚本虽为全量微调,但共享环境依赖);ftfy、Jinja2:文本修复与模型卡渲染辅助库。
配置 Accelerate 环境
🤗 Accelerate 是帮助你进行多 GPU/TPU 训练或混合精度训练的库,它会根据你的硬件和环境自动配置训练方案。初始化方式有三种:
方式一:交互式配置
accelerate config按提示选择设备类型、显存、混合精度、分布式策略等。
方式二:使用默认配置
accelerate config default跳过所有交互选项,直接生成默认配置。
方式三:在 Notebook 等非交互环境中配置
from accelerate.utils import write_basic_config write_basic_config()💡 若要训练自有数据集,请先阅读 创建训练数据集指南,了解如何构造能被训练脚本直接消费的数据集格式(如
imagefolder+metadata.jsonl或直接上传到 Hub)。
脚本参数解析:parse_args()
训练脚本提供了大量参数用于定制训练过程。全部参数及其说明集中在parse_args()函数中(源码见 train_text_to_image.py 的parse_args定义,约 L201-L525)。每个参数都带默认值(如训练 batch size、学习率等),你也可以在启动命令中通过--参数名 值覆盖。
例如,要使用 fp16 混合精度加速训练:
accelerate launch train_text_to_image.py \ --mixed_precision="fp16"基础且重要的参数
| 参数 | 说明 | 默认值 |
|---|---|---|
--pretrained_model_name_or_path | Hub 上的模型名或本地预训练模型路径(必填) | 无 |
--dataset_name | Hub 上的数据集名,或本地数据集路径 | 无 |
--image_column | 数据集中存放图像的列名 | "image" |
--caption_column | 数据集中存放文本描述的列名 | "text" |
--output_dir | 训练模型与 checkpoint 的保存目录 | "sd-model-finetuned" |
--push_to_hub | 是否将训练好的模型推送到 Hub | 关闭 |
--checkpointing_steps | 每隔多少步保存一次 checkpoint;训练中断时可结合--resume_from_checkpoint断点续训 | 500 |
训练过程核心参数
以下参数在源码中都有明确的默认值与语义(见parse_args):
--resolution(默认512):输入图像统一缩放到的分辨率。注意:若使用 768×768 版本的 stable-diffusion-2,需要改为768;--center_crop/--random_flip:先缩放到目标分辨率,再做中心裁剪或随机水平翻转,作为数据增强;--train_batch_size(默认16):每设备(per device)的 batch size;--num_train_epochs(默认100)与--max_train_steps:二者指定其一,--max_train_steps优先;--gradient_accumulation_steps(默认1):累积多少步再执行一次参数更新,等效放大 batch size;--gradient_checkpointing:开启梯度检查点以节省显存;--learning_rate(默认1e-4);--scale_lr:按 GPU 数、累积步数与 batch size 缩放学习率;--lr_scheduler:可选["linear", "cosine", "cosine_with_restarts", "polynomial", "constant", "constant_with_warmup"],默认"constant";--lr_warmup_steps默认500;--snr_gamma:Min-SNR 损失加权的 γ 值,推荐5.0(详见下文);--use_ema/--offload_ema/--foreach_ema:EMA 权重跟踪及显存优化选项;--use_8bit_adam:使用 bitsandbytes 的 8-bit Adam 优化器;--allow_tf32:在 Ampere 架构 GPU 上允许 TF32 加速矩阵运算;--adam_beta1/--adam_beta2/--adam_weight_decay/--adam_epsilon:Adam 优化器超参;--max_grad_norm(默认1.0):梯度裁剪范数;--prediction_type:"epsilon"或"v_prediction",缺省时沿用 scheduler 配置;--report_to:日志后端,默认"tensorboard",可选"wandb"、"comet_ml"或"all";--checkpoints_total_limit:最多保留的 checkpoint 数量,超出自动删除最旧的;--resume_from_checkpoint:断点续训,传 checkpoint 路径或"latest"自动选择最新;--enable_xformers_memory_efficient_attention:启用 xFormers 内存高效注意力;--noise_offset(默认0,推荐0.1):为低噪声时间步增加偏移噪声,改善暗部细节;--input_perturbation(默认0,推荐0.1):输入扰动,提升小数据量下的收敛质量;--validation_prompts与--validation_epochs(默认5):训练过程中周期性生成验证图像并记录日志;--dream_training与--dream_detail_preservation(默认1.0):启用 DREAM 训练策略;--image_interpolation_mode(默认"lanczos"):resize 时的插值方式;--seed:随机种子,保证训练可复现;--max_train_samples:调试时截断训练样本数;--dataloader_num_workers(默认0):数据加载子进程数;--hub_model_id/--hub_token:推送 Hub 时的仓库名与令牌;--tracker_project_name(默认"text2image-fine-tune"):日志跟踪项目名。
参数校验逻辑位于 parse_args 末尾:--dataset_name与--train_data_dir二者必须提供其一,否则直接抛出ValueError;--non_ema_revision缺省时复用--revision。
Min-SNR 加权:加速收敛的损失重平衡
Min-SNR(Minimum Signal-to-Noise Ratio)加权策略通过重新平衡损失,帮助模型更快收敛。其思想是:扩散模型在不同时间步的信噪比差异巨大,直接对所有时间步等权求 MSE 会主导低信噪比(高噪声)区域的损失;Min-SNR 通过min(SNR, γ)截断高 SNR 步的权重,抑制其对梯度的支配。
- 该脚本既支持预测
epsilon(噪声),也支持v_prediction,而 Min-SNR 对两种预测类型都兼容; - 该加权策略仅由 PyTorch 支持(对应 PyTorch 版训练脚本)。
启用方式(推荐值 5.0):
accelerate launch train_text_to_image.py \ --snr_gamma=5.0从源码看,snr_gamma传入后在训练循环中生效(train_text_to_image.py 约 L1022-L1039):当--snr_gamma为空时使用普通F.mse_loss(..., reduction="mean");否则调用 src/diffusers/training_utils.py 中导出的compute_snr()(定义于 L81-L114),基于 scheduler 的alphas_cumprod计算每个时间步的 SNR = (α/σ)²,再取min(snr, γ)作为逐样本损失权重;当预测类型为epsilon时权重再除以snr,为v_prediction时除以snr + 1,随后对逐元素 MSE 加权取平均。
值得注意的实践结论:
- 对较小数据集,Min-SNR 的效果可能不如大数据集明显;
- 社区在 Weights and Biases 上有不同
snr_gamma取值(如 1.0 与 5.0)的损失面对比实验,可据此观察收敛差异; - 该策略在
epsilon与v_prediction两种预测目标下均有配套的数学形式(详见论文 Section 3.4 与 Section 4.2 的讨论)。
训练脚本内部结构:main() 全流程拆解
数据集预处理代码与训练循环集中在main()函数中(train_text_to_image.py 约 L528-L1165)。若需改造训练脚本,这里是主要修改点。整体流程可拆解为以下阶段:
1. 加载 Scheduler 与 Tokenizer
训练脚本首先加载噪声调度器(noise scheduler)与分词器(tokenizer),见 L592-L595。你可以在此处替换为其他 scheduler:
noise_scheduler = DDPMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") tokenizer = CLIPTokenizer.from_pretrained( args.pretrained_model_name_or_path, subfolder="tokenizer", revision=args.revision )2. 加载 VAE、文本编码器与 UNet
脚本在deepspeed_zero_init_disabled_context_manager()的上下文管理器中加载CLIPTextModel与AutoencoderKL(L616-L622),再加载UNet2DConditionModel(L624-L626)。随后:
vae与text_encoder冻结(requires_grad_(False)),只有unet.train()参与训练(L629-L631);- 若开启
--use_ema,为 UNet 参数创建EMAModel(L634-L643),训练结束时最终保存的权重使用 EMA 版本; - 若开启
--enable_xformers_memory_efficient_attention,校验 xformers 版本(0.0.16 在部分 GPU 上无法用于训练,建议 ≥0.0.17)并启用(L645-L656)。
文档中给出的 UNet 加载写法(用于 checkpoint 恢复场景):
load_model = UNet2DConditionModel.from_pretrained(input_dir, subfolder="unet") model.register_to_config(**load_model.config) model.load_state_dict(load_model.state_dict())这与脚本中load_model_hook的实现一致(L684-L693):从 checkpoint 目录加载unet子文件夹,先同步 config 再加载权重。因为checkpoint 只保存 UNet,从 checkpoint 恢复推理时只需单独加载 UNet(详见下文推理小节)。
3. 数据预处理:tokenize_captions 与 train_transforms
接下来对数据集的文本列与图像列进行预处理:
tokenize_captions:对每个 caption 做分词,max_length=tokenizer.model_max_length,padding="max_length"、truncation=True;当某样本含多个 caption(列表形式)时,训练时随机取一个,验证时取第一个;train_transforms:组合Resize(resolution)、中心/随机裁剪、随机水平翻转、ToTensor()与Normalize([0.5], [0.5]),插值方式由--image_interpolation_mode控制;preprocess_train(L817-L821)将二者打包为数据集变换:
def preprocess_train(examples): images = [image.convert("RGB") for image in examples[image_column]] examples["pixel_values"] = [train_transforms(image) for image in images] examples["input_ids"] = tokenize_captions(examples) return examplescollate_fn(L829-L833)将样本堆叠为连续的pixel_values与input_ids张量,供 DataLoader 消费。
4. 训练循环:latent 编码 → 加噪 → 条件嵌入 → 更新参数
训练循环是脚本的核心(约 L963-L1116),单步迭代逻辑如下:
- 编码到 latent 空间:
vae.encode(...).latent_dist.sample()得到潜在表示,并乘上vae.config.scaling_factor; - 采样噪声并加噪:
torch.randn_like(latents)采样噪声,若开启--noise_offset则叠加偏移噪声;再为每个样本随机采样时间步timesteps,通过noise_scheduler.add_noise(latents, noise, timesteps)完成前向扩散加噪;若开启--input_perturbation则先对噪声做扰动(L967-L990); - 计算文本嵌入:
text_encoder(batch["input_ids"])得到条件向量encoder_hidden_states; - 确定损失目标:根据
prediction_type选择target = noise(epsilon)或target = noise_scheduler.get_velocity(latents, noise, timesteps)(v_prediction);若开启--dream_training,则调用 src/diffusers/training_utils.py 中的compute_dream_and_update_latents执行 DREAM 策略(L996-L1017); - 前向与损失:
unet(noisy_latents, timesteps, encoder_hidden_states)预测噪声残差,按上文是否启用 Min-SNR 计算 MSE 损失; - 反向传播与更新:
accelerator.backward(loss)→ 梯度裁剪accelerator.clip_grad_norm_(..., args.max_grad_norm)→optimizer.step()→lr_scheduler.step()→ 清零梯度(L1045-L1051); - EMA 更新与日志:每个同步梯度步(
accelerator.sync_gradients)后更新 EMA 权重、推进进度条、记录train_loss;每--checkpointing_steps步保存checkpoint-{global_step}(L1054-L1090)。
如果希望深入理解「管线、模型与调度器」在去噪过程中的基本模式,可以参考官方教程 理解 Pipelines、Models 与 Schedulers。
启动训练:以 Naruto 数据集为例
完成参数调整或确认默认配置后即可启动训练。下面以 lambdalabs/naruto-blip-captions(Naruto 角色 + BLIP 自动生成描述)数据集为例,训练一个能生成火影角色的模型。
首先设置环境变量MODEL_NAME与dataset_name,分别指向预训练模型与数据集(Hub 名称或本地路径):
export MODEL_NAME="stable-diffusion-v1-5/stable-diffusion-v1-5" export dataset_name="lambdalabs/naruto-blip-captions" accelerate launch --mixed_precision="fp16" train_text_to_image.py \ --pretrained_model_name_or_path=$MODEL_NAME \ --dataset_name=$dataset_name \ --use_ema \ --resolution=512 --center_crop --random_flip \ --train_batch_size=1 \ --gradient_accumulation_steps=4 \ --gradient_checkpointing \ --max_train_steps=15000 \ --learning_rate=1e-05 \ --max_grad_norm=1 \ --enable_xformers_memory_efficient_attention \ --lr_scheduler="constant" --lr_warmup_steps=0 \ --output_dir="sd-naruto-model" \ --push_to_hub💡 训练本地数据集时,将
TRAIN_DIR与OUTPUT_DIR环境变量分别指向数据集目录与模型保存目录,并把命令中的--dataset_name替换为--train_data_dir=$TRAIN_DIR。本地目录需符合imagefolder结构:图片文件 +metadata.jsonl(每行一条{"file_name": "...", "text": "..."}),详见 创建训练数据集指南。多 GPU 训练时,在
accelerate launch命令中追加--multi_gpu参数即可。
参数组合速查
| 场景 | 建议组合 |
|---|---|
| 单卡 24GB 显存 | --gradient_checkpointing+--mixed_precision="fp16"+--train_batch_size=1+--gradient_accumulation_steps=4 |
| 更大 batch / 更快训练 | 使用 30GB+ 显存 GPU,适当增大--train_batch_size |
| 进一步省显存 | 追加--enable_xformers_memory_efficient_attention(需安装 xformers) |
| 加速收敛 | 追加--snr_gamma=5.0 |
| 提升生成稳定性 | 追加--use_ema(需额外一份全精度参数内存;显存紧张可用--offload_ema将 EMA 权重放到 CPU 固定内存) |
关于 EMA 的补充(来自 examples/text_to_image/README.md):EMA 通过对模型参数维护指数移动平均来平滑更新噪声、提升性能;--foreach_ema使用更快的 foreach 实现;--offload_ema将 EMA 权重驻留于 CPU 固定内存,每个参数更新步非阻塞地搬回 GPU 更新后再搬回 CPU,在主机-设备带宽充足时可做到几乎零额外开销。
训练完成后:加载模型进行推理
训练结束后,模型会被保存到--output_dir(上述例子为sd-naruto-model)。加载微调后的模型进行推理:
from diffusers import StableDiffusionPipeline import torch pipeline = StableDiffusionPipeline.from_pretrained("path/to/saved_model", dtype=torch.float16, use_safetensors=True).to("cuda") # 也可用 "mps"、"xpu"、"cpu" image = pipeline(prompt="yoda").images[0] image.save("yoda-naruto.png")训练脚本在结束时会通过StableDiffusionPipeline.from_pretrained(...)组装完整管线并save_pretrained(output_dir)(L1125-L1133),因此保存目录可直接作为完整 pipeline 加载。
从 checkpoint 恢复推理
由于训练过程的 checkpoint 只保存 UNet 权重,从checkpoint-<N>恢复推理时需单独加载 UNet 再注入管线:
import torch from diffusers import StableDiffusionPipeline, UNet2DConditionModel model_path = "path_to_saved_model" unet = UNet2DConditionModel.from_pretrained(model_path + "/checkpoint-<N>/unet", dtype=torch.float16) pipe = StableDiffusionPipeline.from_pretrained("<initial model>", unet=unet, dtype=torch.float16) pipe.to("cuda") image = pipe(prompt="yoda").images[0] image.save("yoda-naruto.png")中断续训
若训练因意外中断,可在启动命令中追加--resume_from_checkpoint="latest"(自动选择output_dir中最新的checkpoint-*),或显式指定--resume_from_checkpoint=checkpoint-3000。源码中恢复逻辑见 L928-L953:脚本会扫描output_dir下以checkpoint开头的目录并按步数排序,通过accelerator.load_state恢复优化器、调度器与模型状态,并同步重置global_step与起始 epoch。
进阶方向与延伸阅读
完成基础微调后,可进一步探索:
- LoRA 微调:若训练 LoRA 权重(对应
train_text_to_image_lora.py脚本,本文不做展开),推理时加载 LoRA 权重的方法可参考 使用 PEFT 进行推理(加载 LoRA 权重)。LoRA 只需训练新增的低秩分解矩阵,权重极小、不易灾难性遗忘,可在 T4/V100 等消费级 GPU 上运行,且可使用比全量微调高一个量级的学习率(如1e-4而非1e-5); - 推理控制:关于 guidance scale、prompt weighting 等如何控制生成结果,参见 文生图任务指南;
- DREAM 训练:通过
--dream_training启用,以多一次无梯度的 UNet 前向为代价换取更高的模型保真度,--dream_detail_preservation(默认 1.0)控制细节保留因子; - SDXL 微调:仓库还提供了面向 Stable Diffusion XL 的
train_text_to_image_sdxl.py与train_text_to_image_lora_sdxl.py脚本,详见 examples/text_to_image/README_sdxl.md。
总而言之,train_text_to_image.py是一个结构清晰、参数完备的文生图全量微调参考实现:理解其parse_args的参数面、main()的「加载 → 预处理 → 训练循环 → 保存」链路,以及 Min-SNR、EMA、checkpoint 恢复等机制,你就能把它改造成适合自己数据集与硬件条件的定制化训练方案。
【免费下载链接】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),仅供参考