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

Method zero_shot_eval

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

Source from the content-addressed store, hash-verified

255 return CL_loss
256
257 def zero_shot_eval(self, node_embeds, label_embeds, data):
258 if self.args.if_norm:
259 node_embeds = (node_embeds - node_embeds.mean(0)) / \
260 node_embeds.std(0)
261 label_embeds = (label_embeds - label_embeds.mean(0)
262 ) / label_embeds.std(0)
263
264 # change the adj matrix
265 num_existing_nodes = data.y.shape[0] + 1
266 virtual_node_index = data.y.shape[0]
267 if self.args.test_data in ["Citesee"]:
268 new_edges_to_virtual = [[node_idx, virtual_node_index]
269 for node_idx in range(num_existing_nodes-1)]
270 elif self.args.test_data in ["Cora", "Pubmed", "Citeseer", "Arxiv", "wikics", "facebook", 'home', 'tech','reddit','instagram']:
271 new_edges_to_virtual = []
272 for node_idx in range(num_existing_nodes-1):
273 new_edges_to_virtual.append([node_idx, virtual_node_index])
274 new_edges_to_virtual.append([virtual_node_index, node_idx])
275 new_edge_index = torch.cat([data.edge_index.t(), torch.tensor(
276 new_edges_to_virtual, dtype=torch.long).to(self.args.device)], dim=0).t()
277 adj_normed = self.normalize_adjacency_matrix(
278 new_edge_index, num_existing_nodes)
279
280 # adj_normed = self.normalize_adjacency_matrix(data)
281 for _ in range(self.args.R):
282 node_embeds = torch.mm(adj_normed, node_embeds)
283 node_embeds = node_embeds[:-1, :]
284 node_embeds /= node_embeds.norm(dim=-1,
285 keepdim=True).to(self.args.device)
286 label_embeds /= label_embeds.norm(dim=-1, keepdim=True)
287 dists = torch.einsum('bn,cn->bc', node_embeds, label_embeds)
288 preds = torch.argmax(dists, dim=1)
289 labels = data.y.long().to(self.args.device)
290 if len(data.test_mask) == 10 :
291 data.test_mask = data.test_mask[0]
292 test_mask = data.test_mask
293 test_acc = accuracy_score(labels[test_mask].cpu(), preds[test_mask].cpu())
294 test_f1 = f1_score(labels[test_mask].cpu(), preds[test_mask].cpu())
295 return [test_acc,test_f1]
296
297 def normalize_adjacency_matrix(self, edge_index, num_nodes):
298 # edge_index = data.edge_index

Callers

nothing calls this directly

Calls 1

Tested by

no test coverage detected