I use VAE to reconstruct images. For the first experiment I use VAE to reconstruct MNIST images and it works properly as expected. But using VAE to reconstruct CIFAR-10 and other RGB images, I get in the output something like "noise". The code below:
Download data
transform = transforms.Compose([transforms.Resize((32, 32)), transforms.ToTensor()])
train_data = datasets.CIFAR10(root='data', train=True, download=True, transform=transform)
test_data = datasets.CIFAR10(root='data', train=False, download=True, transform=transform)
Define Encoder&Decoder
import torch.nn as nn
class Encoder(nn.Module):
def __init__(self, latent_dim):
super().__init__()
self.latent_dim = latent_dim # latent space size
hidden_dims = [32, 64, 128, 256, 512] # num of filters in layers
modules = []
in_channels = 3 # initial value of channels
for h_dim in hidden_dims[:-1]: # conv layers
modules.append(
nn.Sequential(
nn.Conv2d(
in_channels=in_channels, # num of input channels
out_channels=h_dim, # num of output channels
kernel_size=3,
stride=2, # convolution kernel step
padding=1, # save shape
),
nn.BatchNorm2d(h_dim),
nn.LeakyReLU(),
)
)
in_channels = h_dim # changing number of input channels for next iteration
modules.append(
nn.Sequential(
nn.Conv2d(in_channels=256, out_channels=512, kernel_size=1), # changing the kernel size, because size of the array (2*2)
nn.BatchNorm2d(512),
nn.LeakyReLU(),
)
)
modules.append(nn.Flatten()) # to vector, size 512 * 2*2 = 2048
modules.append(nn.Linear(512 * 2 * 2, latent_dim))
self.encoder = nn.Sequential(*modules)
def forward(self, x):
x = self.encoder(x)
return x
class Decoder(nn.Module):
def __init__(self, latent_dim):
super().__init__()
hidden_dims = [512, 256, 128, 64, 32] # num of filters in layers
self.linear = nn.Linear(in_features=latent_dim, out_features=512)
modules = []
for i in range(len(hidden_dims) - 1): # define ConvTransopse layers
modules.append(
nn.Sequential(
nn.ConvTranspose2d(
in_channels=hidden_dims[i],
out_channels=hidden_dims[i + 1],
kernel_size=3,
stride=2,
padding=1,
output_padding=1,
),
nn.BatchNorm2d(hidden_dims[i + 1]),
nn.LeakyReLU(),
)
)
modules.append(
nn.Sequential(
nn.ConvTranspose2d(
in_channels=hidden_dims[-1],
out_channels=hidden_dims[-1],
kernel_size=3,
stride=2,
padding=1,
output_padding=1,
),
nn.BatchNorm2d(hidden_dims[-1]),
nn.LeakyReLU(),
nn.Conv2d(in_channels=hidden_dims[-1], out_channels=3, kernel_size=5, padding=2),
nn.Sigmoid(),
)
)
self.decoder = nn.Sequential(*modules)
def forward(self, x):
x = self.linear(x) # from latents space to Linear
x = x.view(-1, 512, 1, 1) # reshape
x = self.decoder(x) # reconstruction
return x
Define extra functions and reassign a VAE Encoder
torch.manual_seed(42)
class VAEEncoder(Encoder):
def __init__(self, latent_dim):
if latent_dim % 2 != 0: # check for the parity of the latent space
raise Exception("Latent size for VAEEncoder must be even")
super().__init__(latent_dim)
def vae_split(latent):
size = latent.shape[1] // 2 # divide the latent representation into mu and log_var
# mu = latent
# log_var = latent
mu = latent[:, :size]
log_var = latent[:, size:]
return mu, log_var
def vae_reparametrize(mu, log_var):
sigma = torch.exp(0.5 * log_var) #0.5 * log_var ???
z = torch.randn(mu.shape[0], mu.shape[1]).to(device)
return z * sigma + mu
def vae_pass_handler(encoder, decoder, data, *args, **kwargs):
latent = encoder(data)
mu, log_var = vae_split(latent)
sample = vae_reparametrize(mu, log_var)
recon = decoder(sample)
return latent, recon
def kld_loss(mu, log_var):
var = log_var.exp()
# kl_loss = torch.mean(-0.5 * torch.sum(1 + log_var - mu ** 2 - var, dim=1), dim=0)
kl_loss = 0.5 * torch.mean(torch.sum(mu ** 2 + var - log_var - 1., dim=-1))
# The same result
return kl_loss
def vae_loss_handler(data, recons, latent, kld_weight=0.005, *args, **kwargs):
mu, log_var = vae_split(latent)
kl_loss = kld_loss(mu, log_var)
bce_loss = F.binary_cross_entropy(recons, data)
loss = kld_weight * kl_loss + bce_loss
return kl_loss, bce_loss, loss # add bce loss(reconstruction)
Define models, cuda etc
import torch.optim as optim
from itertools import chain
torch.manual_seed(42)
latent_dim = 100
learning_rate = 1e-4
encoder = VAEEncoder(latent_dim=latent_dim*2)
decoder = Decoder(latent_dim=latent_dim)
device = 'cuda' if torch.cuda.is_available else 'cpu'
encoder = encoder.to(device)
decoder = decoder.to(device)
optimizer = optim.Adam(
chain(encoder.parameters(), decoder.parameters()), lr=learning_rate
)
Define train function
from tqdm.notebook import tqdm
def train(
enc,
dec,
loader,
optimizer,
single_pass_handler,
loss_handler,
epoch,
log_interval=500,
):
for batch_idx, (data, lab) in enumerate(tqdm(loader)):
batch_size = data.size(0)
optimizer.zero_grad()
data = data.to(device)
lab = lab.to(device)
latent, output = single_pass_handler(encoder, decoder, data, lab) # reconstructed image drom decoder
kl_loss, bce_loss, loss = loss_handler(data, output, latent) # compute loss
# loss = loss_handler(data, output, latent)
loss.backward()
optimizer.step()
if batch_idx % log_interval == 0:
print(
"Train Epoch: {} [{}/{} ({:.0f}%)]".format(
epoch,
batch_idx * len(data),
len(loader.dataset),
100.0 * batch_idx / len(loader),
).ljust(40),
"Loss: {:.6f}".format(loss.item()),
"BCELoss: {:.6f}".format(bce_loss.item()),
"KL_Loss: {:.6f}".format(kl_loss.item()),
)
Training
for i in range(1, 101):
train(
enc=encoder,
dec=decoder,
optimizer=optimizer,
loader=train_loader,
epoch=i,
single_pass_handler=vae_pass_handler,
loss_handler=vae_loss_handler,
log_interval=450,
)
I train almost 100 epochs and KL loss inreases while BCELoss decreases.
KL loss and BCELoss does not change if compare the first or second epoch and the last I can not figure out, why I do not get satisfying result?
encoder = encoder.eval()
decoder = decoder.eval()
Want to return construction images
def run_eval(
encoder,
decoder,
loader,
single_pass_handler,
return_real=True,
return_recon=True,
return_latent=True,
return_labels=True,
):
if return_real:
real = []
if return_recon:
reconstr = []
if return_latent:
latent = []
if return_labels:
labels = []
with torch.no_grad():
for batch_idx, (data, label) in enumerate(loader):
if return_labels:
labels.append(label.numpy())
if return_real:
real.append(data.numpy())
data = data.to(device)
label = label.to(device)
rep, rec = single_pass_handler(encoder, decoder, data, label)
if return_latent:
latent.append(rep.cpu().numpy())
if return_recon:
reconstr.append(rec.cpu().numpy())
result = {}
if return_real:
real = np.concatenate(real)
result["real"] = real.squeeze()
if return_latent:
latent = np.concatenate(latent)
result["latent"] = latent
if return_recon:
reconstr = np.concatenate(reconstr)
result["reconstr"] = reconstr.squeeze()
if return_labels:
labels = np.concatenate(labels)
result["labels"] = labels
return result
Start
run_res = run_eval(encoder, decoder, test_loader, vae_pass_handler)
Plot manifold
plot_manifold(run_res['latent'], run_res['labels'])
I got this picture, I suppose it is not normal, because using KL-divergence we want to make our classes as clasters. Apllying this code on the MNIST data I get
Visualize reconstruction images
And I got this picture...
visualize(run_res['reconstr'], run_res['labels'], num_imgs=16)
Why does it happen? I make a wrong VAE? Although this model, this code for MNIST data works not so bad, but for RGB images, I get nothing





