MCPcopy Create free account
hub / github.com/IBM/Project_CodeNet / GCNConv

Class GCNConv

model-experiments/gnn-based-experiments/src/model/conv.py:38–67  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

36
37### GCN convolution along the graph structure
38class 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

Callers 2

__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected