-
Notifications
You must be signed in to change notification settings - Fork 12
Expand file tree
/
Copy pathmain_niser_ode.py
More file actions
160 lines (139 loc) · 4.52 KB
/
Copy pathmain_niser_ode.py
File metadata and controls
160 lines (139 loc) · 4.52 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
import argparse
import sys
import torch
import random
import numpy as np
import os
import wandb
def seed_torch(seed=42):
seed = int(seed)
random.seed(seed)
os.environ['PYTHONHASHSEED'] = str(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
torch.backends.cudnn.enabled = True
seed_torch(123)
def get_freer_gpu():
os.system('nvidia-smi -q -d Memory |grep -A4 GPU|grep Free >tmp')
memory_available = [int(x.split()[2]) for x in open('tmp', 'r').readlines()]
# memory_available = memory_available[1:6]
if len(memory_available) == 0:
return -1
return int(np.argmax(memory_available))
os.environ["CUDA_VISIBLE_DEVICES"] = str(get_freer_gpu())
sys.path.append('..')
sys.path.append('../..')
parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)
parser.add_argument(
'--dataset-dir', default='datasets/sample', help='the dataset directory'
)
parser.add_argument(
'--solver', default='rk4', help='The neural ordinary equation solver.'
)
parser.add_argument('--embedding-dim', type=int, default=256, help='the embedding size')
parser.add_argument('--num-layers', type=int, default=1, help='the number of layers')
parser.add_argument(
'--feat-drop', type=float, default=0.1, help='the dropout ratio for features'
)
parser.add_argument('--lr', type=float, default=1e-3, help='the learning rate')
parser.add_argument(
'--batch-size', type=int, default=512, help='the batch size for training'
)
parser.add_argument(
'--epochs', type=int, default=30, help='the number of training epochs'
)
parser.add_argument(
'--weight-decay',
type=float,
default=1e-4,
help='the parameter for L2 regularization',
)
parser.add_argument(
'--patience',
type=int,
default=2,
help='the number of epochs that the performance does not improves after which the training stops',
)
parser.add_argument(
'--num-workers',
type=int,
default=10,
help='the number of processes to load the input graphs',
)
parser.add_argument(
'--valid-split',
type=float,
default=None,
help='the fraction for the validation set',
)
parser.add_argument(
'--log-interval',
type=int,
default=100,
help='print the loss after this number of iterations',
)
args = parser.parse_args()
print(args)
wandb.init(config=vars(args))
from pathlib import Path
import torch as th
from torch.utils.data import DataLoader, SequentialSampler
from src.utils.data.dataset import read_dataset, AugmentedDataset
from src.utils.data.collate import (
seq_to_temporal_session_graph,
collate_fn_factory_temporal,
)
from src.utils.train import TrainRunner
from src.models import NISER_ODE
dataset_dir = Path(args.dataset_dir)
print('reading dataset')
train_sessions, test_sessions, train_timestamps, test_timestamps, num_items = read_dataset(dataset_dir)
if args.valid_split is not None:
num_valid = int(len(train_sessions) * args.valid_split)
test_sessions = train_sessions[-num_valid:]
train_sessions = train_sessions[:-num_valid]
test_timestamps = train_timestamps[-num_valid:]
train_timestamps = train_timestamps[:-num_valid]
dataset = args.dataset_dir.strip().split("/")[-1]
train_set = AugmentedDataset(dataset, train_sessions, train_timestamps)
test_set = AugmentedDataset(dataset, test_sessions, test_timestamps)
collate_fn = collate_fn_factory_temporal(seq_to_temporal_session_graph)
train_loader = DataLoader(
train_set,
batch_size=args.batch_size,
shuffle=True,
# drop_last=True,
num_workers=args.num_workers,
collate_fn=collate_fn,
pin_memory=True,
# sampler=SequentialSampler(train_set)
)
test_loader = DataLoader(
test_set,
batch_size=args.batch_size,
shuffle=True,
num_workers=args.num_workers,
collate_fn=collate_fn,
)
model = NISER_ODE(num_items, args.embedding_dim, args.num_layers, feat_drop=args.feat_drop, solver=args.solver)
device = th.device('cuda:0' if th.cuda.is_available() else 'cpu')
model = model.to(device)
print(model)
runner = TrainRunner(
args.dataset_dir,
model,
train_loader,
test_loader,
device=device,
lr=args.lr,
weight_decay=args.weight_decay,
patience=args.patience,
)
print('start training')
mrr10, mrr20, hit10, hit20 = runner.train(args.epochs, args.log_interval)
print('MRR@20\tHR@20')
print(f'{mrr10 * 100:.3f}%\t{mrr20 * 100:.3f}%\t{hit10 * 100:.3f}%\t{hit20 * 100:.3f}%')