news 2026/9/10 19:55:26

在 16GB 显存内完成 Stable Diffusion 3 DreamBooth LoRA 训练:Diffusers 官方教育项目源码全解

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
在 16GB 显存内完成 Stable Diffusion 3 DreamBooth LoRA 训练:Diffusers 官方教育项目源码全解

在 16GB 显存内完成 Stable Diffusion 3 DreamBooth LoRA 训练:Diffusers 官方教育项目源码全解

【免费下载链接】diffusers🤗 Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers

Stable Diffusion 3(SD3)凭借三文本编码器架构实现了出色的文本理解能力,但也带来了高达 20GB 以上的显存占用,令消费级 GPU 用户望而却步。本文以 🤗 Diffusers 仓库中examples/research_projects/sd3_lora_colab教育项目(README 见 examples/research_projects/sd3_lora_colab/README.md)为骨架,逐行拆解其"预计算文本嵌入 + 微型 LoRA 训练"两阶段方案:从 8bit T5 量化、FP16 混合精度、8bit Adam 到梯度检查点与 Flash Attention,完整还原其在 16GB 显存内(含免费版 Colab T4)跑通 SD3 DreamBooth LoRA 训练的全过程,并给出可直接复制的命令行与推理代码。

一、为什么 SD3 训练如此吃显存

SD3 与 SD1.x/SDXL 最大的不同在于其文本编码架构:模型同时使用CLIP-L、CLIP-G 与 T5-XXL 三个文本编码器,其中仅 T5-XXL 就有约 46 亿参数。当三个编码器以 FP32 精度同时加载并计算文本嵌入时,显存占用可飙升至约 20GB;即便降到 FP16,也仍需要约 12GB。这一量级显然超出了 16GB 消费级显卡(尤其是 Colab 免费版 T4 的 15GB)的承受范围。

本项目的核心洞察是:文本嵌入在训练中是可以"预先算好、序列化复用"的。DreamBooth 的训练数据往往只有寥寥数张到数十张实例图片,其 prompt 高度固定。既然每张图对应的prompt_embeds在训练全程不变,就没必要在每一步训练时反复加载三个编码器去前向计算,完全可以离线算一次、存成文件、训练时直接读取。这正是该教育项目能够把训练显存压进 16GB 的第一块基石。

二、整体方案:两阶段流水线

