| 145 | |
| 146 | |
| 147 | class 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 |
no outgoing calls
no test coverage detected