Segmind SD-Small 和 SD-Tiny 知识蒸馏发布
TL;DR
Segmind 发布了 SD‑Small 和 SD‑Tiny 的训练代码和预训练检查点,这些通过块移除知识蒸馏压缩的扩散模型在保持与基础模型相当的图像保真度的同时,使用了 35%–55% 更少的参数,并实现了最高 100% 的推理速度提升。
知识蒸馏方法论
发布的模型使用了在 On Architectural Compression of Text‑to‑Image Diffusion Models (Shinkook et al.) 中描述的 Block‑Removal Knowledge Distillation 技术进行训练。
- 教师模型:Realistic‑Vision 4.0,一个高质量的 Stable Diffusion 检查点。
- 学生架构:移除层的 UNet 变体,导致 SD‑Small 减少 35% 参数,SD‑Tiny 减少 55% 参数。
- 损失组成:
- 目标图像和生成图像之间潜在表示的标准扩散损失。
- 潜在级损失,使学生生成的潜在表示与教师生成的潜在表示对齐。
- 特征级损失,匹配教师和学生之间每个 UNet 块的输出(这是最关键的组成部分)。
- 训练数据:LAION Art Aesthetic 数据集,过滤后保留图像得分 > 7.5。
- 训练计划:使用 100 万张图像,SD‑Small 训练 100k 步,SD‑Tiny 训练 125k 步。
完整的蒸馏管道可在 [segmind/distill‑sd](https://github.com/segmind/distill-sd) 仓库中获取,预训练检查点托管在 Hugging Face 的 segmind 命名空间下。
使用 🤗 Diffusers 的模型
两个模型可以直接从 🤗 Diffusers 库的 DiffusionPipeline 加载:
from diffusers import DiffusionPipeline
import torch
pipeline = DiffusionPipeline.from_pretrained(
"segmind/small-sd", torch_dtype=torch.float16
)
prompt = "Portrait of a pretty girl"
negative_prompt = (
"(deformed iris, deformed pupils, semi-realistic, cgi, 3d, render, sketch, "
"cartoon, drawing, anime:1.4), text, close up, cropped, out of frame, "
"worst quality, low quality, jpeg artifacts, ugly, duplicate, morbid, "
"mutilated, extra fingers, mutated hands, poorly drawn hands, poorly drawn "
"face, mutation, deformed, blurry, dehydrated, bad anatomy, bad proportions, "
"extra limbs, cloned face, disfigured, gross proportions, malformed limbs, "
"missing arms, missing legs, extra arms, extra legs, fused fingers, too many "
"fingers, long neck
)
image = pipeline(prompt, negative_prompt=negative_prompt).images[0]
image.save("my_image.png
同样的 API 适用于 segmind/tiny-sd,只需交换模型标识符。
推理速度提升
在相同硬件上的基准测试显示,蒸馏模型相比原始基础检查点延迟降低最高可达 2×。用于这些测量的推理脚本已包含在仓库中(inference.py)。
已知限制
- 这些模型是 早期阶段的发布;视觉质量可能尚未达到生产级扩散模型的水平。
- 它们 未针对通用生成进行优化,在处理复杂的组合提示或多个概念时可能会遇到困难。
- 最佳实践是在特定领域的数据上进行 微调或应用 LoRA,以获得更高的保真度。
在肖像数据集上微调 SD‑Tiny
Segmind 在使用 Realistic‑Vision 4.0 生成的 7k 图像肖像数据集上演示了 SD‑Tiny 的微调。训练超参数:
- 步数:131 000
- 学习率:1e‑4
- 批量大小:32(梯度累积 = 4)
- 图像分辨率:768 px
- 混合精度:fp16
得到的样本在质量上接近原始教师模型,同时保持了 55% 的参数减少。
在蒸馏模型上的 LoRA 训练
对 SD‑Tiny 应用低秩适配(LoRA)由于模型尺寸减小,可获得 更快的 LoRA 收敛。在抽象概念上训练的示例 LoRA 检查点已在仓库中提供(lora_training.py)。
社区邀请
Segmind 鼓励开发者为项目贡献代码、报告问题并分享微调后的检查点。沟通渠道包括 Discord 服务器和 GitHub 仓库,欢迎星标和拉取请求。