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)
| 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 | """ |
no test coverage detected