MCPcopy Create free account
hub / github.com/MotrixLab/MotionDiffuse / TextVAEDecoder

Class TextVAEDecoder

text2motion/datasets/evaluator_models.py:123–184  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

121
122
123class TextVAEDecoder(nn.Module):
124 def __init__(self, text_size, input_size, output_size, hidden_size, n_layers):
125 super(TextVAEDecoder, self).__init__()
126 self.input_size = input_size
127 self.output_size = output_size
128 self.hidden_size = hidden_size
129 self.n_layers = n_layers
130 self.emb = nn.Sequential(
131 nn.Linear(input_size, hidden_size),
132 nn.LayerNorm(hidden_size),
133 nn.LeakyReLU(0.2, inplace=True))
134
135 self.z2init = nn.Linear(text_size, hidden_size * n_layers)
136 self.gru = nn.ModuleList([nn.GRUCell(hidden_size, hidden_size) for i in range(self.n_layers)])
137 self.positional_encoder = PositionalEncoding(hidden_size)
138
139
140 self.output = nn.Sequential(
141 nn.Linear(hidden_size, hidden_size),
142 nn.LayerNorm(hidden_size),
143 nn.LeakyReLU(0.2, inplace=True),
144 nn.Linear(hidden_size, output_size)
145 )
146
147 #
148 # self.output = nn.Sequential(
149 # nn.Linear(hidden_size, hidden_size),
150 # nn.LayerNorm(hidden_size),
151 # nn.LeakyReLU(0.2, inplace=True),
152 # nn.Linear(hidden_size, output_size-4)
153 # )
154
155 # self.contact_net = nn.Sequential(
156 # nn.Linear(output_size-4, 64),
157 # nn.LayerNorm(64),
158 # nn.LeakyReLU(0.2, inplace=True),
159 # nn.Linear(64, 4)
160 # )
161
162 self.output.apply(init_weight)
163 self.emb.apply(init_weight)
164 self.z2init.apply(init_weight)
165 # self.contact_net.apply(init_weight)
166
167 def get_init_hidden(self, latent):
168 hidden = self.z2init(latent)
169 hidden = torch.split(hidden, self.hidden_size, dim=-1)
170 return list(hidden)
171
172 def forward(self, inputs, last_pred, hidden, p):
173 h_in = self.emb(inputs)
174 pos_enc = self.positional_encoder(p).to(inputs.device).detach()
175 h_in = h_in + pos_enc
176 for i in range(self.n_layers):
177 # print(h_in.shape)
178 hidden[i] = self.gru[i](h_in, hidden[i])
179 h_in = hidden[i]
180 pose_pred = self.output(h_in)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected