MCPcopy Create free account
hub / github.com/NineAbyss/ZeroG / forward

Method forward

code/st_model.py:201–255  ·  view source on GitHub ↗
(self, data,args)

Source from the content-addressed store, hash-verified

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:

Callers

nothing calls this directly

Calls 1

Tested by

no test coverage detected