MCPcopy Create free account
hub / github.com/chaoshangcs/GTS / decoder

Method decoder

model/pytorch/model.py:180–206  ·  view source on GitHub ↗

Decoder forward pass :param encoder_hidden_state: (num_layers, batch_size, self.hidden_state_size) :param labels: (self.horizon, batch_size, self.num_nodes * self.output_dim) [optional, not exist for inference] :param batches_seen: global step [optional, not exist fo

(self, encoder_hidden_state, adj, labels=None, batches_seen=None)

Source from the content-addressed store, hash-verified

178 return encoder_hidden_state
179
180 def decoder(self, encoder_hidden_state, adj, labels=None, batches_seen=None):
181 """
182 Decoder forward pass
183 :param encoder_hidden_state: (num_layers, batch_size, self.hidden_state_size)
184 :param labels: (self.horizon, batch_size, self.num_nodes * self.output_dim) [optional, not exist for inference]
185 :param batches_seen: global step [optional, not exist for inference]
186 :return: output: (self.horizon, batch_size, self.num_nodes * self.output_dim)
187 """
188 batch_size = encoder_hidden_state.size(1)
189 go_symbol = torch.zeros((batch_size, self.num_nodes * self.decoder_model.output_dim),
190 device=device)
191 decoder_hidden_state = encoder_hidden_state
192 decoder_input = go_symbol
193
194 outputs = []
195
196 for t in range(self.decoder_model.horizon):
197 decoder_output, decoder_hidden_state = self.decoder_model(decoder_input, adj,
198 decoder_hidden_state)
199 decoder_input = decoder_output
200 outputs.append(decoder_output)
201 if self.training and self.use_curriculum_learning:
202 c = np.random.uniform(0, 1)
203 if c < self._compute_sampling_threshold(batches_seen):
204 decoder_input = labels[t]
205 outputs = torch.stack(outputs)
206 return outputs
207
208 def forward(self, label, inputs, node_feas, temp, gumbel_soft, labels=None, batches_seen=None):
209 """

Callers 1

forwardMethod · 0.95

Calls 1

Tested by

no test coverage detected