:param inputs: shape (batch_size, self.num_nodes * self.output_dim) :param hidden_state: (num_layers, batch_size, self.hidden_state_size) optional, zeros if not provided :return: output: # shape (batch_size, self.num_nodes * self.output_dim) h
(self, inputs, adj, hidden_state=None)
| 104 | filter_type=self.filter_type) for _ in range(self.num_rnn_layers)]) |
| 105 | |
| 106 | def forward(self, inputs, adj, hidden_state=None): |
| 107 | """ |
| 108 | :param inputs: shape (batch_size, self.num_nodes * self.output_dim) |
| 109 | :param hidden_state: (num_layers, batch_size, self.hidden_state_size) |
| 110 | optional, zeros if not provided |
| 111 | :return: output: # shape (batch_size, self.num_nodes * self.output_dim) |
| 112 | hidden_state # shape (num_layers, batch_size, self.hidden_state_size) |
| 113 | (lower indices mean lower layers) |
| 114 | """ |
| 115 | hidden_states = [] |
| 116 | output = inputs |
| 117 | for layer_num, dcgru_layer in enumerate(self.dcgru_layers): |
| 118 | next_hidden_state = dcgru_layer(output, hidden_state[layer_num], adj) |
| 119 | hidden_states.append(next_hidden_state) |
| 120 | output = next_hidden_state |
| 121 | |
| 122 | projected = self.projection_layer(output.view(-1, self.rnn_units)) |
| 123 | output = projected.view(-1, self.num_nodes * self.output_dim) |
| 124 | |
| 125 | return output, torch.stack(hidden_states) |
| 126 | |
| 127 | |
| 128 | class GTSModel(nn.Module, Seq2SeqAttrs): |
nothing calls this directly
no outgoing calls
no test coverage detected