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

Method forward

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

: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)

Source from the content-addressed store, hash-verified

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

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected