137 lines
4.4 KiB
Python
137 lines
4.4 KiB
Python
import torch
|
||||
|
|
from torch import nn, optim
|
|||
|
|
from torch.utils.data import DataLoader
|
|||
|
|
from torchvision import datasets, transforms
|
|||
|
|
import numpy as np
|
|||
|
|
import matplotlib.pyplot as plt
|
|||
|
|
|
|||
|
|
|
|||
|
|
# ---------- 模型定义(修正) ----------
|
|||
|
|
class VAE(nn.Module):
|
|||
|
|
def __init__(self):
|
|||
|
|
super(VAE, self).__init__() # 修正:去掉多余的 self
|
|||
|
|
self.encoder = nn.Sequential(
|
|||
|
|
nn.Linear(784, 256),
|
|||
|
|
nn.ReLU(),
|
|||
|
|
nn.Linear(256, 64),
|
|||
|
|
nn.ReLU(),
|
|||
|
|
nn.Linear(64, 20), # 输出 mu 和 log_var(或 sigma)
|
|||
|
|
)
|
|||
|
|
self.decoder = nn.Sequential(
|
|||
|
|
nn.Linear(10, 64),
|
|||
|
|
nn.ReLU(),
|
|||
|
|
nn.Linear(64, 256),
|
|||
|
|
nn.ReLU(),
|
|||
|
|
nn.Linear(256, 784),
|
|||
|
|
nn.Sigmoid(),
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
def forward(self, x):
|
|||
|
|
hidden = self.encoder(x)
|
|||
|
|
mu, log_var = hidden.chunk(2, dim=1) # 用 log_var 更稳定
|
|||
|
|
# 重参数化:sigma = exp(0.5 * log_var)
|
|||
|
|
std = torch.exp(0.5 * log_var)
|
|||
|
|
eps = torch.randn_like(std)
|
|||
|
|
z = mu + eps * std
|
|||
|
|
x_hat = self.decoder(z)
|
|||
|
|
|
|||
|
|
# KL 散度(按 batch 和像素平均)
|
|||
|
|
KL = -0.5 * torch.sum(1 + log_var - mu.pow(2) - log_var.exp())
|
|||
|
|
KL = KL / (x.size(0) * 28 * 28) # 与重构损失尺度一致
|
|||
|
|
|
|||
|
|
return x_hat, KL
|
|||
|
|
|
|||
|
|
|
|||
|
|
# ---------- 超参数设置 ----------
|
|||
|
|
batch_size = 128
|
|||
|
|
epochs = 20
|
|||
|
|
learning_rate = 1e-3
|
|||
|
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|||
|
|
print(f"Using device: {device}")
|
|||
|
|
|
|||
|
|
# ---------- 数据加载 ----------
|
|||
|
|
transform = transforms.Compose(
|
|||
|
|
[
|
|||
|
|
transforms.ToTensor(), # 将 [0,255] 转为 [0,1] 的 Tensor
|
|||
|
|
transforms.Lambda(lambda x: x.view(-1)), # 展平为 784 维向量
|
|||
|
|
]
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
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)
|
|||
|
|
|
|||
|
|
# ---------- 初始化模型、优化器 ----------
|
|||
|
|
model = VAE().to(device)
|
|||
|
|
optimizer = optim.Adam(model.parameters(), lr=learning_rate)
|
|||
|
|
|
|||
|
|
# 重构损失(二元交叉熵,因为像素在 [0,1])
|
|||
|
|
criterion = nn.BCELoss(reduction="sum") # sum 后与 KL 求和,再除以像素数
|
|||
|
|
|
|||
|
|
# ---------- 训练循环 ----------
|
|||
|
|
for epoch in range(1, epochs + 1):
|
|||
|
|
model.train()
|
|||
|
|
total_loss = 0
|
|||
|
|
total_recon = 0
|
|||
|
|
total_kl = 0
|
|||
|
|
|
|||
|
|
for batch_idx, (data, _) in enumerate(train_loader):
|
|||
|
|
data = data.to(device)
|
|||
|
|
optimizer.zero_grad()
|
|||
|
|
|
|||
|
|
x_hat, kl = model(data)
|
|||
|
|
recon_loss = criterion(x_hat, data) # 按 batch 求和(每个像素的 BCE 之和)
|
|||
|
|
|
|||
|
|
# 总损失 = 重构损失 + KL 散度(都已除以像素数,但 recon 未除,所以需统一)
|
|||
|
|
# 此处将 recon 也除以像素数,使两项量级匹配
|
|||
|
|
recon_loss = recon_loss / (data.size(0) * 28 * 28)
|
|||
|
|
loss = recon_loss + kl
|
|||
|
|
|
|||
|
|
loss.backward()
|
|||
|
|
optimizer.step()
|
|||
|
|
|
|||
|
|
total_loss += loss.item()
|
|||
|
|
total_recon += recon_loss.item()
|
|||
|
|
total_kl += kl.item()
|
|||
|
|
|
|||
|
|
avg_loss = total_loss / len(train_loader)
|
|||
|
|
avg_recon = total_recon / len(train_loader)
|
|||
|
|
avg_kl = total_kl / len(train_loader)
|
|||
|
|
|
|||
|
|
print(
|
|||
|
|
f"Epoch {epoch:2d} | Avg Loss: {avg_loss:.4f} | Recon: {avg_recon:.4f} | KL: {avg_kl:.4f}"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
# 每 5 个 epoch 生成一些样本看看效果(可选)
|
|||
|
|
if epoch % 5 == 0:
|
|||
|
|
model.eval()
|
|||
|
|
with torch.no_grad():
|
|||
|
|
# 从标准正态分布采样 16 个 latent code
|
|||
|
|
sample_z = torch.randn(16, 10).to(device)
|
|||
|
|
generated = model.decoder(sample_z).cpu().numpy()
|
|||
|
|
# 显示
|
|||
|
|
fig, axes = plt.subplots(4, 4, figsize=(6, 6))
|
|||
|
|
for i, ax in enumerate(axes.flat):
|
|||
|
|
ax.imshow(generated[i].reshape(28, 28), cmap="gray")
|
|||
|
|
ax.axis("off")
|
|||
|
|
plt.suptitle(f"Epoch {epoch} Generated Samples")
|
|||
|
|
plt.show()
|
|||
|
|
plt.close()
|
|||
|
|
|
|||
|
|
# ---------- 测试集评估(可选) ----------
|
|||
|
|
model.eval()
|
|||
|
|
test_loss = 0
|
|||
|
|
with torch.no_grad():
|
|||
|
|
for data, _ in test_loader:
|
|||
|
|
data = data.to(device)
|
|||
|
|
x_hat, kl = model(data)
|
|||
|
|
recon = criterion(x_hat, data) / (data.size(0) * 28 * 28)
|
|||
|
|
test_loss += (recon + kl).item()
|
|||
|
|
print(f"Test Average Loss: {test_loss / len(test_loader):.4f}")
|