从 sd3_dreambooth_lora_16gb.ipynb 的 21 个单元格可以看出,整个工作流被清晰地切分为两个阶段:

  1. 阶段一(compute_embeddings.py:加载 SD3 管道(文本编码器以 8bit 量化),计算实例 prompt 的嵌入并序列化为sample_embeddings.parquet,随后释放显存。
  2. 阶段二(train_dreambooth_lora_sd3_miniature.py:完全不加载任何文本编码器,直接读取 parquet 中的嵌入 + 实例图片,仅训练 Transformer(MMDiT)的 LoRA 层。

这种设计将"最占显存的文本编码"与"真正需要反向传播的训练"在时间上彻底解耦,配合训练脚本内的多项显存优化手段,最终在 16GB 显存内完成训练。

三、阶段一:用 8bit T5 预计算并序列化文本嵌入

3.1 源码拆解:compute_embeddings.py

脚本核心位于 examples/research_projects/sd3_lora_colab/compute_embeddings.py,其关键实现如下:

PROMPT = "a photo of sks dog" MAX_SEQ_LENGTH = 77 LOCAL_DATA_DIR = "dog" OUTPUT_PATH = "sample_embeddings.parquet" def load_sd3_pipeline(): id = "stabilityai/stable-diffusion-3-medium-diffusers" text_encoder = T5EncoderModel.from_pretrained(id, subfolder="text_encoder_3", load_in_8bit=True, device_map="auto") pipeline = StableDiffusion3Pipeline.from_pretrained( id, text_encoder_3=text_encoder, transformer=None, vae=None, device_map="balanced" ) return pipeline

这里有两个关键设计:

  • load_in_8bit=True:T5-XXL 以LLM.int8()论文(见 Hugging Face 论文页 2208.07339)引入的 8bit 量化方式加载。这是显存占用从 20GB(FP32)/ 12GB(FP16)进一步压到约 10.5GB的决定性手段,代价是对 T5 编码质量的影响在可接受范围内。
  • transformer=None, vae=None:计算嵌入阶段根本不需要 MMDiT Transformer 和 VAE,直接置空不加载,进一步省显存。

随后通过pipeline.encode_prompt(prompt=..., prompt_2=None, prompt_3=None, max_sequence_length=...)(该方法定义于 src/diffusers/pipelines/stable_diffusion_3/pipeline_stable_diffusion_3.py)得到四组张量:prompt_embedsnegative_prompt_embedspooled_prompt_embedsnegative_pooled_prompt_embeds,其中prompt_embeds形状为[77, 4096]MAX_SEQ_LENGTH× T5 隐藏维度)。

3.2 以"图片哈希"为键的 parquet 序列化

脚本对dog/目录下每张 JPEG 图片计算SHA-256 哈希generate_image_hash),然后把(image_hash, prompt_embeds, negative_prompt_embeds, pooled_prompt_embeds, negative_pooled_prompt_embeds)组装成 DataFrame,并将嵌入展平为列表后写入sample_embeddings.parquet。哈希建键的意义在于:训练阶段的数据集类需要按图片精确对齐到其对应的嵌入,只要图片文件没有改动,哈希就能稳定对应。

3.3 命令行参数

脚本通过argparse暴露了 4 个参数(含默认值):

参数默认值说明
--prompt"a photo of sks dog"实例 prompt,训练时使用的触发词
--max_sequence_length77计算嵌入的最大序列长度,越长计算成本越高
--local_data_dir"dog"包含实例图片的目录(脚本默认假定图片为.jpeg扩展名,可按需修改)
--output_path"sample_embeddings.parquet"parquet 输出路径

默认直接运行python compute_embeddings.py即可完成 dog 示例的嵌入计算;脚本会在结束时打印各张量形状与torch.cuda.max_memory_allocated()统计的峰值显存,便于核对内存预算。

四、阶段二:微型 DreamBooth LoRA 训练脚本

4.1 只训 LoRA,冻结一切

在 examples/research_projects/sd3_lora_colab/train_dreambooth_lora_sd3_miniature.py 中,训练脚本只加载三样东西:

  • FlowMatchEulerDiscreteScheduler(scheduler,用于 flow matching 加噪);
  • AutoencoderKL(VAE,把像素空间编码到潜空间);
  • SD3Transformer2DModel(MMDiT 主模型)。

随后立即transformer.requires_grad_(False)vae.requires_grad_(False),并给 Transformer 的注意力层挂上 PEFT LoRA:

transformer_lora_config = LoraConfig( r=args.rank, lora_alpha=args.rank, init_lora_weights="gaussian", target_modules=["to_k", "to_q", "to_v", "to_out.0"], ) transformer.add_adapter(transformer_lora_config)

target_modules覆盖了自注意力与交叉注意力的 Q/K/V 及输出投影to_out.0rank默认取 4(--rank参数可调),lora_alpha与 rank 相同。文本编码器完全不参与训练(README 明确说明训练中文本编码器被刻意禁用),因此训练图里只有 LoRA 参数需要梯度。

4.2 数据管道:按哈希从 parquet 取嵌入

DreamBoothDataset是理解本脚本的关键:它读取--instance_data_dir下的图片,逐张计算哈希;同时pd.read_parquet(data_df_path)读入阶段一产出的嵌入表,通过map_image_hash_embedding构建{image_hash: (prompt_embeds, pooled_prompt_embeds)}字典。取数据时:

