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

Function augment_edge

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

Input: data: PyG data object Output: data (edges are augmented in the following ways): data.edge_index: Added next-token edge. The inverse edges were also added. data.edge_attr (torch.Long): data.edge_attr[:

(data, num_edge_type)

Source from the content-addressed store, hash-verified

36
37# VT: with ourgraphs the below "AST" edges etc do not have to be AST edges but may be arbitray (data flow etc)
38def augment_edge(data, num_edge_type):
39 '''
40 Input:
41 data: PyG data object
42 Output:
43 data (edges are augmented in the following ways):
44 data.edge_index: Added next-token edge. The inverse edges were also added.
45 data.edge_attr (torch.Long):
46 data.edge_attr[:,0]: whether it is a graph edge (0) for next-token edge (1)
47 data.edge_attr[:,1]: whether it is original direction (0) or inverse direction (1)
48 IF IN DATA: n types for graph edges
49 data.edge_attr[:, 1+n]: types, n > 0
50
51 '''
52 if num_edge_type > 0: # hasattr(data, "edge_type"):
53 # num_edge_type = max(data.edge_type).item()+1
54 idx = data.edge_type ##.view(-1, 1)
55 edge_type = torch.zeros(idx.size()[0], num_edge_type).scatter_(1, idx, 1)
56 # y_one_hot = y_one_hot.view(*y.shape, -1)
57 else:
58 # num_edge_type = 0
59 edge_type = torch.zeros(data.edge_index.size()[1], 0)
60 # print(num_edge_type)
61
62 ##### AST edge
63 edge_index_ast = data.edge_index
64 edge_attr_ast = torch.zeros((edge_index_ast.size(1), 2))
65 if num_edge_type:
66 edge_attr_ast = torch.cat([edge_attr_ast, edge_type], dim=-1)
67
68 ##### Inverse AST edge
69 edge_index_ast_inverse = torch.stack([edge_index_ast[1], edge_index_ast[0]], dim = 0)
70 edge_attr_ast_inverse = torch.cat([torch.zeros(edge_index_ast_inverse.size(1), 1), torch.ones(edge_index_ast_inverse.size(1), 1)], dim = 1)
71 if num_edge_type:
72 edge_attr_ast_inverse = torch.cat([edge_attr_ast_inverse, edge_type], dim=-1)
73
74 ##### Next-token edge
75
76 ## Obtain attributed nodes and get their indices in dfs order
77 # attributed_node_idx = torch.where(data.node_is_attributed.view(-1,) == 1)[0]
78 # attributed_node_idx_in_dfs_order = attributed_node_idx[torch.argsort(data.node_dfs_order[attributed_node_idx].view(-1,))]
79
80 ## Since the nodes are already sorted in dfs ordering in our case, we can just do the following.
81 attributed_node_idx_in_dfs_order = torch.where(data.node_is_attributed.view(-1,) == 1)[0]
82
83 ## build next token edge
84 # Given: attributed_node_idx_in_dfs_order
85 # [1, 3, 4, 5, 8, 9, 12]
86 # Output:
87 # [[1, 3, 4, 5, 8, 9]
88 # [3, 4, 5, 8, 9, 12]
89 edge_index_nextoken = torch.stack([attributed_node_idx_in_dfs_order[:-1], attributed_node_idx_in_dfs_order[1:]], dim = 0)
90 edge_attr_nextoken = torch.cat([torch.ones(edge_index_nextoken.size(1), 1), torch.zeros(edge_index_nextoken.size(1), 1+num_edge_type)], dim = 1)
91
92
93 ##### Inverse next-token edge
94 edge_index_nextoken_inverse = torch.stack([edge_index_nextoken[1], edge_index_nextoken[0]], dim = 0)
95 edge_attr_nextoken_inverse = torch.ones((edge_index_nextoken.size(1), 2))

Callers 1

mainFunction · 0.90

Calls 1

sizeMethod · 0.45

Tested by

no test coverage detected