Files

234 lines
8.6 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import math
import torch
from torch import nn, optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from torch.nn import functional as F
import numpy as np
import matplotlib.pyplot as plt
class SinusoidalEmbedding(nn.Module):
"""DDPM 标准的正弦位置编码,把时间步 t 编码成向量"""
def __init__(self, dim):
super().__init__()
self.dim = dim
def forward(self, t):
half_dim = self.dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, device=t.device) * -emb)
emb = t[:, None] * emb[None, :]
return torch.cat([emb.sin(), emb.cos()], dim=-1)
class ResidualBlock(nn.Module):
def __init__(self, in_ch, out_ch, stride=1, time_dim=None):
super().__init__()
# 第一个卷积:可能改变通道数和步长
self.conv1 = nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=1,
padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(out_ch)
self.relu = nn.ReLU(inplace=True)
self.conv2 = nn.Conv2d(out_ch, out_ch, kernel_size=3, stride=1,
padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(out_ch)
# 时间嵌入投影:把 t 的嵌入映射到 out_ch,加到特征图上
self.time_proj = None
if time_dim is not None:
self.time_proj = nn.Sequential(
nn.SiLU(),
nn.Linear(time_dim, out_ch),
)
# shortcut:如果维度变化,用1x1卷积(也可加BN)
self.shortcut = nn.Sequential()
if in_ch != out_ch :
self.shortcut = nn.Sequential(
nn.Conv2d(in_ch, out_ch, kernel_size=1, stride=1, bias=False),
nn.BatchNorm2d(out_ch)
)
def forward(self, x, t_emb=None):
residual = x
out = self.conv1(x)
out = self.bn1(out)
out = self.relu(out)
out = self.conv2(out)
out = self.bn2(out)
# 添加时间嵌入(FiLM 式相加注入)
if t_emb is not None and self.time_proj is not None:
out = out + self.time_proj(t_emb)[:, :, None, None]
# 添加 shortcut
out += self.shortcut(residual)
out = self.relu(out) # 标准做法:最后加一次激活
return out
class DownSample(nn.Module):
def __init__(self,ch):
super().__init__()
self.conv1=nn.Conv2d(ch,ch,2,2)
def forward(self,x):
return self.conv1(x)
class UpSample(nn.Module):
def __init__(self, ch):
super().__init__()
# 使用 ch//2 确保整数
self.conv1 = nn.ConvTranspose2d(ch, ch//2, 4, 2, 1)
def forward(self, x):
return self.conv1(x)
class UNet(nn.Module):
def __init__(self, time_dim=64):
super().__init__()
self.time_emb = SinusoidalEmbedding(time_dim)
self.time_mlp = nn.Sequential(
nn.Linear(time_dim, time_dim * 4),
nn.SiLU(),
nn.Linear(time_dim * 4, time_dim * 4),
)
t_dim = time_dim * 4
self.conv1 = nn.Conv2d(1, 32, 3, 1, 1)
self.res1 = ResidualBlock(32, 64, time_dim=t_dim)
self.down1 = DownSample(64)
self.res2 = ResidualBlock(64, 128, time_dim=t_dim)
self.down2 = DownSample(128)
self.res3 = ResidualBlock(128, 256, time_dim=t_dim)
self.down3 = DownSample(256)
self.res4 = ResidualBlock(256, 512, time_dim=t_dim)
# --- 上采样部分(重新设计通道匹配)---
self.up1 = UpSample(512) # 512 -> 256
self.res5 = ResidualBlock(512, 256, time_dim=t_dim) # 拼接后 512 -> 256
self.up2 = UpSample(256) # 256 -> 128
self.res6 = ResidualBlock(256, 128, time_dim=t_dim) # 拼接后 256 -> 128
self.up3 = UpSample(128) # 128 -> 64
self.res7 = ResidualBlock(128, 64, time_dim=t_dim) # 拼接后 128 -> 64
self.res8 = ResidualBlock(64, 32, time_dim=t_dim)
self.conv2 = nn.Conv2d(32, 1, 3, 1, 1)
def forward(self, x, t):
t_emb = self.time_mlp(self.time_emb(t)) # [b, time_dim*4]
# ----- 下采样(保存跳跃连接)-----
x1 = self.res1(self.conv1(x), t_emb) # [b, 64, 28, 28]
x2 = self.res2(self.down1(x1), t_emb) # [b, 128, 14, 14]
x3 = self.res3(self.down2(x2), t_emb) # [b, 256, 7, 7]
x4 = self.res4(self.down3(x3), t_emb) # [b, 512, 3, 3]
# ----- 上采样 13×3 → 7×7-----
x4_up = self.up1(x4) # [b, 256, 6, 6]
x4_up = F.interpolate(x4_up, size=7, mode='bilinear') # [b, 256, 7, 7]
x4_cat = torch.cat([x4_up, x3], dim=1) # [b, 512, 7, 7]
x3_new = self.res5(x4_cat, t_emb) # [b, 256, 7, 7]
# ----- 上采样 27×7 → 14×14-----
x3_up = self.up2(x3_new) # [b, 128, 14, 14] 尺寸恰好为14
x3_cat = torch.cat([x3_up, x2], dim=1) # [b, 256, 14, 14]
x2_new = self.res6(x3_cat, t_emb) # [b, 128, 14, 14]
# ----- 上采样 314×14 → 28×28-----
x2_up = self.up3(x2_new) # [b, 64, 28, 28]
x2_cat = torch.cat([x2_up, x1], dim=1) # [b, 128, 28, 28]
x1_new = self.res7(x2_cat, t_emb) # [b, 64, 28, 28]
# ----- 最终输出-----
x1_new = self.res8(x1_new, t_emb) # [b, 32, 28, 28]
out = self.conv2(x1_new) # [b, 1, 28, 28]
return out
if __name__=='__main__':
# ---------- 数据加载 ----------
batch_size = 16
transform = transforms.Compose(
[
transforms.ToTensor(), # 将 [0,255] 转为 [0,1] 的 Tensor
transforms.Normalize((0.5,), (0.5,)), #归一化到[-1,1]
]
)
train_dataset = datasets.MNIST(
root="./data", train=True, download=True, transform=transform
)
test_dataset = datasets.MNIST(
root="./data", train=False, download=True, transform=transform
)
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)
# ---------- DDPM 超参数 ----------
timesteps = 600
betas = torch.linspace(1e-4, 0.02, timesteps) # 线性采样 beta
alphas = 1.0 - betas
alpha_bar = torch.cumprod(alphas, dim=0) # \bar{alpha}_t
def q_sample(x_0, t):
"""前向扩散:q(x_t | x_0) = sqrt(alpha_bar_t) * x_0 + sqrt(1 - alpha_bar_t) * eps"""
alpha_bar_t = alpha_bar[t][:, None, None, None]
eps = torch.randn_like(x_0)
x_t = torch.sqrt(alpha_bar_t) * x_0 + torch.sqrt(1 - alpha_bar_t) * eps
return x_t, eps
def sample(model, n=16, device="cpu"):
"""简易采样:从纯噪声逐步去噪"""
model.eval()
x = torch.randn(n, 1, 28, 28, device=device)
with torch.no_grad():
for t in reversed(range(timesteps)):
t_batch = torch.full((n,), t, device=device, dtype=torch.long)
eps_pred = model(x, t_batch)
alpha_t = alphas[t].to(device)
alpha_bar_t = alpha_bar[t].to(device)
x = (x - (1 - alpha_t) / torch.sqrt(1 - alpha_bar_t) * eps_pred) / torch.sqrt(alpha_t)
if t > 0:
x += torch.sqrt(betas[t].to(device)) * torch.randn_like(x)
model.train()
return x
# ---------- 训练 ----------
device = "cuda" if torch.cuda.is_available() else "cpu"
model = UNet().to(device)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
epochs = 20
for epoch in range(epochs):
total_loss = 0.0
for x_0, _ in train_loader:
x_0 = x_0.to(device)
t = torch.randint(0, timesteps, (x_0.size(0),), device=device) # 随机采样时间
x_t, eps = q_sample(x_0, t)
eps_pred = model(x_t, t)
loss = F.mse_loss(eps_pred, eps) # 预测噪声,MSE 损失
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
print(f"Epoch {epoch+1}/{epochs}, Loss: {total_loss/len(train_loader):.4f}")
# ---------- 采样可视化 ----------
samples = sample(model, n=16, device=device).cpu().clamp(-1, 1)
samples = (samples + 1) / 2 # [-1,1] -> [0,1]
fig, axes = plt.subplots(4, 4, figsize=(8, 8))
for i, ax in enumerate(axes.flat):
ax.imshow(samples[i, 0], cmap="gray")
ax.axis("off")
plt.tight_layout()
plt.savefig("samples.png")
plt.show()