(self, node_dims, edge_dims,
n_message=3, n_feedforward=2, drop_rate=.1,
autoregressive=False,
activations=(F.relu, torch.sigmoid), vector_gate=False)
| 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): |
nothing calls this directly
no test coverage detected