234 lines
8.6 KiB
Python
234 lines
8.6 KiB
Python
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]
|
|||
|
|
|
|||
|
|
# ----- 上采样 1(3×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]
|
|||
|
|
|
|||
|
|
# ----- 上采样 2(7×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]
|
|||
|
|
|
|||
|
|
# ----- 上采样 3(14×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()
|
|||
|
|
|
|||
|
|
|