(self,latent_dims)
| 4 | |
| 5 | class VariationalEncoder(nn.Module): |
| 6 | def __init__(self,latent_dims): |
| 7 | super(VariationalEncoder,self).__init__() |
| 8 | self.linear1 = nn.Linear(256,128) |
| 9 | self.linear2 = nn.Linear(128,64) |
| 10 | self.linear3 = nn.Linear(64,latent_dims) |
| 11 | self.linear4 = nn.Linear(64,latent_dims) |
| 12 | |
| 13 | self.N = torch.distributions.Normal(0,1) |
| 14 | self.N.loc = self.N.loc.cuda() |
| 15 | self.N.scale = self.N.scale.cuda() |
| 16 | self.kl = 0 |
| 17 | |
| 18 | def forward(self,x): |
| 19 | x = F.relu(self.linear1(x)) |