prompt_embeds = np.array(prompt_embeds).reshape(154, 4096) pooled_prompt_embeds = np.array(pooled_prompt_embeds).reshape(2048)

注意这里 reshape 成154 × 4096——因为encode_prompt会同时返回正、负 prompt 的嵌入并拼接(77×2),训练时直接喂给 Transformer 的encoder_hidden_states。图片侧则统一做 Resize、中心/随机裁剪、可选水平翻转,并归一化到[-1, 1]Normalize([0.5], [0.5]))。collate_fn把像素、嵌入分别torch.stack打包成 batch。

4.3 训练循环:flow matching + SD3 时间步采样

训练循环体现了 SD3 的flow matching目标:随机采样时间步,用sigmas * noise + (1.0 - sigmas) * model_input构造带噪潜变量,预测目标直接是干净潜变量model_input(而非 epsilon),损失为加权 MSE。时间步采样与损失加权复用 diffusers 训练工具库 src/diffusers/training_utils.py 中的两个函数:

  • compute_density_for_timestep_sampling(training_utils.py):支持sigma_sqrt/logit_normal/mode/cosmap四种加权方案,其中logit_normal通过正态分布 + sigmoid 让时间步采样偏向中段,mode方案则用mode_scale(默认 1.29)控制分布形状;
  • compute_loss_weighting_for_sd3(training_utils.py):对sigma_sqrtsigmas^-2加权、对cosmap用余弦映射加权,其余方案权重为 1。

4.4 四大显存优化手段(README 的 "How" 部分)

训练脚本同时启用了以下手段,README 对此有明确交代:

  1. 8bit Adam:通过bitsandbytesbnb.optim.AdamW8bit优化器(--use_8bit_adam开关),把优化器状态从 8 字节/参数压到 2 字节/参数,是显存优化的最大头;
  2. 梯度检查点(Gradient checkpointing)transformer.enable_gradient_checkpointing()以少量重计算换取大幅显存节省;
  3. 梯度累积--gradient_accumulation_steps=4,用时间换显存,等效放大 batch size;
  4. FP16 混合精度--mixed_precision="fp16"下所有非训练权重(VAE、非 LoRA 的 Transformer)cast 到torch.float16,而 LoRA 可训练参数通过cast_training_params(见 training_utils.py)保持 FP32 以保证数值稳定;
  5. Flash Attention:脚本通过F.scaled_dot_product_attention()(PyTorch 原生 SDPA)实现高效注意力,无需额外安装 flash-attn 包。

VAE 保持torch.float32,编码潜变量后与噪声混合,再 cast 到weight_dtype送入 Transformer,兼顾精度与显存。

4.5 断点续训、验证与 Hub 推送

脚本还集成了完整的工程化能力:

  • Checkpoint--checkpointing_steps(默认 500 步)通过accelerator.save_state保存;--resume_from_checkpoint支持路径或"latest"自动续训;--checkpoints_total_limit可限制保留数量自动清理旧 checkpoint;
  • 保存/加载钩子save_model_hook/load_model_hook通过StableDiffusion3Pipeline.save_lora_weights/lora_state_dict以标准 LoRA 格式序列化,加载时用convert_unet_state_dict_to_peft转换键名并经set_peft_model_state_dict注入;
  • 验证--validation_prompt+--validation_epochs(默认 50)周期性调用log_validation,期间启用enable_model_cpu_offload()以省显存,结果写入 TensorBoard/wandb;
  • 推送 Hub--push_to_hub会通过save_model_card生成带 Gallery 组件与触发词的模型卡(license 标注 openrail++),并用upload_folder上传output_dir

五、端到端实操:在 Colab 免费版跑通全流程

以下步骤完整继承自 sd3_dreambooth_lora_16gb.ipynb。

5.1 前置条件:SD3 是 gated 模型

