项目文件夹

文件
Quan (Andy) Gan e667545da5 [Feature] Node2vec (#2992)
* add seal example

* 1. add paper infomation in examples/README
2. adjust codes
3. option test

* use latest `to_simple` to replace coalesce graph function

* remove outdated codes

* remove useless comment

* Node2vec
1.implement node2vec random walk c++ op
2.implement node2vec model
3.implement node2vec example

* add CMakeLists file modify

* refine c++ codes

* refine c++ codes

* add missing whitespace

* refine python codes

* add codes

* add node2vec_impl.h

* fix codes

* fix code style problem

* fixes

* remove

* lots of changes

* add benchmark

* fixes

Co-authored-by: smilexuhc <smile.xuhc@gmail.com>
Co-authored-by: Minjie Wang <wmjlyjemaine@gmail.com>
2021-06-23 14:47:56 +08:00

64 行
2.0 KiB
Python

import argparse
from dgl.data import CitationGraphDataset
from ogb.nodeproppred import *
from ogb.linkproppred import *
def load_graph(name):
cite_graphs = ['cora', 'citeseer', 'pubmed']
if name in cite_graphs:
dataset = CitationGraphDataset(name)
graph = dataset[0]
nodes = graph.nodes()
y = graph.ndata['label']
train_mask = graph.ndata['train_mask']
val_mask = graph.ndata['test_mask']
nodes_train, y_train = nodes[train_mask], y[train_mask]
nodes_val, y_val = nodes[val_mask], y[val_mask]
eval_set = [(nodes_train, y_train), (nodes_val, y_val)]
elif name.startswith('ogbn'):
dataset = DglNodePropPredDataset(name)
graph, y = dataset[0]
split_nodes = dataset.get_idx_split()
nodes = graph.nodes()
train_idx = split_nodes['train']
val_idx = split_nodes['valid']
nodes_train, y_train = nodes[train_idx], y[train_idx]
nodes_val, y_val = nodes[val_idx], y[val_idx]
eval_set = [(nodes_train, y_train), (nodes_val, y_val)]
else:
raise ValueError("Dataset name error!")
return graph, eval_set
def parse_arguments():
"""
Parse arguments
"""
parser = argparse.ArgumentParser(description='Node2vec')
parser.add_argument('--dataset', type=str, default='cora')
# 'train' for training node2vec model, 'time' for testing speed of random walk
parser.add_argument('--task', type=str, default='train')
parser.add_argument('--runs', type=int, default=10)
parser.add_argument('--device', type=str, default='cpu')
parser.add_argument('--embedding_dim', type=int, default=128)
parser.add_argument('--walk_length', type=int, default=50)
parser.add_argument('--p', type=float, default=0.25)
parser.add_argument('--q', type=float, default=4.0)
parser.add_argument('--num_walks', type=int, default=10)
parser.add_argument('--epochs', type=int, default=100)
parser.add_argument('--batch_size', type=int, default=128)
args = parser.parse_args()
return args