1. 项目概述
在计算机视觉领域,生成模型一直是研究热点之一。最近,我基于PyTorch框架实现了一个简化版的Stable Diffusion模型,并在MNIST手写数字数据集上进行了训练和测试。这个项目结合了卷积神经网络(CNN)和注意力机制(Attention),能够根据文本描述生成对应的数字图像。
整个项目包含几个关键组件:变分自编码器(VAE)用于图像的特征提取和重建、时间步调度器控制扩散过程、文本编码器将文字描述转换为向量表示,以及核心的U-Net结构负责噪声预测。经过40多轮训练后,模型在MNIST测试集上表现良好,能够生成清晰可辨的数字图像。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 模型架构解析
2.1 核心组件设计
2.1.1 文本编码器
文本编码器采用预训练的CLIP模型,将文本描述转换为512维的向量表示。这个组件在项目中保持冻结状态,不参与训练:
python复制class TextEncoder(nn.Module):
def __init__(self, path: str = None):
super().__init__()
path = path or "openai/clip-vit-base-patch32"
self.tokenizer = CLIPTokenizer.from_pretrained(path)
self.encoder = CLIPTextModel.from_pretrained(path).eval()
def forward(self, texts: list[str]) -> Tensor:
with torch.no_grad():
inputs = self.tokenizer(texts, return_tensors="pt")
outputs = self.encoder(**inputs)
return outputs.last_hidden_state
注意:使用预训练模型时要注意输入文本的预处理方式,确保与原始训练时的格式一致。
2.1.2 时间步嵌入
时间步嵌入将离散的时间步转换为连续的向量表示,采用了类似Transformer的位置编码方式:
python复制class TimestepEmbedding(nn.Module):
def __init__(self, max_step: int = 1000, d_model: int = 512):
super().__init__()
pe = torch.zeros((max_step, d_model))
position = torch.arange(0, max_step).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2) * -(math.log(10000)/d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
self.register_buffer("pe", pe)
self.mlp = nn.Sequential(
nn.Linear(d_model, 4*d_model),
nn.SiLU(),
nn.Linear(4*d_model, d_model),
)
这种设计使得模型能够感知不同时间步的差异,同时保持了连续空间的平滑性。
2.2 U-Net结构设计
2.2.1 残差块与注意力机制
模型的核心是改进的U-Net结构,包含下采样、中间层和上采样三个部分。每个层级都使用了残差连接和注意力机制:
python复制class CrossAttnResNetBlock(nn.Module):
def __init__(self, c1: int, c2: int, d1: int = 512, d2: int = 512):
super().__init__()
self.res = ResNetBlock(c1, c2, d1=d1)
self.norm = nn.GroupNorm(32, c2)
self.attn = CrossAttention(c2, d2)
def forward(self, x: Tensor, e: Tensor, text: Tensor) -> Tensor:
y = self.res(x, e)
residual = y
batch, c, h, w = y.shape
y = self.norm(y)
y = y.view(batch, c, -1).transpose(1, 2)
y = self.attn(y, text)
y = y.transpose(1, 2).view(batch, c, h, w)
return y + residual
这种设计结合了CNN在图像处理上的优势与注意力机制在长距离依赖建模上的优势。
2.2.2 完整的U-Net结构
完整的U-Net包含编码器、中间层和解码器:
python复制class UNet(nn.Module):
def __init__(self, max_step=1000, time_dim=512, text_dim=512):
super().__init__()
self.scheduler = TimeStepScheduler(max_step)
self.embed = TimestepEmbedding(max_step, time_dim)
self.conv_in = nn.Conv2d(4, 64, kernel_size=3, padding=1)
param = {"d1": time_dim, "d2": text_dim}
self.down1 = Down(64, 128, **param)
self.down2 = Down(128, 256, **param)
self.mid = Mid(256, **param)
self.up1 = Up(256, 128, **param)
self.up2 = Up(128, 64, **param)
self.conv_out = nn.Conv2d(64, 4, kernel_size=3, padding=1)
3. 训练流程与技巧
3.1 数据准备与预处理
MNIST数据集包含60,000张28x28的手写数字图像。预处理流程包括:
- 调整大小为28x28像素
- 转换为张量并归一化到[-1,1]范围
- 为每张图像生成对应的文本标签("0"到"9")
python复制tf = T.Compose([
T.Resize((28, 28)),
T.ToTensor(),
T.Normalize((0.1307,), (0.3081,))
])
提示:保持训练和验证集的预处理方式一致非常重要,不一致的预处理会导致模型性能下降。
3.2 训练配置
训练采用以下关键配置:
- 优化器:AdamW,初始学习率5e-4
- 学习率调度:余弦退火
- 批量大小:训练集50,验证集100
- 梯度累积步数:2
- 训练轮数:50
- 时间步数:1000
python复制self.optimizer = AdamW(model.parameters(), lr=5e-4)
self.scheduler = CosineAnnealingLR(
optimizer, T_max=50, eta_min=5e-5)
3.3 训练技巧
- 梯度裁剪:防止梯度爆炸
python复制torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0)
- 混合精度训练:减少显存占用,加快训练速度
python复制scaler = torch.cuda.amp.GradScaler()
with torch.cuda.amp.autocast():
loss = model(x, text)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
- 模型保存策略:只保存验证损失下降的模型
python复制if val_loss < best_loss:
best_loss = val_loss
torch.save(model.state_dict(), "best_model.pt")
4. 结果分析与优化
4.1 训练结果
经过40轮训练后,模型在测试集上达到了以下指标:
- 训练损失:0.0483
- 验证损失:0.0486
生成的数字图像清晰可辨,能够准确反映输入的文本描述。下图展示了模型对数字0-9的生成结果:

4.2 常见问题与解决方案
-
生成图像模糊
- 原因:VAE重建能力不足或扩散步数过多
- 解决:调整VAE的结构或减少扩散步数
-
模式崩溃(生成多样性不足)
- 原因:模型过于关注少数模式
- 解决:增加模型容量或调整损失函数权重
-
训练不稳定
- 原因:学习率过高或梯度爆炸
- 解决:降低学习率,增加梯度裁剪阈值
4.3 性能优化建议
-
架构优化:
- 尝试不同的U-Net深度和宽度
- 调整注意力头的数量
- 实验不同的激活函数
-
训练策略优化:
- 使用渐进式训练,先训练低分辨率,再提高分辨率
- 尝试不同的噪声调度策略
- 引入分类器引导
-
推理优化:
- 使用DDIM采样加速推理
- 尝试不同的CFG(Classifier-Free Guidance)尺度
- 实现批量化生成
5. 扩展与应用
5.1 扩展到其他数据集
虽然本项目使用MNIST数据集,但模型架构可以轻松扩展到其他图像生成任务:
- Fashion-MNIST:替换数据集,调整输入尺寸
- CIFAR-10:修改输入通道数为3
- 自定义数据集:准备图像-文本对,调整预处理流程
5.2 实际应用场景
- 教育领域:自动生成数学题目配图
- 设计领域:快速生成数字艺术素材
- 数据增强:为分类任务生成更多训练样本
5.3 未来改进方向
- 引入潜在扩散模型(LDM)降低计算成本
- 实现更高分辨率的图像生成
- 支持多模态输入(文本+草图)
- 优化推理速度,实现实时生成
在实际使用这个模型时,我发现几个关键点值得注意:首先,文本编码的质量对生成结果影响很大,需要确保文本描述与训练时的格式一致;其次,扩散步数的选择需要在生成质量和速度之间取得平衡;最后,适当的后处理(如锐化)可以显著改善生成图像的视觉效果。
