| 36 | |
| 37 | ### GCN convolution along the graph structure |
| 38 | class GCNConv(MessagePassing): |
| 39 | def __init__(self, emb_dim, edge_attr_dim): |
| 40 | super(GCNConv, self).__init__(aggr='add') |
| 41 | |
| 42 | self.linear = torch.nn.Linear(emb_dim, emb_dim) |
| 43 | self.root_emb = torch.nn.Embedding(1, emb_dim) |
| 44 | |
| 45 | # edge_attr is two dimensional after augment_edge transformation |
| 46 | self.edge_encoder = torch.nn.Linear(edge_attr_dim, emb_dim) |
| 47 | |
| 48 | def forward(self, x, edge_index, edge_attr): |
| 49 | x = self.linear(x) |
| 50 | edge_embedding = self.edge_encoder(edge_attr) |
| 51 | |
| 52 | row, col = edge_index |
| 53 | |
| 54 | #edge_weight = torch.ones((edge_index.size(1), ), device=edge_index.device) |
| 55 | deg = degree(row, x.size(0), dtype = x.dtype) + 1 |
| 56 | deg_inv_sqrt = deg.pow(-0.5) |
| 57 | deg_inv_sqrt[deg_inv_sqrt == float('inf')] = 0 |
| 58 | |
| 59 | norm = deg_inv_sqrt[row] * deg_inv_sqrt[col] |
| 60 | |
| 61 | return self.propagate(edge_index, x=x, edge_attr = edge_embedding, norm=norm) + F.relu(x + self.root_emb.weight) * 1./deg.view(-1,1) |
| 62 | |
| 63 | def message(self, x_j, edge_attr, norm): |
| 64 | return norm.view(-1, 1) * F.relu(x_j + edge_attr) |
| 65 | |
| 66 | def update(self, aggr_out): |
| 67 | return aggr_out |
| 68 | |
| 69 | |
| 70 | ### GNN to generate node embedding |