Views
No views yet
1import torch
2
3class ImageGenerationVAE(nn.Module, PyTorchModelHubMixin):
4 def __init__(self, hidden_dim=3000):
5 super(ImageGenerationVAE, self).__init__()
6
7 self.encoder = nn.Sequential(
8 nn.Conv2d(1, 64, kernel_size=3, stride=1, padding=1),
9 nn.MaxPool2d(kernel_size=2, stride=2),
10 nn.ReLU(),
11 nn.BatchNorm2d(64),
12 nn.Conv2d(64, 32, kernel_size=3, stride=1, padding=1),
13 nn.MaxPool2d(kernel_size=2, stride=2),
14 nn.ReLU(),
15 nn.BatchNorm2d(32),
16 nn.Flatten(),
17 )
18
19 self.mu_nn = nn.Linear(32 * 7 * 7, hidden_dim)
20 self.sigma_nn = nn.Linear(32 * 7 * 7, hidden_dim)
21
22 self.decoder = nn.Sequential(
23 nn.Linear(hidden_dim, 32 * 7 * 7),
24 nn.ReLU(),
25 nn.Unflatten(1, (32, 7, 7)),
26 nn.ConvTranspose2d(32, 64, kernel_size=2, stride=2),
27 nn.ReLU(),
28 nn.BatchNorm2d(64),
29 nn.ConvTranspose2d(64, 1, kernel_size=2, stride=2),
30 nn.Sigmoid(),
31 )
32
33 self.kl = 0
34
35 def forward(self, x):
36 # encoding head
37 x = self.encoder(x)
38 mu = self.mu_nn(x)
39 log_var = self.sigma_nn(x)
40
41 std = torch.exp(log_var) # More stable than full exp
42 epsilon = torch.randn_like(std)
43 z = mu + std * epsilon
44
45 # decoding head
46 output = self.decoder(z)
47
48 # KL divergence
49 self.kl = -0.5 * torch.sum(1 + log_var - mu.pow(2) - log_var.exp())
50 self.kl = self.kl / x.size(0) # normalize by batch size
51
52 return output