| 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 | ''' |