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}")
|