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

Class EncoderModel

model/pytorch/model.py:60–91  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

58
59
60class 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
94class DecoderModel(nn.Module, Seq2SeqAttrs):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected