MCPcopy Create free account
hub / github.com/Tencent/embedx / EmbedxGraph

Class EmbedxGraph

demo/data/scripts/embedx_graph.py:128–219  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

126
127
128class EmbedxGraph:
129
130 def __init__(self, nx_graph):
131 self.data = json_graph.node_link_data(nx_graph)
132 self.context = Context(json_graph.node_link_graph(self.data))
133 self.has_feature = self.__has_feature()
134
135 def next_node(self):
136 assert "nodes" in self.data
137 for json_node in self.data["nodes"]:
138 yield EmbedxNode(json_node)
139
140 def statistics(self):
141 stats = [("Total edge number", self.context.total_edge_num())]
142 total_node_num, train_node_num, test_node_num = 0, 0, 0
143 max_label, num_label = 0, 0
144 total_feature_num, total_valid_feature_num = 0, 0
145 for node in self.next_node():
146 total_node_num += 1
147 if node.stage == "train":
148 train_node_num += 1
149 if node.stage == "test":
150 test_node_num += 1
151
152 if isinstance(node.label, int):
153 max_label = max(max_label, node.label)
154 num_label = max(max_label, node.label) + 1
155 else:
156 max_label = len(node.label) - 1
157 num_label = len(node.label)
158 if self.has_feature:
159 total_feature_num += len(node.feature)
160 total_valid_feature_num += len(node.valid_feature)
161 stats.append(("Total node number", total_node_num))
162 stats.append(("Total train node number", train_node_num))
163 stats.append(("Total test node number", test_node_num))
164 stats.append(("Total feature number", total_feature_num))
165 stats.append(("Total valid feature number", total_valid_feature_num))
166 if total_feature_num > 0:
167 ratio = 100.0 * (float(total_valid_feature_num) / total_feature_num)
168 stats.append(("Valid feature ratio", ratio))
169 stats.append(("Max label", max_label))
170 stats.append(("Label number", num_label))
171
172 return stats
173
174 def write_context(self, context_path):
175 with open(context_path, 'w', encoding='utf-8') as fout:
176 for src_id in self.context.keys():
177 neighbor = self.context.find_neighbor(src_id)
178 neighbor_str = Context.get_neighbor_str(src_id, neighbor)
179 fout.write(neighbor_str)
180 fout.write("\n")
181
182 def write_node_feature(self, node_feature_path, group_config_path):
183 assert self.has_feature
184 max_feature_num = 0
185 with open(node_feature_path, 'w', encoding='utf-8') as fout:

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected