(self, node_embeds, label_embeds, data)
| 255 | return CL_loss |
| 256 | |
| 257 | def zero_shot_eval(self, node_embeds, label_embeds, data): |
| 258 | if self.args.if_norm: |
| 259 | node_embeds = (node_embeds - node_embeds.mean(0)) / \ |
| 260 | node_embeds.std(0) |
| 261 | label_embeds = (label_embeds - label_embeds.mean(0) |
| 262 | ) / label_embeds.std(0) |
| 263 | |
| 264 | # change the adj matrix |
| 265 | num_existing_nodes = data.y.shape[0] + 1 |
| 266 | virtual_node_index = data.y.shape[0] |
| 267 | if self.args.test_data in ["Citesee"]: |
| 268 | new_edges_to_virtual = [[node_idx, virtual_node_index] |
| 269 | for node_idx in range(num_existing_nodes-1)] |
| 270 | elif self.args.test_data in ["Cora", "Pubmed", "Citeseer", "Arxiv", "wikics", "facebook", 'home', 'tech','reddit','instagram']: |
| 271 | new_edges_to_virtual = [] |
| 272 | for node_idx in range(num_existing_nodes-1): |
| 273 | new_edges_to_virtual.append([node_idx, virtual_node_index]) |
| 274 | new_edges_to_virtual.append([virtual_node_index, node_idx]) |
| 275 | new_edge_index = torch.cat([data.edge_index.t(), torch.tensor( |
| 276 | new_edges_to_virtual, dtype=torch.long).to(self.args.device)], dim=0).t() |
| 277 | adj_normed = self.normalize_adjacency_matrix( |
| 278 | new_edge_index, num_existing_nodes) |
| 279 | |
| 280 | # adj_normed = self.normalize_adjacency_matrix(data) |
| 281 | for _ in range(self.args.R): |
| 282 | node_embeds = torch.mm(adj_normed, node_embeds) |
| 283 | node_embeds = node_embeds[:-1, :] |
| 284 | node_embeds /= node_embeds.norm(dim=-1, |
| 285 | keepdim=True).to(self.args.device) |
| 286 | label_embeds /= label_embeds.norm(dim=-1, keepdim=True) |
| 287 | dists = torch.einsum('bn,cn->bc', node_embeds, label_embeds) |
| 288 | preds = torch.argmax(dists, dim=1) |
| 289 | labels = data.y.long().to(self.args.device) |
| 290 | if len(data.test_mask) == 10 : |
| 291 | data.test_mask = data.test_mask[0] |
| 292 | test_mask = data.test_mask |
| 293 | test_acc = accuracy_score(labels[test_mask].cpu(), preds[test_mask].cpu()) |
| 294 | test_f1 = f1_score(labels[test_mask].cpu(), preds[test_mask].cpu()) |
| 295 | return [test_acc,test_f1] |
| 296 | |
| 297 | def normalize_adjacency_matrix(self, edge_index, num_nodes): |
| 298 | # edge_index = data.edge_index |
nothing calls this directly
no test coverage detected