SD3 模型仓库(stabilityai/stable-diffusion-3-medium-diffusers)是受控访问的,必须先在其模型页同意共享联系方式以通过 gating,然后登录本机:

hf auth login

登录凭证同时也用于训练后将模型推送至 Hugging Face Hub。

5.2 安装依赖并克隆仓库

!pip install -q -U git+https://github.com/huggingface/diffusers !pip install -q -U \ transformers \ accelerate \ wandb \ bitsandbytes \ peft !git clone https://github.com/huggingface/diffusers %cd diffusers/examples/research_projects/sd3_lora_colab

5.3 下载实例图片数据集

项目使用 Hugging Face 上的diffusers/dog-example数据集(仅 5 张狗的照片):

from huggingface_hub import snapshot_download local_dir = "./dog" snapshot_download( "diffusers/dog-example", local_dir=local_dir, repo_type="dataset", ignore_patterns=".gitattributes", )

下载后清理.cache目录即可进入嵌入计算。

5.4 计算并序列化嵌入

!python compute_embeddings.py

使用默认实例 prompt"a photo of sks dog",产物为sample_embeddings.parquet。接着手动释放显存:

import torch import gc def flush(): torch.cuda.empty_cache() gc.collect() flush()

5.5 启动训练

训练命令(与 notebook 完全一致):

!accelerate launch train_dreambooth_lora_sd3_miniature.py \ --pretrained_model_name_or_path="stabilityai/stable-diffusion-3-medium-diffusers" \ --instance_data_dir="dog" \ --data_df_path="sample_embeddings.parquet" \ --output_dir="trained-sd3-lora-miniature" \ --mixed_precision="fp16" \ --instance_prompt="a photo of sks dog" \ --resolution=1024 \ --train_batch_size=1 \ --gradient_accumulation_steps=4 --gradient_checkpointing \ --use_8bit_adam \ --learning_rate=1e-4 \ --report_to="wandb" \ --lr_scheduler="constant" \ --lr_warmup_steps=0 \ --max_train_steps=500 \ --seed="0"

各关键参数的语义与默认值如下(均可在脚本parse_args中查到):

参数本示例取值默认值说明
--data_df_pathsample_embeddings.parquet无(必填)阶段一序列化的嵌入文件
--mixed_precisionfp16可选no/fp16/bf16(bf16 需 Ampere+ GPU)
--resolution1024512训练分辨率,SD3 原生 1024
--train_batch_size14单设备 batch size
--gradient_accumulation_steps41累积步数,等效放大 batch
--use_8bit_adam开启关闭bitsandbytes 8bit 优化器
--learning_rate1e-41e-4初始学习率
--lr_schedulerconstantconstant可选linear/cosine/cosine_with_restarts/polynomial/constant/constant_with_warmup
--lr_warmup_steps0500warmup 步数
--max_train_steps500无(按 epoch 推导)总训练步数,提供后覆盖num_train_epochs
--weighting_scheme未指定logit_normalSD3 时间步加权方案
--rank未指定4LoRA 秩
--max_sequence_length未指定77T5 编码最大序列长度,须与阶段一致
--checkpointing_steps未指定500每 N 步存一次 checkpoint
--seed0随机种子,保证可复现

notebook 提示:完整训练约需1 小时(取决于数据集规模)。训练期间可通过 wandb 实时监控 loss 与验证图。

5.6 推理验证

训练完成后加载 LoRA 并进行推理(注意训练结束会自动把 LoRA 权重以 diffusers 格式写入output_dir):

from diffusers import DiffusionPipeline import torch pipeline = DiffusionPipeline.from_pretrained( "stabilityai/stable-diffusion-3-medium-diffusers", torch_dtype=torch.float16 ) pipeline.load_lora_weights("trained-sd3-lora-miniature") pipeline.enable_sequential_cpu_offload() image = pipeline("a photo of sks dog in a bucket").images[0] image.save("bucket_dog.png")

