(self, node_embeds, label_embeds, data)
| 79 | return CL_loss |
| 80 | |
| 81 | def zero_shot_eval(self, node_embeds, label_embeds, data): |
| 82 | if self.args.if_norm: |
| 83 | node_embeds = (node_embeds - node_embeds.mean(0)) / \ |
| 84 | node_embeds.std(0) |
| 85 | label_embeds = (label_embeds - label_embeds.mean(0) |
| 86 | ) / label_embeds.std(0) |
| 87 | |
| 88 | # change the adj matrix |
| 89 | num_existing_nodes = data.y.shape[0] + 1 |
| 90 | virtual_node_index = data.y.shape[0] |
| 91 | if self.args.test_data in ["Citeseer"]: |
| 92 | new_edges_to_virtual = [[node_idx, virtual_node_index] |
| 93 | for node_idx in range(num_existing_nodes-1)] |
| 94 | elif self.args.test_data in ["Cora", "Pubmed", "Citeseer", "Arxiv", "wikics", "facebook", 'home', 'tech']: |
| 95 | new_edges_to_virtual = [] |
| 96 | for node_idx in range(num_existing_nodes-1): |
| 97 | new_edges_to_virtual.append([node_idx, virtual_node_index]) |
| 98 | new_edges_to_virtual.append([virtual_node_index, node_idx]) |
| 99 | new_edge_index = torch.cat([data.edge_index.t(), torch.tensor( |
| 100 | new_edges_to_virtual, dtype=torch.long).to(self.args.device)], dim=0).t() |
| 101 | adj_normed = self.normalize_adjacency_matrix( |
| 102 | new_edge_index, num_existing_nodes) |
| 103 | |
| 104 | # adj_normed = self.normalize_adjacency_matrix(data) |
| 105 | for _ in range(self.args.R): |
| 106 | node_embeds = torch.mm(adj_normed, node_embeds) |
| 107 | node_embeds = node_embeds[:-1, :] |
| 108 | node_embeds /= node_embeds.norm(dim=-1, |
| 109 | keepdim=True).to(self.args.device) |
| 110 | label_embeds /= label_embeds.norm(dim=-1, keepdim=True) |
| 111 | dists = torch.einsum('bn,cn->bc', node_embeds, label_embeds) |
| 112 | preds = torch.argmax(dists, dim=1) |
| 113 | labels = data.y.long().to(self.args.device) |
| 114 | test_acc = accuracy_score(labels.cpu(), preds.cpu()) |
| 115 | return test_acc |
| 116 | |
| 117 | def normalize_adjacency_matrix(self, edge_index, num_nodes): |
| 118 | # edge_index = data.edge_index |
no test coverage detected