VAE 탐구 - 6/6

Tetrapod·2024년 5월 22일

VAE 탐구

목록 보기
6/6
post-thumbnail

이 글에서는 마지막으로 코드 구현을 해보고자 한다.


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

MNIST 가져오기

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
    

Loss 정의

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

Loss History

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)

테스트 이미지 2차원에 나열


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 근방에 모이긴 했지만 생각보다 학습이 제대로 이루어 지진 않은 듯하다.


Reference

X

0개의 댓글