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

Method forward

code/st_model.py:33–79  ·  view source on GitHub ↗
(self, data)

Source from the content-addressed store, hash-verified

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:

Callers

nothing calls this directly

Calls 1

Tested by

no test coverage detected