-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathvae_model.py
More file actions
78 lines (66 loc) · 2.5 KB
/
Copy pathvae_model.py
File metadata and controls
78 lines (66 loc) · 2.5 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
import torch
import torch.nn as nn
class Unflatten(nn.Module):
def __init__(self, channels=128, height=16, width=16):
super(Unflatten, self).__init__()
self.channels = channels
self.height = height
self.width = width
def forward(self, input):
return input.view(input.size(0), self.channels, self.height, self.width)
class ConvolutionnalVAE(nn.Module):
def __init__(self, image_channels=3, z_dim=32, input_size=128):
super(ConvolutionnalVAE, self).__init__()
self.h_dim = 128 * 16 * 16
# Encoder
self.encoder = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=4, stride=2, padding=1),
nn.LayerNorm([32,64,64]),
nn.ReLU(),
nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=1),
nn.LayerNorm([64,32,32]),
nn.ReLU(),
nn.Conv2d(64, 128, kernel_size=4, stride=2, padding=1),
nn.LayerNorm([128,16,16]),
nn.ReLU(),
nn.Flatten()
)
# Latent space layers
self.l_mu = nn.Linear(self.h_dim, z_dim)
self.l_logvar = nn.Linear(self.h_dim, z_dim)
self.dec_projection = nn.Linear(z_dim, self.h_dim)
# Decoder
self.decoder = nn.Sequential(
Unflatten(channels=128, height=16, width=16),
nn.ConvTranspose2d(128, 32, kernel_size=4, stride=2, padding=1),
nn.LayerNorm([32,32,32]),
nn.ReLU(),
nn.ConvTranspose2d(32, 16, kernel_size=4, stride=2, padding=1),
nn.LayerNorm([16,64,64]),
nn.ReLU(),
nn.ConvTranspose2d(16, 3, kernel_size=4, stride=2, padding=1),
)
def reparametrize(self, mu, log_var):
std = log_var.mul(0.5).exp_()
epsilon = torch.randn_like(mu)
z = mu + std * epsilon
return z
def bottleneck(self, h):
mu = self.l_mu(h)
log_var = self.l_logvar(h)
z = self.reparametrize(mu, log_var)
return z, mu, log_var
def encode(self, x):
h = self.encoder(x)
z, mu, log_var = self.bottleneck(h)
return z, mu, log_var
def decode(self, z):
z = self.dec_projection(z)
z = self.decoder(z)
z = torch.clamp(z, min=-10, max=10)
z = torch.sigmoid(z)
return z
def forward(self, x):
z, mu, logvar = self.encode(x)
recon_x = self.decode(z)
return recon_x, mu, logvar