| 58 | |
| 59 | |
| 60 | class EncoderModel(nn.Module, Seq2SeqAttrs): |
| 61 | def __init__(self, **model_kwargs): |
| 62 | nn.Module.__init__(self) |
| 63 | Seq2SeqAttrs.__init__(self, **model_kwargs) |
| 64 | self.input_dim = int(model_kwargs.get('input_dim', 1)) |
| 65 | self.seq_len = int(model_kwargs.get('seq_len')) # for the encoder |
| 66 | self.dcgru_layers = nn.ModuleList( |
| 67 | [DCGRUCell(self.rnn_units, self.max_diffusion_step, self.num_nodes, |
| 68 | filter_type=self.filter_type) for _ in range(self.num_rnn_layers)]) |
| 69 | |
| 70 | def forward(self, inputs, adj, hidden_state=None): |
| 71 | """ |
| 72 | Encoder forward pass. |
| 73 | :param inputs: shape (batch_size, self.num_nodes * self.input_dim) |
| 74 | :param hidden_state: (num_layers, batch_size, self.hidden_state_size) |
| 75 | optional, zeros if not provided |
| 76 | :return: output: # shape (batch_size, self.hidden_state_size) |
| 77 | hidden_state # shape (num_layers, batch_size, self.hidden_state_size) |
| 78 | (lower indices mean lower layers) |
| 79 | """ |
| 80 | batch_size, _ = inputs.size() |
| 81 | if hidden_state is None: |
| 82 | hidden_state = torch.zeros((self.num_rnn_layers, batch_size, self.hidden_state_size), |
| 83 | device=device) |
| 84 | hidden_states = [] |
| 85 | output = inputs |
| 86 | for layer_num, dcgru_layer in enumerate(self.dcgru_layers): |
| 87 | next_hidden_state = dcgru_layer(output, hidden_state[layer_num], adj) |
| 88 | hidden_states.append(next_hidden_state) |
| 89 | output = next_hidden_state |
| 90 | |
| 91 | return output, torch.stack(hidden_states) # runs in O(num_layers) so not too slow |
| 92 | |
| 93 | |
| 94 | class DecoderModel(nn.Module, Seq2SeqAttrs): |