MCPcopy Create free account
hub / github.com/chengsen/PyTorch_TextGCN / BuildGraph

Class BuildGraph

build_graph.py:91–182  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

89
90
91class BuildGraph:
92 def __init__(self, dataset):
93 clean_corpus_path = "data/text_dataset/clean_corpus"
94 self.graph_path = "data/graph"
95 if not os.path.exists(self.graph_path):
96 os.makedirs(self.graph_path)
97
98 self.word2id = dict() # 单词映射
99 self.dataset = dataset
100 print(f"\n==> 现在的数据集是:{dataset}<==")
101
102 self.g = nx.Graph()
103
104 self.content = f"{clean_corpus_path}/{dataset}.txt"
105
106 self.get_tfidf_edge()
107 self.get_pmi_edge()
108 self.save()
109
110 def get_pmi_edge(self):
111 pmi_edge_lst, self.pmi_time = get_pmi_edge(self.content, window_size=20, threshold=0.0)
112 print("pmi time:", self.pmi_time)
113
114 for edge_item in pmi_edge_lst:
115 word_indx1 = self.node_num + self.word2id[edge_item[0]]
116 word_indx2 = self.node_num + self.word2id[edge_item[1]]
117 if word_indx1 == word_indx2:
118 continue
119 self.g.add_edge(word_indx1, word_indx2, weight=edge_item[2])
120
121 print_graph_detail(self.g)
122
123 def get_tfidf_edge(self):
124 # 获得tfidf权重矩阵(sparse)和单词列表
125 tfidf_vec = self.get_tfidf_vec()
126
127 count_lst = list() # 统计每个句子的长度
128 for ind, row in tqdm(enumerate(tfidf_vec),
129 desc="generate tfidf edge"):
130 count = 0
131 for col_ind, value in zip(row.indices, row.data):
132 word_ind = self.node_num + col_ind
133 self.g.add_edge(ind, word_ind, weight=value)
134 count += 1
135 count_lst.append(count)
136
137 print_graph_detail(self.g)
138
139 def get_tfidf_vec(self):
140 """
141 学习获得tfidf矩阵,及其对应的单词序列
142 :param content_lst:
143 :return:
144 """
145 start = time()
146 text_tfidf = Pipeline([
147 ("vect", CountVectorizer(min_df=1,
148 max_df=1.0,

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected