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

Class DecoderModel

model/pytorch/model.py:94–125  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

92
93
94class 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
128class GTSModel(nn.Module, Seq2SeqAttrs):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected