-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathoffline_graph_creation.py
More file actions
88 lines (74 loc) · 2.66 KB
/
Copy pathoffline_graph_creation.py
File metadata and controls
88 lines (74 loc) · 2.66 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
import json
import os
import torch
import shutil
import copy
import random
import argparse
import numpy as np
import configparser
class Graph_sampling():
def __init__(self, path_in, path_out):
self.path_in = path_in
self.path_out = path_out
self.files = self.find_files()
def find_files(self):
files = []
for r, d, f in os.walk(self.path_in):
for file in f:
if '.json' in file:
files.append(os.path.join(r, file))
files.sort()
return files
def sampling_graphs(self, r_tar, r_all, length):
possible_pairs_drug = []
possible_pairs_ae = []
for i in range(length):
if ([r_tar[0], i] not in r_all) and ([i, r_tar[0]] not in r_all) and (r_tar[0] != i):
possible_pairs_drug.append([r_tar[0], i])
if ([i, r_tar[1]] not in r_all) and ([r_tar[1], i] not in r_all) and (r_tar[1] != i):
possible_pairs_ae.append([i, r_tar[1]])
return possible_pairs_drug, possible_pairs_ae
def execute_sampling(self):
counter = 0
for f in self.files:
filename = f.split('/')[-1]
with open(f) as json_file:
data = json.load(json_file)
dict_out = {}
for i, r in enumerate(data['relation pairs']):
emb_1 = data['node initialization'][0][r[0]]
emb_2 = data['node initialization'][0][r[1]]
sampled_pairs_drug, sampled_pairs_ae = self.sampling_graphs(r, data['relation pairs'], len(data['tokens']))
dict_out['r_' + str(i+1)] = [[emb_1, emb_2]]
dict_out['r_' + str(i+1) + '_indexes'] = [r]
dict_out['r_' + str(i+1) + '_rev'] = []
dict_out['r_' + str(i+1) + '_indexes_rev'] = []
for p in sampled_pairs_drug:
emb_1_t = data['node initialization'][0][p[0]]
emb_2_t = data['node initialization'][0][p[1]]
dict_out['r_' + str(i+1)].append([emb_1_t, emb_2_t])
dict_out['r_' + str(i+1) + '_indexes'].append(p)
for p in sampled_pairs_ae:
emb_1_t = data['node initialization'][0][p[0]]
emb_2_t = data['node initialization'][0][p[1]]
dict_out['r_' + str(i+1) + '_rev'].append([emb_1_t, emb_2_t])
dict_out['r_' + str(i+1) + '_indexes_rev'].append(p)
with open(self.path_out + filename, 'w') as fp:
json.dump(dict_out, fp)
counter += 1
if counter % 100 == 0:
print('{} files processed.' .format(counter))
if __name__ == "__main__":
parser = configparser.ConfigParser()
parser.read("./configs/offline_graphs.conf")
path_in = parser.get("config", "input_path")
path_out = parser.get("config", "output_path")
if not os.path.isdir(path_in):
print('The path to processed data does not exist.')
sys.exit()
if not os.path.isdir(path_out):
print('Creating output path.')
os.makedirs(path_out)
graph_sampling_obj = Graph_sampling(path_in, path_out)
graph_sampling_obj.execute_sampling()