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

Method __init__

s3f/gvp_layer.py:314–337  ·  view source on GitHub ↗
(self, node_dims, edge_dims,
                 n_message=3, n_feedforward=2, drop_rate=.1,
                 autoregressive=False, 
                 activations=(F.relu, torch.sigmoid), vector_gate=False)

Source from the content-addressed store, hash-verified

312 (vector_act will be used as sigma^+ in vector gating if `True`)
313 '''
314 def __init__(self, node_dims, edge_dims,
315 n_message=3, n_feedforward=2, drop_rate=.1,
316 autoregressive=False,
317 activations=(F.relu, torch.sigmoid), vector_gate=False):
318
319 super(GVPConvLayer, self).__init__()
320 self.conv = GVPConv(node_dims, node_dims, edge_dims, n_message,
321 aggr="add" if autoregressive else "mean",
322 activations=activations, vector_gate=vector_gate)
323 GVP_ = functools.partial(GVP,
324 activations=activations, vector_gate=vector_gate)
325 self.norm = nn.ModuleList([GVPLayerNorm(node_dims) for _ in range(2)])
326 self.dropout = nn.ModuleList([Dropout(drop_rate) for _ in range(2)])
327
328 ff_func = []
329 if n_feedforward == 1:
330 ff_func.append(GVP_(node_dims, node_dims, activations=(None, None)))
331 else:
332 hid_dims = 4*node_dims[0], 2*node_dims[1]
333 ff_func.append(GVP_(node_dims, hid_dims))
334 for i in range(n_feedforward-2):
335 ff_func.append(GVP_(hid_dims, hid_dims))
336 ff_func.append(GVP_(hid_dims, node_dims, activations=(None, None)))
337 self.ff_func = nn.Sequential(*ff_func)
338
339 def forward(self, x, edge_index, edge_attr,
340 autoregressive_x=None, node_mask=None):

Callers

nothing calls this directly

Calls 4

GVPConvClass · 0.85
GVPLayerNormClass · 0.85
DropoutClass · 0.85
__init__Method · 0.45

Tested by

no test coverage detected