(self, data,args)
| 199 | self.criteria = nn.CrossEntropyLoss() |
| 200 | |
| 201 | def forward(self, data,args): |
| 202 | # insert a description node |
| 203 | virtual_node_description = self.descriptions[data.dataset_name] |
| 204 | all_node_texts = data.raw_text + [virtual_node_description] |
| 205 | |
| 206 | tokens = self.tokenizer(all_node_texts, max_length=256, return_tensors='pt', |
| 207 | truncation=True, padding=True).to(self.args.device) |
| 208 | if args.text_encoder == 'llama': |
| 209 | outputs = self.lora_model(**tokens, output_hidden_states=True) |
| 210 | node_embeds = outputs.hidden_states[-1][:, 0, :] |
| 211 | else: |
| 212 | node_embeds = self.lora_model(**tokens)[0][:, 0, :] |
| 213 | |
| 214 | tokens = self.tokenizer(data.label_text, max_length=256, return_tensors='pt', |
| 215 | truncation=True, padding=True).to(self.args.device) |
| 216 | if args.text_encoder == 'llama': |
| 217 | outputs_label = self.lora_model(**tokens, output_hidden_states=True) |
| 218 | label_embeds = outputs_label.hidden_states[-1][:, 0, :] |
| 219 | else: |
| 220 | label_embeds = self.lora_model(**tokens)[0][:, 0, :] |
| 221 | |
| 222 | if self.args.if_norm: |
| 223 | node_embeds = (node_embeds - node_embeds.mean(0)) / \ |
| 224 | node_embeds.std(0) |
| 225 | label_embeds = (label_embeds - label_embeds.mean(0) |
| 226 | ) / label_embeds.std(0) |
| 227 | |
| 228 | # change the adj matrix |
| 229 | num_existing_nodes = data.y.shape[0] + 1 |
| 230 | virtual_node_index = data.y.shape[0] |
| 231 | if data.dataset_name in ["Citeseer", "Arxiv"]: |
| 232 | new_edges_to_virtual = [[node_idx, virtual_node_index] |
| 233 | for node_idx in range(num_existing_nodes-1)] |
| 234 | elif data.dataset_name in ["Cora", "Pubmed", "wikics", "home", "tech","reddit","instagram"]: |
| 235 | new_edges_to_virtual = [] |
| 236 | for node_idx in range(num_existing_nodes-1): |
| 237 | new_edges_to_virtual.append([node_idx, virtual_node_index]) |
| 238 | new_edges_to_virtual.append([virtual_node_index, node_idx]) |
| 239 | new_edge_index = torch.cat([data.edge_index.t(), torch.tensor( |
| 240 | new_edges_to_virtual, dtype=torch.long).to(self.args.device)], dim=0).t() |
| 241 | |
| 242 | adj_normed = self.normalize_adjacency_matrix( |
| 243 | new_edge_index, num_existing_nodes) |
| 244 | for _ in range(self.args.R): |
| 245 | node_embeds = torch.mm(adj_normed, node_embeds) |
| 246 | new_node_embeds = node_embeds[:-1, :] |
| 247 | logits = torch.mm(new_node_embeds, label_embeds.transpose(1, 0)) |
| 248 | logits = torch.div(logits, 1) |
| 249 | # 11*7 -> 10*7 |
| 250 | # 10*1 forever |
| 251 | labels = data.y.long().to(self.args.device) if data.y.dim( |
| 252 | ) == 1 else data.y.squeeze(1).long().to(self.args.device) |
| 253 | CL_loss = self.criteria(logits, labels) |
| 254 | |
| 255 | return CL_loss |
| 256 | |
| 257 | def zero_shot_eval(self, node_embeds, label_embeds, data): |
| 258 | if self.args.if_norm: |
nothing calls this directly
no test coverage detected