이 글에서는 마지막으로 코드 구현을 해보고자 한다.
import numpy as np
import matplotlib.pyplot as plt
import torch
import torch.nn as nn
import torch.nn.functional as F
from torchvision import transforms, datasets
batch_size = 512
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transforms.ToTensor())
test_dataset = datasets.MNIST(root='./data', train=False, transform=transforms.ToTensor())
train_loader = torch.utils.data.DataLoader(dataset=train_dataset, batch_size=batch_size, shuffle=True, drop_last=True)
test_loader = torch.utils.data.DataLoader(dataset=test_dataset, batch_size=batch_size, shuffle=False, drop_last=True)
class Encoder(nn.Module):
def __init__(self, output_dim=2):
super().__init__()
self.conv2d_1 = nn.Conv2d(1, 32, kernel_size=3, stride=2, padding=1)
self.conv2d_2 = nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1)
self.conv2d_3 = nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1)
self.linear_mean = nn.Linear(2048, output_dim)
self.linear_logvar = nn.Linear(2048, output_dim)
def forward(self, inputs):
# (batch, 1, 28, 28)
x = self.conv2d_1(inputs).relu()
# (batch, 32, 14, 14)
x = self.conv2d_2(x).relu()
# (batch, 64, 7, 7)
x = self.conv2d_3(x).relu()
# (batch, 128, 4, 4)
x = x.reshape(-1, 2048).relu()
# (batch, 2048)
z_mean = self.linear_mean(x)
z_logvar = self.linear_logvar(x) # 분산에 로그
return z_mean, z_logvar
class Decoder(nn.Module):
def __init__(self, input_dim=2):
super().__init__()
self.linear = nn.Linear(input_dim, 2048)
self.convt_2d_1 = nn.ConvTranspose2d(128, 64, kernel_size=3, padding=1, stride=2)
self.convt_2d_2 = nn.ConvTranspose2d(64, 32, kernel_size=4, padding=1, stride=2)
self.convt_2d_3 = nn.ConvTranspose2d(32, 1, kernel_size=4, padding=1, stride=2)
# self.convt_2d_4 = nn.Conv2d(16, 1, kernel_size=1)
def forward(self, inputs):
# (batch, 2)
x = self.linear(inputs).relu()
# (batch, 2048)
x = x.reshape(-1, 128, 4, 4)
# (batch, 128, 4, 4)
x = self.convt_2d_1(x).relu()
# (batch, 64, 7, 7)
x = self.convt_2d_2(x).relu()
# (batch, 32, 14, 14)
x = self.convt_2d_3(x).sigmoid()
return x
class VAE(nn.Module):
def __init__(self, emb_dim=2):
super().__init__()
self.encoder = Encoder(emb_dim)
self.decoder = Decoder(emb_dim)
def forward(self, inputs):
z_mean, z_logvar = self.encoder(inputs)
z = z_mean + (z_logvar/2).exp() * torch.randn_like(z_logvar) * self.training
outputs = self.decoder(z)
return outputs, z_mean, z_logvar
def vae_loss(x_pred, x_true, z_mean, z_logvar):
recon_loss = F.binary_cross_entropy(x_pred, x_true)
kl_loss = 0.5 * torch.mean(z_mean**2 + z_logvar.exp() - z_logvar - 1)
return recon_loss, kl_loss
from torch import optim
import numpy as np
# torch.autograd.set_detect_anomaly(True)
epochs = 500
model = VAE(2).cuda()
# 모델, 손실함수, 옵티마이저 초기화
optimizer = optim.AdamW(model.parameters(), lr=0.002)
scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.5, min_lr=1e-5, patience=20,)
# Early stopping을 위한 변수 초기화
best_val_loss = float('inf')
patience = 50
early_stopping_counter = 0
# history
loss_history = []
recon_loss_history = []
kl_loss_history = []
val_loss_history = []
val_recon_loss_history = []
val_kl_loss_history = []
# 학습 루틴
for epoch in range(epochs):
# 학습
model.train()
loss = 0.0
recon_loss = 0.0
kl_loss = 0.0
for inputs, _ in train_loader:
inputs = inputs.cuda()
optimizer.zero_grad()
outputs, z_mean, z_logvar = model(inputs)
_recon_loss, _kl_loss = vae_loss(outputs, inputs, z_mean, z_logvar)
_loss = _recon_loss + _kl_loss*5e-3
_loss.backward()
optimizer.step()
loss += _loss.item()
recon_loss += _recon_loss.item()
kl_loss += _kl_loss.item()
# 검증
model.eval()
val_loss = 0.0
val_recon_loss = 0.0
val_kl_loss = 0.0
with torch.no_grad():
for inputs, _ in test_loader:
inputs = inputs.cuda()
outputs, z_mean, z_logvar = model(inputs)
_recon_loss, _kl_loss = vae_loss(outputs, inputs, z_mean, z_logvar)
_loss = _recon_loss + _kl_loss*5e-3
val_loss += _loss.item()
val_recon_loss += _recon_loss.item()
val_kl_loss += _kl_loss.item()
# 출력
loss /= len(train_loader)
recon_loss /= len(train_loader)
kl_loss /= len(train_loader)
val_loss /= len(test_loader)
val_recon_loss /= len(test_loader)
val_kl_loss /= len(test_loader)
loss_history.append(loss)
recon_loss_history.append(recon_loss)
kl_loss_history.append(kl_loss)
val_loss_history.append(val_loss)
val_recon_loss_history.append(val_recon_loss)
val_kl_loss_history.append(val_kl_loss)
if (epoch+1)%5 == 0:
print(f"Epoch[{epoch+1:03d}]\tTrain Loss[all,recon,kl]: [{loss:.5f} {recon_loss:.5f} {kl_loss:.5f}]\t" +
f"Valid Loss[all,recon,kl]: [{val_loss:.5f} {val_recon_loss:.5f} {val_kl_loss:.5f}]\t" +
f"lr: {optimizer.param_groups[0]['lr']:.5f}")
scheduler.step(val_loss)
# 검증 손실이 이전보다 크면 early stopping counter를 증가시킴
if val_loss > best_val_loss:
early_stopping_counter += 1
else:
best_val_loss = val_loss
early_stopping_counter = 0
if best_val_loss < 0.35:
torch.save(model.state_dict(), './models/best.pt')
# early stopping 조건 충족 시 학습 중지
if early_stopping_counter >= patience:
print("Early stopping! No improvement in validation loss.")
print(f"Epoch {epoch+1:03d}\tTrain Loss: {loss:.6f}" +
f"\tValid Loss: {val_loss:.6f}\tlr: {optimizer.param_groups[0]['lr']:.6f}")
print(f"Best_val_loss : {best_val_loss:.6f}")
break
# model = VAE(2).cuda()
model.load_state_dict(torch.load('./models/best.pt'))
print("Restored to best model.")
import matplotlib.pyplot as plt
# Loss history 그래프 그리기
plt.plot(loss_history, label='Train Loss')
plt.plot(val_loss_history, label='Validation Loss')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.title('Loss History')
plt.legend()
plt.show()
import matplotlib.pyplot as plt
fig, axs = plt.subplots(nrows=2, figsize=(7,2))
model = model.cpu()
for inputs, label in test_loader:
inputs = inputs[:10]
with torch.no_grad():
model.eval()
outputs, _, _ = model(inputs)
break
inputs = torch.hstack(list(inputs.squeeze())).unsqueeze(-1).tile(1,1,3)
outputs = torch.hstack(list(outputs.squeeze())).unsqueeze(-1).tile(1,1,3)
axs[0].imshow(inputs)
axs[1].imshow(outputs)
import matplotlib.pyplot as plt
import matplotlib.patches as patches
colors = plt.cm.get_cmap('tab10', 10)
fig, ax = plt.subplots(figsize=(5,4))
latent_space = []
labels = []
model = model.cpu()
for inputs, label in test_loader:
with torch.no_grad():
model.eval()
outputs, z_mean, z_logvar = model(inputs)
latent_space.append(z_mean)
labels.append(label)
latent_space = torch.cat(latent_space)
labels = torch.cat(labels)
for i in range(10):
idx = torch.where(labels==i)[0]
ax.scatter(latent_space[idx, 0], latent_space[idx, 1], c=[colors(i)], label=f'Class {i}')
rect = patches.Rectangle((-2, -2), 4, 4, linewidth=1, edgecolor='r', facecolor='none')
ax.add_patch(rect)
ax.legend()
plt.show()

x = torch.arange(-2, 2.1, step=0.4)
y = torch.arange(2.0, -2.1, step=-0.4)
z = torch.stack(torch.meshgrid([y,x]), dim=-1).reshape(-1, 2)
z = z[:, [1,0]]
with torch.no_grad():
model.eval()
outputs = model.decoder(z)
outputs = outputs.squeeze() # (121, 28, 28)
image = torch.zeros([28*11, 28*11], dtype=torch.float32)
for i, img in enumerate(outputs):
image[i//11*28:i//11*28+28, i%11*28:i%11*28+28] = img
image = image.unsqueeze(-1).tile(1,1,3)
fig, axs = plt.subplots(figsize=(4,4))
axs.imshow(image)

latent_space가 0 근방에 모이긴 했지만 생각보다 학습이 제대로 이루어 지진 않은 듯하다.
X