| 80 | |
| 81 | |
| 82 | class COND_NET(nn.Module): #not chnaged yet |
| 83 | # some code is modified from vae examples |
| 84 | # (https://github.com/pytorch/examples/blob/master/vae/main.py) |
| 85 | def __init__(self): |
| 86 | super(COND_NET, self).__init__() |
| 87 | self.t_dim = 14 |
| 88 | self.c_dim = 10 |
| 89 | self.fc = nn.Linear(self.t_dim, self.c_dim, bias=True) |
| 90 | self.relu = nn.PReLU()#nn.ReLU() |
| 91 | |
| 92 | def encode(self, full_embed): |
| 93 | x = self.relu(self.fc(full_embed)) |
| 94 | # mu = x[:, :self.c_dim] |
| 95 | # logvar = x[:, self.c_dim:] |
| 96 | return x |
| 97 | |
| 98 | # def reparametrize(self, mu, logvar): |
| 99 | # std = logvar.mul(0.5).exp_() |
| 100 | # if cfg.CUDA: |
| 101 | # eps = torch.cuda.FloatTensor(std.size()).normal_() |
| 102 | # else: |
| 103 | # eps = torch.FloatTensor(std.size()).normal_() |
| 104 | # eps = Variable(eps) |
| 105 | # return eps.mul(std).add_(mu) |
| 106 | |
| 107 | def forward(self, full_embed): |
| 108 | c_code = self.encode(full_embed) |
| 109 | # c_code = self.reparametrize(mu, logvar) |
| 110 | return c_code #, mu, logvar |
| 111 | |
| 112 | class MESH_NET(nn.Module): |
| 113 | def __init__(self): |