MCPcopy Create free account
hub / github.com/NineAbyss/ZeroG / kHopSubgraphDataset_Arxiv

Class kHopSubgraphDataset_Arxiv

code/SubgraphDataset.py:147–206  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

145
146
147class kHopSubgraphDataset_Arxiv(Dataset):
148 def __init__(self, data, num_hops=1, max_nodes=100, dataset_name="Arxiv", transform=None, pre_transform=None):
149 super(kHopSubgraphDataset_Arxiv, self).__init__(None, transform, pre_transform)
150 self.data = data
151 self.num_hops = num_hops
152 self.unique_classes = data.y.unique()
153 self.k_over_2 = len(self.unique_classes) // 2
154 if dataset_name == "Citeseer":
155 self.k_over_2 = 2
156 elif dataset_name == "Arxiv":
157 self.k_over_2 = 10
158 self.max_nodes = max_nodes
159 self.dataset_name = dataset_name
160 self.file_path = self.dataset_name + '_index.pkl'
161 if os.path.exists(self.file_path):
162 with open(self.file_path, 'rb') as f:
163 self.valid_subgraphs = pickle.load(f)
164 else:
165 self.valid_subgraphs = self._find_valid_subgraphs()
166
167 def _find_valid_subgraphs(self):
168 valid_subgraphs = []
169 for idx in tqdm(range(self.data.num_nodes)):
170 subgraph_node_idx, subgraph_edge_index, mapping, edge_mask = k_hop_subgraph(
171 node_idx=idx,
172 num_hops=self.num_hops,
173 edge_index=self.data.edge_index,
174 relabel_nodes=True,
175 num_nodes=self.data.num_nodes
176 )
177 unique_classes_in_subgraph = np.unique(self.data.y[subgraph_node_idx].cpu().numpy())
178 print("idx: {}, subgraph: {}, k: {}, nodes: {}".format(idx, len(valid_subgraphs), len(unique_classes_in_subgraph), len(subgraph_node_idx)))
179 if len(unique_classes_in_subgraph) >= self.k_over_2 and len(subgraph_node_idx) <= self.max_nodes:
180 valid_subgraphs.append(idx)
181 with open(self.file_path, 'wb') as f:
182 pickle.dump(valid_subgraphs, f)
183 return valid_subgraphs
184
185 def len(self):
186 return len(self.valid_subgraphs)
187
188 def get(self, idx):
189 subgraph_idx = self.valid_subgraphs[idx]
190 subgraph_node_idx, subgraph_edge_index, mapping, edge_mask = k_hop_subgraph(
191 node_idx=subgraph_idx,
192 num_hops=self.num_hops,
193 edge_index=self.data.edge_index,
194 relabel_nodes=True,
195 num_nodes=self.data.num_nodes
196 )
197 sub_data = Data(edge_index=subgraph_edge_index)
198 sub_data.y = self.data.y[subgraph_node_idx]
199 sub_data.raw_text = [self.data.raw_texts[i] for i in subgraph_node_idx.tolist()]
200 sub_data.label_text = self.data.label_text
201 sub_data.adjacency_matrix = to_dense_adj(subgraph_edge_index, max_num_nodes=mapping.size(0))[0]
202 sub_data.dataset_name = self.dataset_name
203 return sub_data
204

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected