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

Method forward

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

Encoder forward pass. :param inputs: shape (batch_size, self.num_nodes * self.input_dim) :param hidden_state: (num_layers, batch_size, self.hidden_state_size) optional, zeros if not provided :return: output: # shape (batch_size, self.hidden_state_size)

(self, inputs, adj, hidden_state=None)

Source from the content-addressed store, hash-verified

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

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected