(self, data)
| 31 | self.criteria = nn.CrossEntropyLoss() |
| 32 | |
| 33 | def forward(self, data): |
| 34 | # insert a description node |
| 35 | virtual_node_description = self.descriptions[data.dataset_name] |
| 36 | all_node_texts = data.raw_text + [virtual_node_description] |
| 37 | |
| 38 | tokens = self.tokenizer(all_node_texts, max_length=256, return_tensors='pt', |
| 39 | truncation=True, padding=True).to(self.args.device) |
| 40 | node_embeds = self.textmodel(**tokens)[0][:, 0, :] |
| 41 | |
| 42 | tokens = self.tokenizer(data.label_text, max_length=256, return_tensors='pt', |
| 43 | truncation=True, padding=True).to(self.args.device) |
| 44 | label_embeds = self.textmodel(**tokens)[0][:, 0, :] |
| 45 | |
| 46 | if self.args.if_norm: |
| 47 | node_embeds = (node_embeds - node_embeds.mean(0)) / \ |
| 48 | node_embeds.std(0) |
| 49 | label_embeds = (label_embeds - label_embeds.mean(0) |
| 50 | ) / label_embeds.std(0) |
| 51 | |
| 52 | # change the adj matrix |
| 53 | num_existing_nodes = data.y.shape[0] + 1 |
| 54 | virtual_node_index = data.y.shape[0] |
| 55 | if data.dataset_name in ["Citeseer", "Arxiv"]: |
| 56 | new_edges_to_virtual = [[node_idx, virtual_node_index] |
| 57 | for node_idx in range(num_existing_nodes-1)] |
| 58 | elif data.dataset_name in ["Cora", "Pubmed", "wikics", "home", "tech"]: |
| 59 | new_edges_to_virtual = [] |
| 60 | for node_idx in range(num_existing_nodes-1): |
| 61 | new_edges_to_virtual.append([node_idx, virtual_node_index]) |
| 62 | new_edges_to_virtual.append([virtual_node_index, node_idx]) |
| 63 | new_edge_index = torch.cat([data.edge_index.t(), torch.tensor( |
| 64 | new_edges_to_virtual, dtype=torch.long).to(self.args.device)], dim=0).t() |
| 65 | |
| 66 | adj_normed = self.normalize_adjacency_matrix( |
| 67 | new_edge_index, num_existing_nodes) |
| 68 | for _ in range(self.args.R): |
| 69 | node_embeds = torch.mm(adj_normed, node_embeds) |
| 70 | new_node_embeds = node_embeds[:-1, :] |
| 71 | logits = torch.mm(new_node_embeds, label_embeds.transpose(1, 0)) |
| 72 | logits = torch.div(logits, 1) |
| 73 | # 11*7 -> 10*7 |
| 74 | # 10*1 forever |
| 75 | labels = data.y.long().to(self.args.device) if data.y.dim( |
| 76 | ) == 1 else data.y.squeeze(1).long().to(self.args.device) |
| 77 | CL_loss = self.criteria(logits, labels) |
| 78 | |
| 79 | return CL_loss |
| 80 | |
| 81 | def zero_shot_eval(self, node_embeds, label_embeds, data): |
| 82 | if self.args.if_norm: |
nothing calls this directly
no test coverage detected