Wrapper functions for node level data loader
(root, name, args)
| 572 | # ====================================================================== |
| 573 | |
| 574 | def load_node_dataset(root, name, args): |
| 575 | """ |
| 576 | Wrapper functions for node level data loader |
| 577 | """ |
| 578 | |
| 579 | dataset = None |
| 580 | if name in ['yelp', 'amazon']: |
| 581 | dataset = FraudDataset(root, name) |
| 582 | elif name in ['weibo', ]: |
| 583 | dataset = TextDataset(root, name) |
| 584 | elif name in ['tfinance', 'tsocial']: |
| 585 | dataset = TDataset(root, name) |
| 586 | elif name == 'elliptic': |
| 587 | dataset = EllipticBitcoinDataset(osp.join(root, name)) |
| 588 | timestep = pd.read_csv(dataset.raw_paths[0], header=None).iloc[:, 1] |
| 589 | timestep = torch.tensor(timestep, dtype=dataset.x.dtype) |
| 590 | dataset.x = torch.concat( |
| 591 | [timestep.unsqueeze(dim=0).T, dataset.x], dim=1) |
| 592 | elif name == 'dgraphfin': |
| 593 | dataset = DGraphFin(osp.join(root, name)) |
| 594 | elif name in ['questions', 'tolokers']: |
| 595 | dataset = HeterophilousGraphDataset(root, name.capitalize()) |
| 596 | elif name in ['Arxiv', 'Cora', 'Pubmed', 'Citeseer', 'wikics','reddit','instagram']: |
| 597 | |
| 598 | # dataset = CitationDataset(root, name, args) |
| 599 | |
| 600 | data = torch.load(f"datasets/{name.lower()}.pt") |
| 601 | data.label_text = data.label_name |
| 602 | dataset = DataWrapper(data, args) |
| 603 | if name in ['Cora', 'Pubmed', 'Citeseer']: |
| 604 | dataset.test_masks = dataset.data.test_mask[0].unsqueeze(1) |
| 605 | |
| 606 | return dataset |
| 607 | |
| 608 | |
| 609 |
no test coverage detected