MCPcopy Create free account
hub / github.com/DeepGraphLearning/S3F / forward

Method forward

s3f/gvp.py:222–242  ·  view source on GitHub ↗
(self, graph, input, all_loss=None, metric=None)

Source from the content-addressed store, hash-verified

220 raise ValueError("Unknown readout `%s`" % readout)
221
222 def forward(self, graph, input, all_loss=None, metric=None):
223 h_node = self.residue_embdding(input)
224
225 edge_index = graph.edge_list.t()[:2]
226 node_in, node_out = edge_index
227 pos_in, pos_out = graph.node_position[node_in], graph.node_position[node_out]
228 vec_edge = (pos_out - pos_in).unsqueeze(-2) # [n_edge, 1, 3]
229 h_edge = rbf((pos_out - pos_in).norm(dim=-1), dim=self.rbf_dim), vec_edge
230
231 h_node = self.W_v(h_node)
232 h_edge = self.W_e(h_edge)
233 for layer in self.layers:
234 h_node = layer(h_node, edge_index, h_edge)
235 node_feature = self.W_out(h_node)
236
237 graph_feature = self.readout(graph, node_feature)
238
239 return {
240 "graph_feature": graph_feature,
241 "node_feature": node_feature
242 }

Callers

nothing calls this directly

Calls 1

rbfFunction · 0.85

Tested by

no test coverage detected