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

Method __init__

s3f/gvp_layer.py:246–271  ·  view source on GitHub ↗
(self, in_dims, out_dims, edge_dims,
                 n_layers=3, module_list=None, aggr="mean", 
                 activations=(F.relu, torch.sigmoid), vector_gate=False)

Source from the content-addressed store, hash-verified

244 (vector_act will be used as sigma^+ in vector gating if `True`)
245 '''
246 def __init__(self, in_dims, out_dims, edge_dims,
247 n_layers=3, module_list=None, aggr="mean",
248 activations=(F.relu, torch.sigmoid), vector_gate=False):
249 super(GVPConv, self).__init__(aggr=aggr)
250 self.si, self.vi = in_dims
251 self.so, self.vo = out_dims
252 self.se, self.ve = edge_dims
253
254 GVP_ = functools.partial(GVP,
255 activations=activations, vector_gate=vector_gate)
256
257 module_list = module_list or []
258 if not module_list:
259 if n_layers == 1:
260 module_list.append(
261 GVP_((2*self.si + self.se, 2*self.vi + self.ve),
262 (self.so, self.vo), activations=(None, None)))
263 else:
264 module_list.append(
265 GVP_((2*self.si + self.se, 2*self.vi + self.ve), out_dims)
266 )
267 for i in range(n_layers - 2):
268 module_list.append(GVP_(out_dims, out_dims))
269 module_list.append(GVP_(out_dims, out_dims,
270 activations=(None, None)))
271 self.message_func = nn.Sequential(*module_list)
272
273 def forward(self, x, edge_index, edge_attr):
274 '''

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected