(self, root, name, args, transform=None, pre_transform=None)
| 462 | "Pubmed": "Pubmed.pt", "Citeseer": "Citeseer.pt"} |
| 463 | |
| 464 | def __init__(self, root, name, args, transform=None, pre_transform=None): |
| 465 | |
| 466 | self.name = name |
| 467 | assert self.name in ['Cora', 'Pubmed', 'Citeseer'] |
| 468 | |
| 469 | self.args = args |
| 470 | super().__init__(root, transform, pre_transform) |
| 471 | # if self.name == 'Citeseer': |
| 472 | # self.data = torch.load(osp.join(self.processed_dir, 'data.pt')) |
| 473 | # else: |
| 474 | |
| 475 | self.load(self.processed_paths[0]) |
| 476 | |
| 477 | if args.if_text: |
| 478 | if not osp.exists(osp.join(self.processed_dir, 'data_text.pt')): |
| 479 | self.text_emb() |
| 480 | torch.save(self.data, osp.join( |
| 481 | self.processed_dir, 'data_text.pt')) |
| 482 | else: |
| 483 | self.data = torch.load( |
| 484 | osp.join(self.processed_dir, 'data_text.pt')) |
| 485 | if self.name == 'Citeseer': |
| 486 | self.data.edge_index = torch.load( |
| 487 | "datasets/Citeseer/raw/data.pt").edge_index |
| 488 | self.label_emb_func(root) |
| 489 | |
| 490 | def label_emb_func(self, root): |
| 491 | if self.name == "Cora": |
nothing calls this directly
no test coverage detected