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)
| 36 | |
| 37 | # VT: with ourgraphs the below "AST" edges etc do not have to be AST edges but may be arbitray (data flow etc) |
| 38 | def 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)) |