notebook 特别提醒:此处推理会明显偏慢,因为enable_sequential_cpu_offload()会在每个子模块用前加载、用后卸载,引入大量数据搬运开销;这是 16GB 显存下的合理取舍。

六、已知限制(README 的 "Gotchas")

README 明确强调这是一个**教育性质(EDUCATIONAL)**项目,用于证明"在消费级 GPU 上微调大型扩散系统是可能的",但要达到 SOTA 效果还需自行补充组件。用户必须了解的已知限制:

  • 文本编码器训练被刻意禁用:只微调 Transformer 的 LoRA,不训练任何文本编码器(更不涉及文本编码器 LoRA);
  • 不支持 prior-preservation(先验保留):即不支持 class 图片的正则化损失,训练数据只有实例图片,容易出现过拟合/语言漂移;
  • 不支持实例图片的自定义 caption:所有图片共享同一个实例 prompt,README 指出这相对容易扩展——只需为每张图单独算嵌入并按哈希存储,训练侧DreamBoothDataset已天然支持按图取嵌入,改动点集中在数据准备阶段。

七、扩展思路:以本项目为模板继续演进

从源码结构看,本项目刻意把"嵌入计算"与"LoRA 训练"解耦,这为二次开发留下了清晰接口:

  1. 自定义 caption 支持:修改compute_embeddings.py为每张图计算专属嵌入(而非全局一份),parquet 按image_hash索引的存储结构无需改动,训练侧DreamBoothDataset.__getitem__已按哈希取嵌入,可直接复用;
  2. 提升效果:可尝试在LoraConfig中调整rank、切换--weighting_scheme(如cosmap)观察收敛差异,或补上 prior-preservation 所需的 class 数据与损失项;
  3. 复用到其他数据:将--instance_data_dir换成自己的图片目录(注意 JPEG 扩展名假设)、修改--instance_prompt--max_sequence_length后重新跑两个阶段即可。

总而言之,这个教育项目用约 10.5GB(嵌入阶段)与更低(训练阶段)的显存预算,展示了"大模型消费级微调"的完整方法论:把昂贵的文本编码离线化、把可训练参数最小化(LoRA)、把数值精度与优化器状态做精细化取舍。对于想在 16GB 显卡上尝试 SD3 微调的开发者,它是极具参考价值的起点模板。

【免费下载链接】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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 19:55:06

cann/ge HCCL TP图Python样例

样例使用指导 【免费下载链接】ge GE(Graph Engine)是面向昇腾的图编译器和执行器,提供了计算图优化、多流并行、内存复用和模型下沉等技术手段,加速模型执行效率,减少模型内存占用。 GE 提供对 PyTorch、TensorFlow 前…

作者头像 李华
网站建设 2026/9/10 19:54:56

虚拟现实交互设计入门:从原理到项目实战的完整心得

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/10 19:53:23

WSL 容器 C API 端到端实战:用 WslcSDK 驱动容器完整生命周期

WSL 容器 C API 端到端实战:用 WslcSDK 驱动容器完整生命周期 【免费下载链接】WSL Windows Subsystem for Linux 项目地址: https://gitcode.com/GitHub_Trending/ws/WSL WSL 容器(WSLC)在 Windows Subsystem for Linux 项目中提供了…

作者头像 李华
网站建设 2026/9/10 19:53:22

LeetCode 496. Next Greater Element I 题解:Go 单调栈与哈希表实战

LeetCode 496. Next Greater Element I 题解:Go 单调栈与哈希表实战 【免费下载链接】LeetCode-Go ✅ Solutions to LeetCode by Go, 100% test coverage, runtime beats 100% | LeetCode 题解 项目地址: https://gitcode.com/GitHub_Trending/le/LeetCode-Go …

作者头像 李华
网站建设 2026/9/10 19:53:20

SpringBoot+Vue城市公交调度系统设计与实现:从排班到实时监控

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华