| 92 | |
| 93 | |
| 94 | class DecoderModel(nn.Module, Seq2SeqAttrs): |
| 95 | def __init__(self, **model_kwargs): |
| 96 | # super().__init__(is_training, adj_mx, **model_kwargs) |
| 97 | nn.Module.__init__(self) |
| 98 | Seq2SeqAttrs.__init__(self, **model_kwargs) |
| 99 | self.output_dim = int(model_kwargs.get('output_dim', 1)) |
| 100 | self.horizon = int(model_kwargs.get('horizon', 1)) # for the decoder |
| 101 | self.projection_layer = nn.Linear(self.rnn_units, self.output_dim) |
| 102 | self.dcgru_layers = nn.ModuleList( |
| 103 | [DCGRUCell(self.rnn_units, self.max_diffusion_step, self.num_nodes, |
| 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): |