MCPcopy Create free account
hub / github.com/anton-jeran/MESH2IR / COND_NET

Class COND_NET

evaluate/model.py:82–110  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

80
81
82class 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
112class MESH_NET(nn.Module):
113 def __init__(self):

Callers 1

define_moduleMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected