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

Method zero_shot_eval

code/st_model.py:81–115  ·  view source on GitHub ↗
(self, node_embeds, label_embeds, data)

Source from the content-addressed store, hash-verified

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

Callers 1

evalFunction · 0.45

Calls 1

Tested by

no test coverage detected