From 71ca6baa8f104a9e47cf057f27e4271eb852fa09 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Sat, 8 Jan 2022 22:43:44 +0800 Subject: [PATCH 01/57] ads --- src/utils/data/preprocess.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/src/utils/data/preprocess.py b/src/utils/data/preprocess.py index fcf9efd..7f9a3dd 100644 --- a/src/utils/data/preprocess.py +++ b/src/utils/data/preprocess.py @@ -103,6 +103,7 @@ def train_test_split(df, test_split=0.2): def save_sessions(df, filepath): df = reorder_sessions_by_endtime(df) + print(df.heads()) sessions = df.groupby('sessionId').itemId.apply(lambda x: ','.join(map(str, x))) sessions.to_csv(filepath, sep='\t', header=False, index=False) @@ -117,15 +118,15 @@ def save_dataset(dataset_dir, df_train, df_test): # update itemId train_itemId_new, uniques = pd.factorize(df_train.itemId) - df_train = df_train.assign(itemId=train_itemId_new) - oid2nid = {oid: i for i, oid in enumerate(uniques)} - test_itemId_new = df_test.itemId.map(oid2nid) - df_test = df_test.assign(itemId=test_itemId_new) + df_train = df_train.assign(itemId=train_itemId_new) + oid2nid = {oid: i for i, oid in enumerate(uniques)} + test_itemId_new = df_test.itemId.map(oid2nid) + df_test = df_test.assign(itemId=test_itemId_new) print(f'saving dataset to {dataset_dir}') dataset_dir.mkdir(parents=True, exist_ok=True) save_sessions(df_train, dataset_dir / 'train.txt') - save_sessions(df_test, dataset_dir / 'test.txt') + save_sessions(df_test, dataset_dir / 'test.txt') num_items = len(uniques) with open(dataset_dir / 'num_items.txt', 'w') as f: f.write(str(num_items)) @@ -154,7 +155,6 @@ def preprocess_diginetica(dataset_dir, csv_file): df_train, df_test = split_by_time(df, pd.Timedelta(days=7)) save_dataset(dataset_dir, df_train, df_test) - def preprocess_gowalla_lastfm(dataset_dir, csv_file, usecols, interval, n): print(f'reading {csv_file}...') df = pd.read_csv( From a7d46d76a72daf8df081bc9ed85175418e49e08f Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Mon, 10 Jan 2022 01:22:53 +0800 Subject: [PATCH 02/57] first --- src/models/__init__.py | 3 +- src/models/niser_ode.py | 218 ++++++++++++++++++ src/preprocess.py | 52 +++++ src/scripts/main_niser_ode.py | 122 ++++++++++ src/utils/__pycache__/train.cpython-38.pyc | Bin 3542 -> 3642 bytes .../data/__pycache__/collate.cpython-38.pyc | Bin 9036 -> 11224 bytes .../data/__pycache__/dataset.cpython-38.pyc | Bin 2199 -> 2908 bytes .../__pycache__/preprocess.cpython-38.pyc | Bin 5645 -> 6269 bytes src/utils/data/collate.py | 47 ++++ src/utils/data/dataset.py | 27 ++- src/utils/data/preprocess.py | 17 +- src/utils/train.py | 15 +- 12 files changed, 482 insertions(+), 19 deletions(-) create mode 100644 src/models/niser_ode.py create mode 100644 src/preprocess.py create mode 100644 src/scripts/main_niser_ode.py diff --git a/src/models/__init__.py b/src/models/__init__.py index 836757f..cb519cb 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -1,4 +1,5 @@ from .lessr import LESSR from .msgifsr import MSGIFSR from .niser import NISER -from .srgnn import SRGNN \ No newline at end of file +from .srgnn import SRGNN +from .niser_ode import NISER_ODE \ No newline at end of file diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py new file mode 100644 index 0000000..ccfb02c --- /dev/null +++ b/src/models/niser_ode.py @@ -0,0 +1,218 @@ +import math + +import torch as th +import torch.nn as nn +import torch.nn.functional as F + +import dgl +import dgl.ops as F +import dgl.function as fn + +import torchcde + +class CDEFunc(nn.Module): + def __init__(self, input_channels, hidden_channels): + ###################### + # input_channels is the number of input channels in the data X. (Determined by the data.) + # hidden_channels is the number of channels for z_t. (Determined by you!) + ###################### + super(CDEFunc, self).__init__() + self.input_channels = input_channels + self.hidden_channels = hidden_channels + + self.linear1 = nn.Linear(hidden_channels, 128) + self.linear2 = nn.Linear(128, input_channels * hidden_channels) + + ###################### + # For most purposes the t argument can probably be ignored; unless you want your CDE to behave differently at + # different times, which would be unusual. But it's there if you need it! + ###################### + def forward(self, t, z): + # z has shape (batch, hidden_channels) + z = self.linear1(z) + z = z.relu() + z = self.linear2(z) + ###################### + # Easy-to-forget gotcha: Best results tend to be obtained by adding a final tanh nonlinearity. + ###################### + z = z.tanh() + ###################### + # Ignoring the batch dimension, the shape of the output tensor must be a matrix, + # because we need it to represent a linear map from R^input_channels to R^hidden_channels. + ###################### + z = z.view(z.size(0), self.hidden_channels, self.input_channels) + return z + +class SRGNNLayer(nn.Module): + def __init__(self, input_dim, output_dim, batch_norm=False, feat_drop=0.0, activation=None): + super().__init__() + self.dropout = nn.Dropout(feat_drop) + self.gru = nn.GRUCell(2 * input_dim, output_dim) + self.W1 = nn.Linear(input_dim, output_dim, bias=False) + self.W2 = nn.Linear(input_dim, output_dim, bias=False) + self.activation = activation + + def messager(self, edges): + + return {'m': edges.src['ft'] * edges.data['w'].unsqueeze(-1), 'w': edges.data['w']} + + def reducer(self, nodes): + + m = nodes.mailbox['m'] + w = nodes.mailbox['w'] + hn = m.sum(dim=1) / w.sum(dim=1).unsqueeze(-1) + + return {'neigh': hn} + + def forward(self, mg, feat): + with mg.local_scope(): + mg.ndata['ft'] = self.dropout(feat) + if mg.number_of_edges() > 0: + mg.update_all(self.messager, self.reducer) + neigh1 = mg.ndata['neigh'] + mg1 = mg.reverse(copy_edata=True) + mg1.update_all(self.messager, self.reducer) + neigh2 = mg1.ndata['neigh'] + neigh1 = self.W1(neigh1) + neigh2 = self.W2(neigh2) + hn = th.cat((neigh1, neigh2), dim=1) + rst = self.gru(hn, feat) + else: + rst = feat + if self.activation is not None: + rst = self.activation(rst) + return rst + +class AttnReadout(nn.Module): + def __init__( + self, + input_dim, + hidden_dim, + output_dim, + batch_norm=True, + feat_drop=0.0, + activation=None, + ): + super().__init__() + self.batch_norm = nn.BatchNorm1d(input_dim) if batch_norm else None + self.feat_drop = nn.Dropout(feat_drop) + self.fc_u = nn.Linear(input_dim, hidden_dim, bias=False) + self.fc_v = nn.Linear(input_dim, hidden_dim, bias=True) + self.fc_e = nn.Linear(hidden_dim, 1, bias=False) + self.fc_out = ( + nn.Linear(input_dim, output_dim, bias=False) + if output_dim != input_dim + else None + ) + self.activation = activation + + def forward(self, g, feat, last_nodes): + if self.batch_norm is not None: + feat = self.batch_norm(feat) + feat = self.feat_drop(feat) + feat_u = self.fc_u(feat) + feat_v = self.fc_v(feat[last_nodes]) + feat_v = dgl.broadcast_nodes(g, feat_v) + e = self.fc_e(th.sigmoid(feat_u + feat_v)) + alpha = F.segment.segment_softmax(g.batch_num_nodes(), e) + feat_norm = feat * alpha + rst = F.segment.segment_reduce(g.batch_num_nodes(), feat_norm, 'sum') + if self.fc_out is not None: + rst = self.fc_out(rst) + if self.activation is not None: + rst = self.activation(rst) + return rst + +class NISER_ODE(nn.Module): + + def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=True, scale=12): + super().__init__() + self.embedding = nn.Embedding(num_items, embedding_dim) + self.register_buffer('indices', th.arange(num_items, dtype=th.long)) + self.embedding_dim = embedding_dim + self.num_layers = num_layers + self.layers = nn.ModuleList() + self.norm = norm + self.scale = scale + input_dim = embedding_dim + for i in range(num_layers): + layer = SRGNNLayer( + input_dim, + embedding_dim, + batch_norm=None, + feat_drop=feat_drop + ) + self.layers.append(layer) + self.readout = AttnReadout( + input_dim, + embedding_dim, + embedding_dim, + batch_norm=None, + feat_drop=feat_drop, + activation=None, + ) + + self.feat_drop = nn.Dropout(feat_drop) + self.fc_sr = nn.Linear(input_dim + embedding_dim, embedding_dim, bias=False) + + self.reduce = nn.Linear(embedding_dim, 32) + self.recover = nn.Linear(32, embedding_dim) + self.cde_func = CDEFunc(33, 32) + self.initial = nn.Linear(33, 32) + + self.reset_parameters() + + def reset_parameters(self): + stdv = 1.0 / math.sqrt(self.embedding_dim) + for weight in self.parameters(): + weight.data.uniform_(-stdv, stdv) + + def forward(self, mg, embeds_ids, times, num_nodes): + iid = mg.ndata['iid'] + + feat = self.feat_drop(self.embedding(iid)) + if self.norm: + feat = feat.div(th.norm(feat, p=2, dim=-1, keepdim=True) + 1e-12) + out = feat + for i, layer in enumerate(self.layers): + out = layer(mg, out) + + feat_ode = self.reduce(self.feat_drop(self.embedding(embeds_ids))) + if self.norm: + feat_ode = feat_ode.div(th.norm(feat_ode, p=2, dim=-1, keepdim=True) + 1e-12) + + feat_ode = th.cat([times.unsqueeze(-1), feat_ode], dim=-1) + X = torchcde.LinearInterpolation(feat_ode) + X0 = X.evaluate(X.interval[0]) + z0 = self.initial(X0) + # print(X.interval) + z_T = torchcde.cdeint(X=X, + z0=z0, + func=self.cde_func, + t=X.interval) + + time_embeds = self.recover(z_T[:, 1]) + + out += time_embeds + + last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) + if self.norm: + feat = feat.div(th.norm(feat, p=2, dim=-1, keepdim=True)) + sr_g = self.readout(mg, feat, last_nodes) + sr_l = feat[last_nodes] + sr = th.cat([sr_l, sr_g], dim=1) + sr = self.fc_sr(sr) + if self.norm: + sr = sr.div(th.norm(sr, p=2, dim=-1, keepdim=True) + 1e-12) + target = self.embedding(self.indices) + if self.norm: + target = target.div(th.norm(target, p=2, dim=-1, keepdim=True) + 1e-12) + logits = sr @ target.t() + if self.scale: + logits = th.log(nn.functional.softmax(self.scale * logits, dim=-1)) + else: + logits = th.log(nn.functional.softmax(logits, dim=-1)) + return logits# , 0 + + + \ No newline at end of file diff --git a/src/preprocess.py b/src/preprocess.py new file mode 100644 index 0000000..9f29202 --- /dev/null +++ b/src/preprocess.py @@ -0,0 +1,52 @@ +from pathlib import Path +import argparse +import time + +parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter) + +optional = parser._action_groups.pop() +required = parser.add_argument_group('required arguments') +required.add_argument( + '-d', + '--dataset', + choices=['diginetica', 'gowalla', 'lastfm'], + required=True, + help='the dataset name', +) +required.add_argument( + '-f', + '--filepath', + required=True, + help='the file for the dataset, i.e., "train-item-views.csv" for diginetica, ' + '"loc-gowalla_totalCheckins.txt" for gowalla, ' + '"userid-timestamp-artid-artname-traid-traname.tsv" for lastfm', +) +optional.add_argument( + '-t', + '--dataset-dir', + default='datasets/{dataset}', + help='the folder to save the preprocessed dataset', +) +parser._action_groups.append(optional) +args = parser.parse_args() + +dataset_dir = Path(args.dataset_dir.format(dataset=args.dataset)) + +if args.dataset == 'diginetica': + from utils.data.preprocess import preprocess_diginetica + + preprocess_diginetica(dataset_dir, args.filepath) +else: + from pandas import Timedelta + from utils.data.preprocess import preprocess_gowalla_lastfm + + csv_file = args.filepath + if args.dataset == 'gowalla': + usecols = [0, 1, 4] + interval = Timedelta(days=1) + n = 30000 + else: + usecols = [0, 1, 2] + interval = Timedelta(hours=8) + n = 40000 + preprocess_gowalla_lastfm(dataset_dir, csv_file, usecols, interval, n) \ No newline at end of file diff --git a/src/scripts/main_niser_ode.py b/src/scripts/main_niser_ode.py new file mode 100644 index 0000000..255cc81 --- /dev/null +++ b/src/scripts/main_niser_ode.py @@ -0,0 +1,122 @@ +import argparse +import sys + +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('--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=0, + 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) + + +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] + +train_set = AugmentedDataset(train_sessions, train_timestamps) +test_set = AugmentedDataset(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) +device = th.device('cuda' 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') +mrr, hit = runner.train(args.epochs, args.log_interval) +print('MRR@20\tHR@20') +print(f'{mrr * 100:.3f}%\t{hit * 100:.3f}%') diff --git a/src/utils/__pycache__/train.cpython-38.pyc b/src/utils/__pycache__/train.cpython-38.pyc index fd7111afa62afa8d6db32843be6c107bd4ad608b..415b8f5eae0684341bdf6907b685daca2dc54fa8 100644 GIT binary patch delta 1148 zcmY*Z&2Jk;6rVRcv)*ri#C4L!fJ2g|4iqYJu#^(b5GG#vJX0^oQ!-y z`k+0<#uih2Y;om|Em!3fdMISt+?;cM0N?!ORG%0Gt?er<(OcgVy#Y_Rwp+b{-~DKN zy}Kh?CS)fX2_``xp(O}LFO`eP{82{OV0Ll*6mhocpTW_4Xl}44Z+!j1+1`;d3}PhoMFKQv0a|4!4d((DUx5hbseB22Hoh)So15;RaN1nW{XG8)QgYJ5 zn-3u;Q(aP}Y+@W6WaXiiL+((YC^}FXQ0p?Sm7pvA*v_b@2qX!_X<{N)=Zc_NjPv0XVP@V|OL`3(Ph?W1?PR#LO@zmjA z%zYpSaQ$HyPfqi6Mmfsc$BspdDIYDS2*cfk*C3RvPhVd~57wq@q@Fi-f@N4TzXz3P z)=;E9=eZZ!d;Y?+wO6;&+KaV+V4+?eFx!?x)khMaH?^=jGZfk$ntTmm;P18M&ibZo z3VF<23|~LmuuHRciSNp-ZYtlf#9J0FSk~EY-l$K=x6PyQY<=0%D~}y)n0(s8oQ3~3 zOe2~-@;;Imar9OZ$f#Lt)6fcxZb=23rEETlrhu7mqk8ESx=FBb8lety&HNfI&Anq8 zyP0(V5B9GhPa>?C*?3~fS{QWHPI+nn3ImRJbob_fc6agC?#PX+@)**W%!l#()Fu9L Xvn#I1!^jNp>PWkj=1#og#6kIA17`G5 delta 1078 zcmYjQ&u`pB6rML8d+qU$-K6P;WVa=OjiOBp2(=>65SofYX+vpI+Z#}pOcSypn{?Kw z5Iw6vkU$7f(Hv8Rec;klIkoBmAt8F;zyS#k`v+9w(koKoz#9`GY`xFleDmhbo98!o zn!h)a3vnzNJqu6#*tuW*E&*rGqoK2F{ewi1y~Z=1IiIj)P@wq7KJw#i8!~}T&OT)Y zWG)JHm$^xuZ^1L;qG`^HxhH*E%)sZvAh~W4|6_z;NmzOJN$4ApJx|vE-rN=YuYL{; z;W&H!gRGJ|SV6wcX@5TpkfHx2%OuKND4<bBSD|Hoa!){FV($gh@%F-F=888> z>Nj5fp@dE(!G{`1s6zxTuAimsm%bVrGxxmvaMb)(xz(}<(xVnmSy;62u!RK+Ll#b( zoBkv`V}A8dLCx$7I#++r!Zg4TSyz;Yjwn^ul{PIZ3hQ!=r~pYC9jq3!YW6n(x_NA@h-vRY==gFysw;=nBQED_9fPKu4A%uQLVS8n=SkBUcWM z<>MHjuY&7TIbi9ilH}OfVJajw8euBl5h}spKVCIo%jsjyGR6&H*WmImJ^D|Xo9;sg z{cY+XVgdtV0-2y_%AaMNU2I%FLLOq7pJVNsal<3=G}Wig_0zUYvt}WzA2>^N%X$3J z@>Vc)aN?!abYgzu0Zg^`4@A1UeyRKg|JHlad>A$xWz}uZ=@$tG!PX_cadAa=l%6rS z!^Ojg?bM{5ayRvAFV!bYex=h>Zw;JxHdge@*5kBQo%POp?V5hY9EeVg&sw=mtaPv; z^dlA?H9OJdu?s}cllASk3`F_*Ab1_<8qjOoMiQqsqH9s(O_(7jc-c-gHh)F!;pb?r z1`9_C+Mv&xr{bA6&sfFopzS67AJ{vr=Lu|LeLM1fY8haE%0{VEh$&LrReg>=1Go3? r`sVus-rF3w>4si{(B>2q5KCs?D^e~o@bL8Q~c)P005|OHg86*;w zs9Xt|A@L5IVeyWb5qK+=NHuCk!A7Nlb;*oTjq333p#}}YyO)M&7~Xv}LZk4G(-`f6 zcR%f=eeh1uaoSH4>%<(O19T81gLH@{*V~dgL{l^kf~1|I!?)zG5|*JE8iyxKN9ZWD zMraPP6gJcJ!1fm{guE0qhwY?20KJ0`WHV!@X}l!c!(cC$WHW1z*cm&!MyQm61E!>n zQ98Cs)(9L$o{{AIQ(F^cjF9lwYjO`Eef%xu9`X1em6`q|(A6}MG?E|&=_M??byaQ>gj~LWe|8yKR?C^G=t{ zN3u_*_&EiSi5(v8X0u$c`7};T{14HI_D5m1q71O}2}2_Qj{;9qgJU_;0&yiusQf^_ zA-Cj~LKUj6OOEWyYviJ|M!qOH3btG7J!y{60M*tFR}vcHsP_oSbd)&(umgK)Bmq&s z)+Vih8`#h`rIrRt4V1xKD*KwFyV{o}SWa&ljH5`2w~inK?njG zg?_&qx-Y?*V=0LBHJO0vjFj)WdM!`92(z6fRx6Zi)L!Ag)6bpSFEmZTAb=+PWcBFw z&*0eAMx9v|{_Akz#D0Ew69@0O*pr5)w&Quaz|MfJAGl|+^i`W&8D_mH3Z&-M*7jVi|h; z+hP=w>kR+f;28N5Sp49W(8>(c7keEo6*hxyOvUUf03@7GV!N{xyAIt4{9i-%*=wkH z69Hwu2_C#_jLEh-Tdvt(0C7_TFo>ZT3cJC-le|2FY0Xp0bcMaPXBfj=|6LwTy)yRS zed&i_m*r>DLH_O3OYyz-aGV?bFR5$*C;SnQr|si=4IZiN1}^XxwE57x)!CofrYm12yMP!xny=&v;6;Tw?<`YFDwsZUVs5tVYAGQHH6H?ecEy zosAhFhi-aSaL5tEttxmzQC;9us;~C*B_-mT*dsr9C zMaR?XjN0rdo8pPg>*fnM+pB0067v*9`hP-te=H#Ky6+u=Ls2rjM2ba^@}FhWWQl)} z`5u|)?_@Wh6CP)`UF8j9pp9YJMAs<6-fu=)#3dWj zru0x=RA+nw@^GmHJrIL&ByKxCF?+o-0yVSam8%zu^-9HZ?802FFlQB=I=fTYDFBma zEA^sPX--cT>--nFY#Z}^5pNPZAN0;3Jc3xv(_(WcT0|6OvLK5KDY=gYc)Pbok`h$D z9i0ppR^@@>A0w}_98BV0Rxn#Wf}mNpDogew)gzatJf&e#KaY#Rxh+ggy=tu#YD?9} z!Y@X6D3xjdKl_8-?Ao8$8|+W*>@UDR-q~LcTewNKg0PB!8)7Yl4(Dx=yn1$%o{x0h zot4pAc(({^jA#sj6;E4kI#%(PU&R*CN(2H1I$UsXp`e2B34|&_4WW+EKtQiIgKfJy zYtv?-OxYqz8G_g%)-4e}I&LlrJ&DAKD*gqe0h90zkqA7iNB&1I*>NTR_)ECJX+Jix zi3@teA2JnQ=XPQ@#@^mKGx{PKMQ8R%0SXUhtS+&F*Bs&B$j!7n*~6z3!^BUm&tgmP zCN2YlMn9E+0)B@sja>!>c_kbuxD;dXy5q^S^L|tzkq~BLKTT-&?b;RQqa@+02R<`^ zD&MG^q}KVL4&-`op|7wY!e=WsvBgQqph#ak6yOgC#W1!VkN z2g4V#XmK3jGy-0}*=2-zgk^-c5$+;<8UeGqslw%F%Fo0V{@a7oC-C|o0T(jx50Lmq g6hkv~c>0VUBW+|uZ7Erp;IHR@0Q1ryRsaA1 delta 2077 zcmZ`)&2Jk;6yI6f>-AUs5j&0@eeo~`Z1iIcJ0 zvY5Pz~zyeG1ex;YB`J6Hh?oXAENDuF_CrO;2u&2o=&)RFf(HQ)*DNfJcH@(M1O;o4t%Mj?g4;Dw{^BP^`5SL2eWw!=HE( zn4CJK|(|)zWHgk+Z=Vy7SD^1D}@zuCkm01TS{{}`{*eu$p6Lt;&iu^j- z%~!FDFs<-J%$U4@3vVHyuMslWv`~Vk9r9W(trP>3goA$*KNC0@4|Eo|FA=xlfuHb^gfVt7 z;J%eD;F0elEFxqO_Sxdb&-u5B($>L^?DGHP9)PQA1;DEVHMgzRZmkcQeo6!!)%C;rQXrU@U+)3z$wnjiYxFO ziL+NRa{rRs8$uD~6xV&frh6Cv=Twv|@;_5wkSlz*dv{QXPi4pP9ApP2%B>3~rfoS_ z%NN-aE{PiS@k&pT{Ky~o%(R1sS8$<>AX;^eXVNk64=@+?mx$ZRR_k{TpCn14F_Hxx zBkTWk_E(8NJu<)5%z))E1{>1ol!lf_LVOgKaoR`+EX~dG=%5G&!c~MKLKdNbAog4W z%aoUk;(N(TFuN}|@2A}Lpd7nVg9?`ZiV|=A=V5$N> zQ?8CcaTGA$)S{*ei~Uk!iGma@s$Y833;ve$xK^cUx3NwZr3Ay=v6CCD~Mu8dc>e K>9DxrhwmSl>W1zB diff --git a/src/utils/data/__pycache__/dataset.cpython-38.pyc b/src/utils/data/__pycache__/dataset.cpython-38.pyc index df4e821b43d478091c494b663386ae925e022bf9..c19cb1a2ee478c09f7196f90dbce53ef48082afe 100644 GIT binary patch literal 2908 zcmbVOTW=gS6t+E=o!xBGkRsHyi2Fq&Xtq2cfhbWC6-XdO&~jaYMw7KS-7q^lwY{iW zO`g&Uc&HHnAbHGh;UAbsp8CuK5-;!_Pm*TS3L$1?`}p{F{GH?P?B|P%9ft4E)^GW* z31h!g=kyn#a}C9Ofl4yT3)ZJr6nrfDf;YOLANsMVAIm_7cUiwBBN?MjWJ@M!+p;Y? zXgeP=xgfiDne6U~{(|H;S-SWDd$Wy{yHKlK4mfp>z6foM@;-_wP+7Lkt*}AO_5|cm z^7~?!Tecse#9q%2c+I7_zz)Q&u%fZ3=6hl{u&wQ6zg-6h95QyQsc!awwBM;k9Sj9V z!#adz5g8RpyHlwLWYwH$q#$^1q8LKa0v<983&n;Gp`q0}FSir~UK}#rg7UAgd!HMn zO>b*^Cog-~Cxhu&mDcn=(UWa8fKVu7bnhnW(WJbg2A8WHJJEw%J);M`wKd;N?Wiz4 zncLh~M%lG$$0b*be7q_1b+YFV6gGQi08p*UMx|7D)~*(lL0*`32+hQF!HKeR$xu(m zBdfFv#(Cv}LX|EoCgqlkwW_dzi*tVy*QQx^GAWFU2DkE2=@w~iM(TXCP_8?e6w`5; z8JiEjOCy(<<{}0gs?ziqj{4b1-s!iCd{Y%9x&W02sHqT*411k2p1?CpbSeH6zsQ$( zmoKB%tI*7OMnG{AG_RoJT@CM0;L6~BxQhVU7c!L5U3S2?1M0(T@y1`g{n=sQ0;4Jy z-$F0~ke0bURcfZ@jv(jL5#ISv`XemjvS#~%MEvcEJ#2U=u;J@zh>MSm)r**hTb9;E zrYc7E5J7SF#=4D3=TS6yaP}-ci0l}^N4<9#I#EfN_-O<4y^HdyDt27_eO^qJITUd0 zs8CgIZ|SFKSHkcVzKe7FKY>9*w_wSVE$wKm5R>D|Jc^^HDABJyx|2J1W-Xz38lxWi za~P2&tg_nf|S;)Z4ftN!shZIFJy2NC>Cw^YpjlgNqet1CZ z<1g8p>{E7weZ^`{Tvdm4Jd6;!t&LgRYP`lZd&kZai<^4syjNx?clXK_8a|_2ET{er zj}1+6T&I!uzorYia8jvKzd)jowPIANwB;6ve?UE%9BH@I9CaM5Zqc7|-0hz`V@ZE$ z-U`Quu6<&}8Hc*Acf-+f6h%iw&;<$5ZyA)J--3RCGWQWhUqdnEV#WF(fxbXC0ukKh zePk;l0pzb(w$Mg00h%4R{N8knSXs&I%@scw5TPG`^hhW*c_)7a0q$TKgG*ROG{E7k zunzP!8|)*MfOptEA^A|Of^VoN#l2va?M855gcftKGLqnvG-z_igTL{mv46I%bK*eb z=}s@;OxMcwmYHbqi6{3ocA-(luzBh4U`Bt5hO^C`49vR)4tSTAfr8{&e3J|bVq~-I zWz0CjHs-NNL@1u)37@T=1!V1b#ei7h_@n42#HkpH{+ewAoK^M>a`7&*HgYaNkgN{w zp~c0IG<5?9+ZJ={ac{eucKiA$ifXxL^Agd$trRgk*Ogr z*n%zJE)&)N$S%Nwi|iM1o3cC~Ex&q1A*pa;?TX{X&7oD}>J(RYt`dt6CgW?54@geqctDG6tnYqi(y%r$`t15NPb7 zu+u(4E@ausQEx!oxK1pAz^NpBMa-W5H~Z!=BM=EUAJ#;j#8=x0W{;J z#q+%DS{X%JmR-d>d7Ohn?Y3Ot&RTYS# zLk`iiNDE#(=+TRZy?FNI(UX^A{{jC6L3ZD(aTTVTdhb>B@4Z*A-mm_df3dH7r_*T@ z81<*W^Xtj8?xEggSSg}tN%~_=`z)+^KT=Hb$D|*tNX4KNm8cfz)^9`!l|Cja-D7>B z=xvgABxqjRTaQ7OHZ zG;i5{dZ0NR-DqZbM6+g&Uth^s?+ZPVMLsGE>&!H|@uG4*wxhCem}F}l@q1V!05pr4 z(&i;U27%u5#rqlaVo;QNl)F_khkqvy(}A-nX|=V*dSzz-VG%*-ES;g|I9LxUuD^{K zbPDqdOna@#ce&M0s>1kg^Lw&z z(bIuXzlWKuaXQra&1)V{9ymd#;hUlJ)SPZUCiix+kN*#>6KAONDi}})k;KuOLVdhu z<`>6a@tOm5$P=dMfE^^SDXtWI!Vi+2sOEJ98Zh7)DFQdZdj^wy{uf3TkNa@|HVmGp z?_0SU+I0h64%C2BBdg0n!1cTMM?QEYOI;12nLmLGazag3DB?z!b=y(~5RVB5#Qm(ine^*$NpmCCI- of{vjh$b54e1;&|7a5Ow*@$1^U+AQ@7zL0R!g2Kl|TO^|OAE_?>UjP6A diff --git a/src/utils/data/__pycache__/preprocess.cpython-38.pyc b/src/utils/data/__pycache__/preprocess.cpython-38.pyc index d791b541da79a6b0f9625e6d814bf2b7a58206c5..d2715d6a0585b17b42ae035798fe3055b887bd24 100644 GIT binary patch delta 2906 zcmZ`*OKclO7~a`k+w1rd=iMYt`bg8(Z4x&vr4OYcDKt;Lsirija_c;)KK*Cgs;I>O&pJ=2Er0*aKmW}C zoPT_0;K{!9bSl-Lz-Q+t!Zf<>r8n2UfEiUNDJgj(shZa$)$=-4sdih*N2pFCK%+EE zW7N1!@-e!HHqZo&47!&lX$nShIzThD2}TXHnYMs@g0|8&ph?CfW%!O}l6} z&_+5)5BAVr*pi`J=vH7g(QR})&}Q04cK~gn{d6bLR=SJs2HHl+l#(5~_kEQJeu_MI z9KCrwotyS8=I3VS7ECAif?KSVEhq4EXP7%@6=6`ae4pi}0BzSfXB7v%g}`OS*__Xc z{8REAN%CLGWHt$hiKq!uGfq}xF__)lm z075HSQWY{f5w;+7A#6iftNkDS2{*q z0IdKa!M~35t+V9!NRI5~y-_RKh==HPvAfZZbz%>rTgfp0Av)RHCv`CmNEk%)w7!qU z`Dkq0x{Y(OTy#H36bT=XY~o+WXdX?TU>G)*LXg%c!T{n}Xk~Mr?ZQ!lS-uq%Y=>I& zLZ@~y5ZVpf3l(FWQPyM9bj(u47+$CSp>dcT<+rpR9*ZBk&;uK?(bY@|!zr0;#`1&z z$r&C_9JUo`J|4JU!49ml4>#s=0cPv6@lkviImW+=zdw$yfd=$~9YBy(Aqzq#N)|*s zjBz}$%NDgtftfeJ;Q~$tf3IO|od-WQ43OiTB<{V`jAzGJ!xSozVbW|TU^`IX234_8 zRYL$i7JRQ{2kbBxiEV9TC_Ij^wyZMAK{CNdlRcXb!&0DBl|^zB@IImB0lBGGRqiHt z!v*;`IechIoqS9NmNcQyxwgX&!49E8D3^#J^J|aA1!JUSmfxi2=#>vsgEiQs`IaZ5 zvzAFMCas1FJ&N|oh+$2Mi0OHy1rZJ0Leamr6y9E8T((P=X9lx*!w+x2zfoDrH?LWR z<<(g@%f~yf5Q8+33}NIV`dFI;)~b(`bfTZKKRzH;=$|9#xk~lW^EnMxBiKFbbG=Gv z6uKw$&pSGSjtgxpJ}GlwB&e?b!Lbg%DfszV1p(-RlWXb&IBH5^;f)fyEP!bFilnkS zf{ItN*vqhIgNqDpzl!|R05`Jh9F~>j2WHtD8Y#I&v*eG$APm2}IiY#{L2`GEy@-dL zM0lDSe_0JX1EUFc7H~OP=W$JDOqTpKPG3P-Q&lh(9ZxZ^&4pDTg)uqpVR|m{&r>~W z7vbNe9-U2BF&Bw$Wl<$ckglpQOV($Rst)xT3mSo!>Q-oaQN7!G)2JHvGeJ`|wg@## zlncs1<*af}xfo`sk$;!&C%yd7^x*?5MndmYhxQ3#=QY=+8;x%m!NyZ`%iicM{(7TV zYYCWXJ8Qxm%m+7`MH}V!Uhqr!-=>B=UHsr46N|FMaoyG&DcehxP%hA>sgM7TxZa%*eT#-<03)Lz{D}nof=h9c&81{)ty*iLgCD ze#5sbOW*}C7yN9BQJ69ejxZQ>5Nr=Dz$vdbch;~a%Am5nIKs!39mm!9`Q?r)jesJa zL$(3i3)od;v#f^*H?7nMOQ2kv$1IPzMd&&Z%Ng6T0=sB_082jVP{28zsD!t+bdyEC zujQv2RyXqy&LChH2z|r;b{-yzbI8HQAR>+lmBT8?xQ;N7uz*lScpe~2*a;+S0>Em( z5K+op&oS9BN=o-LIGVzTHw?lt9Q7kiAYcx%hIF63jSQ(*{xKTCDaMWBW0IXO4cYgW*==dp!XOm P|D&xl7BgZ-m(lSr@DGG# delta 2288 zcmZ`)-A@!(6rVd^JF{Qx;sOf@f{3sRtcn}4P_S*{XJS&tP{%au$Q^WGc6NGaRwY?O z8jTNaOk?kRqfbqn{txXRP#^l#bkg{&?@gMRw&^){uO-r;7WQO+B0gzc*rh_1Jbchav%+tekgpR_A9y+;? zj=@raj??`hi*$k>0NG0?=|PZtD7h?EW*$Fp6N_IXE5jzZ`)aOA?Z_4|U%k1xUGe6f-Ntys6i(7owb6sE(Jr4Vumct-+h&gBWc2j(wjK5mZiU*si{i#o zYO9CHIR9E*tKo%g93WOaN7%)MGQudrAi@|zrz`W?NpgTUv};unod)|ACC-FR7S(Ut zUdsuEKbEoHYOFU~euQ3tlEr`5Cb~SC)T?BgU(ubpJYJ$EYTAbDcr?~M`-bC%vA%W3 zW;eqw7hdRNWR{z$i=+F6b*6wwl8v$N-oy;vN{x3}e446~IsQ}X)5GZdX@>b@Sp*SA zCdnj@W2Ir=Wd~to@q0#jK*WmK#f^Xohp|xK8ArP;yf79>4QjNywilEu`kUIshVMZn zP8j`1mGENXvwdjKjG95+jhsdZ4X{y%Sy%8snp0$f|7L#v9{L0p&|OwV5IhhR5bg>J zm`PU6$Za^(@gmzYAhJ46E&fCLRG0Ue%nVuN3z^4fi+H#BCLf6b%tx%ae#9ow-dm

xtVkeJ-^)h z*wkUdWguh?oe}9m!dRLe1m3o#N6q?rcpF=oKgpHnI=scJbD@3PsfW<~ZqpBM!BU9C zmWV<62|Lff?-`N*lKAtUXEm*@LNKW)-#?52wp^Zc5?iW!Xz2t`-=4?wK9eD{-!(n#hBze{vThwLjJnpH@ z!0}`44g+$s_hF@C#ume4EXprictbQta{^qKIz+M+XpA5JY5#h>UIG!gU1k%{zf3?2G@2adBXje8sJS z+D-$Mz@>p{R+a0bhb`OZEeXpI;^%Mw%z35UvW( zMNA?P(exq^EEM}dM1sr~;jk?np-&y`6MdM5HF#GVP1^FDWdgGhzdHso6#j8}P}7aF GQT`iKjNifl diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index db55c7d..58d9dc6 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -1,6 +1,7 @@ from collections import Counter import numpy as np import torch as th +import torch.nn.functional as F import dgl import pickle import numba @@ -84,6 +85,35 @@ def seq_to_session_graph(seq): return g +def seq_to_temporal_session_graph(seq, times): + items, indices = np.unique(seq, return_index=True) + iid2nid = {iid: i for i, iid in enumerate(items)} + num_nodes = len(items) + + seq_nid = [iid2nid[iid] for iid in seq] + counter = Counter( + [(seq_nid[i], seq_nid[i+1]) for i in range(len(seq)-1)] + ) + edges = counter.keys() + if len(edges) > 0: + src, dst = zip(*edges) + weight = th.tensor(list(counter.values())) + else: + src, dst = [0], [0] + weight = th.ones(1).long() + + g = dgl.graph((src, dst), num_nodes=num_nodes) + + g.edata['w'] = weight + # print(len(times), g.number_of_nodes()) + g.ndata['t'] = th.tensor(times)[indices] + # print(g.edata) + + g.ndata['iid'] = th.from_numpy(items) + label_last(g, iid2nid[seq[-1]]) + + return g + def seq_to_ccs_graph(seq, order=1, coaDict=None): order1 = order @@ -229,6 +259,23 @@ def collate_fn(samples): return collate_fn +def collate_fn_factory_temporal(*seq_to_graph_fns): + def collate_fn(samples): + seqs, times, labels = zip(*samples) + inputs = [] + for seq_to_graph in seq_to_graph_fns: + graphs = list(map(seq_to_graph, seqs, times)) + num_nodes = th.tensor([graph.number_of_nodes() for graph in graphs], dtype=th.long) + max_num = max(num_nodes) + embeds_id = th.vstack([F.pad(graph.ndata['iid'], (0, max_num-len(graph.ndata['iid'])), value=graph.ndata['iid'][-1]) for graph in graphs]) + times = th.vstack([F.pad(graph.ndata['t'], (0, max_num-len(graph.ndata['iid'])), value=graph.ndata['iid'][-1]) for graph in graphs]) + bg = dgl.batch(graphs) + inputs.append(bg) + labels = th.LongTensor(labels) + return inputs, labels, embeds_id, times, num_nodes + + return collate_fn + def collate_fn_factory_ccs(seq_to_graph_fns, order): def collate_fn(samples): seqs, labels = zip(*samples) diff --git a/src/utils/data/dataset.py b/src/utils/data/dataset.py index 5e3e38c..1224e10 100644 --- a/src/utils/data/dataset.py +++ b/src/utils/data/dataset.py @@ -1,4 +1,5 @@ import itertools +from os import read import numpy as np import pandas as pd @@ -18,17 +19,24 @@ def read_sessions(filepath): sessions = sessions.apply(lambda x: list(map(int, x.split(',')))).values return sessions +def read_timestamps(filepath): + sessions = pd.read_csv(filepath, sep='\t', header=None, squeeze=True) + sessions = sessions.apply(lambda x: list(map(float, x.split(',')))).values + return sessions def read_dataset(dataset_dir): - train_sessions = read_sessions(dataset_dir / 'train.txt') - test_sessions = read_sessions(dataset_dir / 'test.txt') + train_sessions = read_sessions(dataset_dir / 'train.txt') + test_sessions = read_sessions(dataset_dir / 'test.txt') + train_timestamp = read_timestamps(dataset_dir / 'train_timestamp.txt') + test_timestamp = read_timestamps(dataset_dir / 'test_timestamp.txt') with open(dataset_dir / 'num_items.txt', 'r') as f: num_items = int(f.readline()) - return train_sessions, test_sessions, num_items + return train_sessions, test_sessions, train_timestamp, test_timestamp, num_items class AugmentedDataset: - def __init__(self, sessions, sort_by_length=False): - self.sessions = sessions + def __init__(self, sessions, timestamps, sort_by_length=False): + self.sessions = sessions + self.timestamps = timestamps # self.graphs = graphs index = create_index(sessions) # columns: sessionId, labelIndex @@ -41,10 +49,13 @@ def __init__(self, sessions, sort_by_length=False): def __getitem__(self, idx): #print(idx) sid, lidx = self.index[idx] - seq = self.sessions[sid][:lidx] - label = self.sessions[sid][lidx] + seq = self.sessions[sid][:lidx] + label = self.sessions[sid][lidx] + times = self.timestamps[sid][:lidx]# - self.sessions[sid][0] + temp = times[0] + times = [(t - temp) // 10000 for t in times] - return seq, label #,seq + return seq, times, label #,seq def __len__(self): return len(self.index) diff --git a/src/utils/data/preprocess.py b/src/utils/data/preprocess.py index 7f9a3dd..fba6394 100644 --- a/src/utils/data/preprocess.py +++ b/src/utils/data/preprocess.py @@ -1,6 +1,6 @@ import pandas as pd import numpy as np - +import time def get_session_id(df, interval): df_prev = df.shift() @@ -103,10 +103,17 @@ def train_test_split(df, test_split=0.2): def save_sessions(df, filepath): df = reorder_sessions_by_endtime(df) - print(df.heads()) - sessions = df.groupby('sessionId').itemId.apply(lambda x: ','.join(map(str, x))) + sessions = df.groupby('sessionId')# . + sessions = sessions.itemId.apply(lambda x: ','.join(map(str, x))) sessions.to_csv(filepath, sep='\t', header=False, index=False) - + # sessions_timestamp = sessions.timestamp.apply(lambda x: ','.join(map(str, x))) + +def save_sessions_timestamp(df, filepath): + df = reorder_sessions_by_endtime(df) + df['timestamp'] = df['timestamp'].apply(lambda x: time.mktime(x.timetuple())) + sessions = df.groupby('sessionId')# . + sessions = sessions.timestamp.apply(lambda x: ','.join(map(str, x))) + sessions.to_csv(filepath, sep='\t', header=False, index=False) def save_dataset(dataset_dir, df_train, df_test): # filter items in test but not in train @@ -127,6 +134,8 @@ def save_dataset(dataset_dir, df_train, df_test): dataset_dir.mkdir(parents=True, exist_ok=True) save_sessions(df_train, dataset_dir / 'train.txt') save_sessions(df_test, dataset_dir / 'test.txt') + save_sessions_timestamp(df_train, dataset_dir / 'train_timestamp.txt') + save_sessions_timestamp(df_test, dataset_dir / 'test_timestamp.txt') num_items = len(uniques) with open(dataset_dir / 'num_items.txt', 'w') as f: f.write(str(num_items)) diff --git a/src/utils/train.py b/src/utils/train.py index 4de090a..0b6c7a9 100644 --- a/src/utils/train.py +++ b/src/utils/train.py @@ -24,12 +24,15 @@ def fix_weight_decay(model): def prepare_batch(batch, device): - inputs, labels = batch + inputs, labels, embeds_ids, times, num_nodes = batch # inputs, labels = batch inputs_gpu = [x.to(device) for x in inputs] labels_gpu = labels.to(device) + embeds_ids = embeds_ids.to(device) + times = times.to(device) + num_nodes = num_nodes.to(device) - return inputs_gpu, labels_gpu + return inputs_gpu, labels_gpu, embeds_ids, times, num_nodes # return inputs_gpu, 0, labels_gpu, 0 @@ -41,8 +44,8 @@ def evaluate(model, data_loader, device, cutoff=20): with th.no_grad(): for batch in data_loader: - inputs, labels = prepare_batch(batch, device) - logits = model(*inputs) + inputs, labels, embeds_ids, times, num_nodes = prepare_batch(batch, device) + logits = model(*inputs, embeds_ids, times, num_nodes) batch_size = logits.size(0) num_samples += batch_size @@ -92,9 +95,9 @@ def train(self, epochs, log_interval=100): for epoch in tqdm(range(epochs)): self.model.train() for batch in self.train_loader: - inputs, labels = prepare_batch(batch, self.device) + inputs, labels, embeds_ids, times, num_nodes = prepare_batch(batch, self.device) self.optimizer.zero_grad() - scores = self.model(*inputs) + scores = self.model(*inputs, embeds_ids, times, num_nodes) assert not th.isnan(scores).any() loss = nn.functional.nll_loss(scores, labels) loss.backward() From f9a39d40a31b36dfef9a063713114efc47e64fd4 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Mon, 10 Jan 2022 09:40:50 +0800 Subject: [PATCH 03/57] first --- src/models/niser_ode.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index ccfb02c..84a2fb0 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -193,7 +193,7 @@ def forward(self, mg, embeds_ids, times, num_nodes): time_embeds = self.recover(z_T[:, 1]) - out += time_embeds + # out += time_embeds last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) if self.norm: @@ -201,7 +201,7 @@ def forward(self, mg, embeds_ids, times, num_nodes): sr_g = self.readout(mg, feat, last_nodes) sr_l = feat[last_nodes] sr = th.cat([sr_l, sr_g], dim=1) - sr = self.fc_sr(sr) + sr = self.fc_sr(sr) + time_embeds if self.norm: sr = sr.div(th.norm(sr, p=2, dim=-1, keepdim=True) + 1e-12) target = self.embedding(self.indices) From 5fb2c6fd68a322270b2d2e0ec1fa0b2821c0c616 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Mon, 10 Jan 2022 22:07:22 +0800 Subject: [PATCH 04/57] pre --- src/preprocess.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/preprocess.py b/src/preprocess.py index 9f29202..221c177 100644 --- a/src/preprocess.py +++ b/src/preprocess.py @@ -1,6 +1,6 @@ from pathlib import Path import argparse -import time + parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter) From 45ce627d9d0828115881fd39d0a6631741bcdef3 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 16:13:40 +0800 Subject: [PATCH 05/57] ode --- src/models/niser_ode.py | 121 +++++++++++++++++++++++++++++++++------- 1 file changed, 102 insertions(+), 19 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 84a2fb0..d327173 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -8,7 +8,95 @@ import dgl.ops as F import dgl.function as fn +from dgl.nn.pytorch import GraphConv + import torchcde +from torchdiffeq import odeint_adjoint + +odeint_adjoint() + +class GraphGRUODE(nn.Module): + + def __init__(self, in_dim, hid_dim, device=th.device('cpu'), gnn='GCNConv', bias=True, **kwargs): + + super(GraphGRUODE, self).__init__() + + self.in_dim = in_dim + self.hid_dim = hid_dim + self.device = device + self.gnn = gnn + self.bias = bias + + if self.gnn == 'GCNConv': + # self.lin_xx = GCNConv(self.in_dim+self.hid_dim, self.hid_dim, bias=self.bias) + # self.lin_hx = nn.Linear(self.hid_dim, self.in_dim, bias=self.bias) + self.lin_xz = GraphConv(self.in_dim, self.hid_dim, bias=self.bias) + self.lin_xr = GraphConv(self.in_dim, self.hid_dim, bias=self.bias) + self.lin_xh = GraphConv(self.in_dim, self.hid_dim, bias=self.bias) + self.lin_hz = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias) + self.lin_hr = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias) + self.lin_hh = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias) + elif self.gnn == 'Linear': + self.lin_xx = nn.Linear(self.in_dim, self.hid_dim, bias=self.bias) + self.lin_hx = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) + self.lin_xz = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) + self.lin_xr = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) + self.lin_xh = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) + self.lin_hz = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) + self.lin_hr = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) + self.lin_hh = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) + else: + raise NotImplementedError + + self.edge_index = None + self.x = None + + self.reset_parameters() + + def reset_parameters(self): + + for module in self.modules(): + if hasattr(module, "reset_parameters"): + module.reset_parameters() + + def set_graph(self, graph: dgl.DGLGraph): + + self.graph = graph + + def set_x(self, x): + self.x = x + + def forward(self, t, h): + + # x = torch.zeros_like(h).to(self.device) + + # edge_index = self.edge_index_batchs[0] + + edge_idx = self.graph.filter_edges(self.graph.edata['t'] <= t) + edge_index = self.graph.edges() + graph = dgl.graph((edge_index[0][edge_idx], edge_index[1][edge_idx]), num_nodes=self.graph.number_of_nodes(), device=self.device) + + x = self.x + + if self.gnn != 'Linear': + # x = self.lin_xx(torch.cat((self.x.to(self.device), h), dim=1), edge_index).to(self.device) + xr, xz, xh = self.lin_xr(graph, x), self.lin_xz(graph, x), self.lin_xh(graph, x) + r = th.sigmoid(xr + self.lin_hr(h, edge_index)) + z = th.sigmoid(xz + self.lin_hz(h, edge_index)) + u = th.tanh(xh + self.lin_hh(r * h, edge_index)) + else: + # print(h.shape) + h = self.lin_hx(h)+self.lin_xx(x) + # x = self.propagate(edge_index=edge_index, x=h, aggr='mean')-h + xr, xz, xh = self.lin_xr(x), self.lin_xz(x), self.lin_xh(x) + r = th.sigmoid(xr + self.lin_hr(h)) + z = th.sigmoid(xz + self.lin_hz(h)) + u = th.tanh(xh + self.lin_hh(r * h)) + + + dh = (1 - z) * (u - h) + # self.x = self.hx(dh, edge_index) + return dh class CDEFunc(nn.Module): def __init__(self, input_channels, hidden_channels): @@ -27,6 +115,11 @@ def __init__(self, input_channels, hidden_channels): # For most purposes the t argument can probably be ignored; unless you want your CDE to behave differently at # different times, which would be unusual. But it's there if you need it! ###################### + + def set_graph(self, graph: dgl.DGLGraph): + + self.graph = graph + def forward(self, t, z): # z has shape (batch, hidden_channels) z = self.linear1(z) @@ -151,15 +244,12 @@ def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=Tru feat_drop=feat_drop, activation=None, ) + + self.ODEFunc = GraphGRUODE(self.embedding_dim, self.embedding_dim // 2, device=self.readout.fc_u.weight.device) self.feat_drop = nn.Dropout(feat_drop) self.fc_sr = nn.Linear(input_dim + embedding_dim, embedding_dim, bias=False) - self.reduce = nn.Linear(embedding_dim, 32) - self.recover = nn.Linear(32, embedding_dim) - self.cde_func = CDEFunc(33, 32) - self.initial = nn.Linear(33, 32) - self.reset_parameters() def reset_parameters(self): @@ -177,23 +267,16 @@ def forward(self, mg, embeds_ids, times, num_nodes): for i, layer in enumerate(self.layers): out = layer(mg, out) - feat_ode = self.reduce(self.feat_drop(self.embedding(embeds_ids))) + feat_ode = self.feat_drop(self.embedding(embeds_ids)) if self.norm: feat_ode = feat_ode.div(th.norm(feat_ode, p=2, dim=-1, keepdim=True) + 1e-12) - feat_ode = th.cat([times.unsqueeze(-1), feat_ode], dim=-1) - X = torchcde.LinearInterpolation(feat_ode) - X0 = X.evaluate(X.interval[0]) - z0 = self.initial(X0) # print(X.interval) - z_T = torchcde.cdeint(X=X, - z0=z0, - func=self.cde_func, - t=X.interval) - - time_embeds = self.recover(z_T[:, 1]) - - # out += time_embeds + self.ODEFunc.set_graph(mg) + self.ODEFunc.set_x(feat_ode) + t_end = mg.edata['t'].max() + t = th.tensor([0., t_end], device=mg.device) + feat = odeint_adjoint(self.ODEFunc, feat_ode, t=t) last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) if self.norm: @@ -201,7 +284,7 @@ def forward(self, mg, embeds_ids, times, num_nodes): sr_g = self.readout(mg, feat, last_nodes) sr_l = feat[last_nodes] sr = th.cat([sr_l, sr_g], dim=1) - sr = self.fc_sr(sr) + time_embeds + sr = self.fc_sr(sr) if self.norm: sr = sr.div(th.norm(sr, p=2, dim=-1, keepdim=True) + 1e-12) target = self.embedding(self.indices) From e1bd59228b6fc8977d191c8824eeb4db11b859df Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 16:21:09 +0800 Subject: [PATCH 06/57] edge --- src/utils/data/collate.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 58d9dc6..0643548 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -107,7 +107,8 @@ def seq_to_temporal_session_graph(seq, times): g.edata['w'] = weight # print(len(times), g.number_of_nodes()) g.ndata['t'] = th.tensor(times)[indices] - # print(g.edata) + g.edata['t'] = th.tensor(times[1:][indices]) + print(g.edata) g.ndata['iid'] = th.from_numpy(items) label_last(g, iid2nid[seq[-1]]) From 33344b838700c2c7dbbeff748ea9078d274319ac Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 17:28:26 +0800 Subject: [PATCH 07/57] edge --- src/models/niser_ode.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index d327173..2692970 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -10,11 +10,8 @@ from dgl.nn.pytorch import GraphConv -import torchcde from torchdiffeq import odeint_adjoint -odeint_adjoint() - class GraphGRUODE(nn.Module): def __init__(self, in_dim, hid_dim, device=th.device('cpu'), gnn='GCNConv', bias=True, **kwargs): From cdd5fa0a6c50395f9fbaec54e68ba3e4051f6648 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 17:33:46 +0800 Subject: [PATCH 08/57] edge --- src/models/niser_ode.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 2692970..4226156 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -48,13 +48,11 @@ def __init__(self, in_dim, hid_dim, device=th.device('cpu'), gnn='GCNConv', bias self.edge_index = None self.x = None - self.reset_parameters() + # self.reset_parameters() - def reset_parameters(self): + # def reset_parameters(self): - for module in self.modules(): - if hasattr(module, "reset_parameters"): - module.reset_parameters() + # self.lin_xz.reset_parameters() def set_graph(self, graph: dgl.DGLGraph): From 0e6d5d387c095566483b7e1969a00e446513feeb Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 17:36:00 +0800 Subject: [PATCH 09/57] edge --- src/utils/data/collate.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 0643548..fe276ab 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -107,7 +107,8 @@ def seq_to_temporal_session_graph(seq, times): g.edata['w'] = weight # print(len(times), g.number_of_nodes()) g.ndata['t'] = th.tensor(times)[indices] - g.edata['t'] = th.tensor(times[1:][indices]) + print(g.ndata['t']) + g.edata['t'] = th.tensor(times)[indices][1:] print(g.edata) g.ndata['iid'] = th.from_numpy(items) From e75a7b50caf31e1fcaaac5a840b75062b327e3b7 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 17:39:13 +0800 Subject: [PATCH 10/57] edge --- src/utils/data/collate.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index fe276ab..627b579 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -108,7 +108,10 @@ def seq_to_temporal_session_graph(seq, times): # print(len(times), g.number_of_nodes()) g.ndata['t'] = th.tensor(times)[indices] print(g.ndata['t']) - g.edata['t'] = th.tensor(times)[indices][1:] + if g.number_of_edges > 0: + g.edata['t'] = th.tensor(times)[indices][1:] + else: + g.edata['t'] = th.tensor([]) print(g.edata) g.ndata['iid'] = th.from_numpy(items) From 3ebeb420eb007b8ff613ef9a347313d239de8d76 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 17:42:50 +0800 Subject: [PATCH 11/57] edge --- src/utils/data/collate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 627b579..dec5d34 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -108,7 +108,7 @@ def seq_to_temporal_session_graph(seq, times): # print(len(times), g.number_of_nodes()) g.ndata['t'] = th.tensor(times)[indices] print(g.ndata['t']) - if g.number_of_edges > 0: + if g.number_of_edges() > 0: g.edata['t'] = th.tensor(times)[indices][1:] else: g.edata['t'] = th.tensor([]) From 407a243d137d6b6f39ff70d37adcd8097b77ec33 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:08:02 +0800 Subject: [PATCH 12/57] edge --- src/utils/data/collate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index dec5d34..138f15d 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -106,7 +106,7 @@ def seq_to_temporal_session_graph(seq, times): g.edata['w'] = weight # print(len(times), g.number_of_nodes()) - g.ndata['t'] = th.tensor(times)[indices] + g.ndata['t'] = th.tensor(times)[indices[::-1]] print(g.ndata['t']) if g.number_of_edges() > 0: g.edata['t'] = th.tensor(times)[indices][1:] From a641e770b1a9196c999123ef9f55d9c896f9cf8d Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:09:15 +0800 Subject: [PATCH 13/57] edge --- src/utils/data/collate.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 138f15d..1a526c6 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -108,10 +108,8 @@ def seq_to_temporal_session_graph(seq, times): # print(len(times), g.number_of_nodes()) g.ndata['t'] = th.tensor(times)[indices[::-1]] print(g.ndata['t']) - if g.number_of_edges() > 0: - g.edata['t'] = th.tensor(times)[indices][1:] - else: - g.edata['t'] = th.tensor([]) + g.edata['t'] = th.tensor(times)[1:] + print(g.edata) g.ndata['iid'] = th.from_numpy(items) From e0080d887d48d0bb809baf3c8f8560bea890d677 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:10:27 +0800 Subject: [PATCH 14/57] edge --- src/utils/data/collate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 1a526c6..dc40e7b 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -106,7 +106,7 @@ def seq_to_temporal_session_graph(seq, times): g.edata['w'] = weight # print(len(times), g.number_of_nodes()) - g.ndata['t'] = th.tensor(times)[indices[::-1]] + g.ndata['t'] = th.tensor(times)[indices] print(g.ndata['t']) g.edata['t'] = th.tensor(times)[1:] From defd652e63681690ea28fc3f1e9fee7a52b536b2 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:11:51 +0800 Subject: [PATCH 15/57] edge --- src/utils/data/collate.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index dc40e7b..cdb36c1 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -108,6 +108,7 @@ def seq_to_temporal_session_graph(seq, times): # print(len(times), g.number_of_nodes()) g.ndata['t'] = th.tensor(times)[indices] print(g.ndata['t']) + print(times) g.edata['t'] = th.tensor(times)[1:] print(g.edata) From 6ed6c2c9020e83ed23cfd7ff2d7bd036519b8765 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:12:11 +0800 Subject: [PATCH 16/57] edge --- src/utils/data/collate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index cdb36c1..c532448 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -108,7 +108,7 @@ def seq_to_temporal_session_graph(seq, times): # print(len(times), g.number_of_nodes()) g.ndata['t'] = th.tensor(times)[indices] print(g.ndata['t']) - print(times) + print(times, g.number_of_nodes(), g.number_of_edges()) g.edata['t'] = th.tensor(times)[1:] print(g.edata) From f730f19c614c88ad03d2e4a5271a822e53565329 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:15:17 +0800 Subject: [PATCH 17/57] edge --- src/utils/data/collate.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index c532448..1c99066 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -107,9 +107,11 @@ def seq_to_temporal_session_graph(seq, times): g.edata['w'] = weight # print(len(times), g.number_of_nodes()) g.ndata['t'] = th.tensor(times)[indices] - print(g.ndata['t']) - print(times, g.number_of_nodes(), g.number_of_edges()) - g.edata['t'] = th.tensor(times)[1:] + + if g.number_of_edges() == 1: + g.edata['t'] = th.tensor(times) + else: + g.edata['t'] = th.tensor(times)[1:] print(g.edata) From 06473068495d7fcced8eea01b23a97b1ecce5177 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:16:23 +0800 Subject: [PATCH 18/57] edge --- src/utils/data/collate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 1c99066..74d1b3e 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -109,7 +109,7 @@ def seq_to_temporal_session_graph(seq, times): g.ndata['t'] = th.tensor(times)[indices] if g.number_of_edges() == 1: - g.edata['t'] = th.tensor(times) + g.edata['t'] = th.tensor(times)[0] else: g.edata['t'] = th.tensor(times)[1:] From eba564203950bc1c1cc6730197d235b52d3bb3c4 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:16:57 +0800 Subject: [PATCH 19/57] ode --- src/utils/data/collate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 74d1b3e..16378f1 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -109,7 +109,7 @@ def seq_to_temporal_session_graph(seq, times): g.ndata['t'] = th.tensor(times)[indices] if g.number_of_edges() == 1: - g.edata['t'] = th.tensor(times)[0] + g.edata['t'] = th.tensor(times)[-1] else: g.edata['t'] = th.tensor(times)[1:] From 9ea9ade919439cbf7a267d5c0006efaaae7e5011 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:18:19 +0800 Subject: [PATCH 20/57] ode --- src/utils/data/collate.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 16378f1..23c7d04 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -108,6 +108,7 @@ def seq_to_temporal_session_graph(seq, times): # print(len(times), g.number_of_nodes()) g.ndata['t'] = th.tensor(times)[indices] + print(g.edata, times, g.number_of_edges(), g.number_of_nodes()) if g.number_of_edges() == 1: g.edata['t'] = th.tensor(times)[-1] else: From 9e3538c5a01ffbb0f01c5e0bdcc9a78ffb4b7761 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:19:36 +0800 Subject: [PATCH 21/57] ode --- src/utils/data/collate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 23c7d04..751764a 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -110,7 +110,7 @@ def seq_to_temporal_session_graph(seq, times): print(g.edata, times, g.number_of_edges(), g.number_of_nodes()) if g.number_of_edges() == 1: - g.edata['t'] = th.tensor(times)[-1] + g.edata['t'] = th.tensor(times)[0] else: g.edata['t'] = th.tensor(times)[1:] From b73260823c1c7e94cb7c94ac58e9bd2786dc42c1 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:21:37 +0800 Subject: [PATCH 22/57] ode --- src/utils/data/collate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 751764a..32e0233 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -109,7 +109,7 @@ def seq_to_temporal_session_graph(seq, times): g.ndata['t'] = th.tensor(times)[indices] print(g.edata, times, g.number_of_edges(), g.number_of_nodes()) - if g.number_of_edges() == 1: + if g.number_of_edges() == 1 and g.number_of_nodes() == 1: g.edata['t'] = th.tensor(times)[0] else: g.edata['t'] = th.tensor(times)[1:] From dd627e54e2e989f261da3c021566594ad5dec6a3 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:24:03 +0800 Subject: [PATCH 23/57] ode --- src/utils/data/collate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 32e0233..c6129cb 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -112,7 +112,7 @@ def seq_to_temporal_session_graph(seq, times): if g.number_of_edges() == 1 and g.number_of_nodes() == 1: g.edata['t'] = th.tensor(times)[0] else: - g.edata['t'] = th.tensor(times)[1:] + g.edata['t'] = th.tensor(times)[-g.number_of_edges():] print(g.edata) From 6b87c55965959b56104170a17002491ea7df33fe Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:29:47 +0800 Subject: [PATCH 24/57] ode --- src/utils/data/collate.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index c6129cb..b47e8c5 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -108,13 +108,13 @@ def seq_to_temporal_session_graph(seq, times): # print(len(times), g.number_of_nodes()) g.ndata['t'] = th.tensor(times)[indices] - print(g.edata, times, g.number_of_edges(), g.number_of_nodes()) + # print(g.edata, times, g.number_of_edges(), g.number_of_nodes()) if g.number_of_edges() == 1 and g.number_of_nodes() == 1: - g.edata['t'] = th.tensor(times)[0] + g.edata['t'] = th.tensor(times)[0].unsqueeze(0) else: g.edata['t'] = th.tensor(times)[-g.number_of_edges():] - print(g.edata) + # print(g.edata) g.ndata['iid'] = th.from_numpy(items) label_last(g, iid2nid[seq[-1]]) From e50092dcd1b3e6b668c95345968353e1a4043a98 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:32:19 +0800 Subject: [PATCH 25/57] ode --- src/models/niser_ode.py | 2 +- src/utils/data/collate.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 4226156..fc27f2f 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -67,7 +67,7 @@ def forward(self, t, h): # edge_index = self.edge_index_batchs[0] - edge_idx = self.graph.filter_edges(self.graph.edata['t'] <= t) + edge_idx = self.graph.filter_edges(lambda edges: edges.data['t'] <= t) edge_index = self.graph.edges() graph = dgl.graph((edge_index[0][edge_idx], edge_index[1][edge_idx]), num_nodes=self.graph.number_of_nodes(), device=self.device) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index b47e8c5..6fbea03 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -110,9 +110,9 @@ def seq_to_temporal_session_graph(seq, times): # print(g.edata, times, g.number_of_edges(), g.number_of_nodes()) if g.number_of_edges() == 1 and g.number_of_nodes() == 1: - g.edata['t'] = th.tensor(times)[0].unsqueeze(0) + g.edata['t'] = th.tensor(times)[0].unsqueeze(-1) else: - g.edata['t'] = th.tensor(times)[-g.number_of_edges():] + g.edata['t'] = th.tensor(times)[-g.number_of_edges():].unsqueeze(-1) # print(g.edata) From 2e3ceac7555d46c922ffde5705938a89a7a83948 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:34:34 +0800 Subject: [PATCH 26/57] ode --- src/utils/data/collate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 6fbea03..27b7431 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -112,7 +112,7 @@ def seq_to_temporal_session_graph(seq, times): if g.number_of_edges() == 1 and g.number_of_nodes() == 1: g.edata['t'] = th.tensor(times)[0].unsqueeze(-1) else: - g.edata['t'] = th.tensor(times)[-g.number_of_edges():].unsqueeze(-1) + g.edata['t'] = th.tensor(times)[-g.number_of_edges():] # print(g.edata) From ab074523a5c60790174af16601cc13b358d788a6 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:36:14 +0800 Subject: [PATCH 27/57] ode --- src/models/niser_ode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index fc27f2f..6f6e757 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -70,7 +70,7 @@ def forward(self, t, h): edge_idx = self.graph.filter_edges(lambda edges: edges.data['t'] <= t) edge_index = self.graph.edges() graph = dgl.graph((edge_index[0][edge_idx], edge_index[1][edge_idx]), num_nodes=self.graph.number_of_nodes(), device=self.device) - + graph = dgl.add_self_loop(graph) x = self.x if self.gnn != 'Linear': From 1d1b3b8c438bfa135dc46290d34327d150389bf6 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:38:33 +0800 Subject: [PATCH 28/57] ode --- src/models/niser_ode.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 6f6e757..2039f56 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -27,12 +27,12 @@ def __init__(self, in_dim, hid_dim, device=th.device('cpu'), gnn='GCNConv', bias if self.gnn == 'GCNConv': # self.lin_xx = GCNConv(self.in_dim+self.hid_dim, self.hid_dim, bias=self.bias) # self.lin_hx = nn.Linear(self.hid_dim, self.in_dim, bias=self.bias) - self.lin_xz = GraphConv(self.in_dim, self.hid_dim, bias=self.bias) - self.lin_xr = GraphConv(self.in_dim, self.hid_dim, bias=self.bias) - self.lin_xh = GraphConv(self.in_dim, self.hid_dim, bias=self.bias) - self.lin_hz = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias) - self.lin_hr = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias) - self.lin_hh = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias) + self.lin_xz = GraphConv(self.in_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) + self.lin_xr = GraphConv(self.in_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) + self.lin_xh = GraphConv(self.in_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) + self.lin_hz = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) + self.lin_hr = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) + self.lin_hh = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) elif self.gnn == 'Linear': self.lin_xx = nn.Linear(self.in_dim, self.hid_dim, bias=self.bias) self.lin_hx = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) @@ -70,7 +70,7 @@ def forward(self, t, h): edge_idx = self.graph.filter_edges(lambda edges: edges.data['t'] <= t) edge_index = self.graph.edges() graph = dgl.graph((edge_index[0][edge_idx], edge_index[1][edge_idx]), num_nodes=self.graph.number_of_nodes(), device=self.device) - graph = dgl.add_self_loop(graph) + # graph = dgl.add_self_loop(graph) x = self.x if self.gnn != 'Linear': From f247ff70790bf8e65b9e49239ff9c559d351cd1f Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:42:06 +0800 Subject: [PATCH 29/57] ode --- src/models/niser_ode.py | 15 ++++++--------- 1 file changed, 6 insertions(+), 9 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 2039f56..06a273c 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -73,6 +73,7 @@ def forward(self, t, h): # graph = dgl.add_self_loop(graph) x = self.x + print(x.shape) if self.gnn != 'Linear': # x = self.lin_xx(torch.cat((self.x.to(self.device), h), dim=1), edge_index).to(self.device) xr, xz, xh = self.lin_xr(graph, x), self.lin_xz(graph, x), self.lin_xh(graph, x) @@ -258,20 +259,16 @@ def forward(self, mg, embeds_ids, times, num_nodes): feat = self.feat_drop(self.embedding(iid)) if self.norm: feat = feat.div(th.norm(feat, p=2, dim=-1, keepdim=True) + 1e-12) - out = feat - for i, layer in enumerate(self.layers): - out = layer(mg, out) - - feat_ode = self.feat_drop(self.embedding(embeds_ids)) - if self.norm: - feat_ode = feat_ode.div(th.norm(feat_ode, p=2, dim=-1, keepdim=True) + 1e-12) + # out = feat + # for i, layer in enumerate(self.layers): + # out = layer(mg, out) # print(X.interval) self.ODEFunc.set_graph(mg) - self.ODEFunc.set_x(feat_ode) + self.ODEFunc.set_x(feat) t_end = mg.edata['t'].max() t = th.tensor([0., t_end], device=mg.device) - feat = odeint_adjoint(self.ODEFunc, feat_ode, t=t) + feat = odeint_adjoint(self.ODEFunc, feat, t=t) last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) if self.norm: From 075d6777aed67dd009025623733531c5b162238b Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:45:27 +0800 Subject: [PATCH 30/57] ode --- src/models/niser_ode.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 06a273c..4c0a10f 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -59,7 +59,7 @@ def set_graph(self, graph: dgl.DGLGraph): self.graph = graph def set_x(self, x): - self.x = x + self.x = x.to(self.device) def forward(self, t, h): @@ -254,6 +254,7 @@ def reset_parameters(self): weight.data.uniform_(-stdv, stdv) def forward(self, mg, embeds_ids, times, num_nodes): + iid = mg.ndata['iid'] feat = self.feat_drop(self.embedding(iid)) From dc441aab650887e92fb2332da68203abd119b5ca Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:45:56 +0800 Subject: [PATCH 31/57] ode --- src/models/niser_ode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 4c0a10f..c7c07ce 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -269,7 +269,7 @@ def forward(self, mg, embeds_ids, times, num_nodes): self.ODEFunc.set_x(feat) t_end = mg.edata['t'].max() t = th.tensor([0., t_end], device=mg.device) - feat = odeint_adjoint(self.ODEFunc, feat, t=t) + feat = odeint_adjoint(self.ODEFunc, feat, t=t)[-1] last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) if self.norm: From 0204ac0261da1f0e9f95b81054761c2c357ca6de Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:47:49 +0800 Subject: [PATCH 32/57] ode --- src/models/niser_ode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index c7c07ce..f55c632 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -241,7 +241,7 @@ def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=Tru activation=None, ) - self.ODEFunc = GraphGRUODE(self.embedding_dim, self.embedding_dim // 2, device=self.readout.fc_u.weight.device) + self.ODEFunc = GraphGRUODE(self.embedding_dim, self.embedding_dim // 2, device=self.embedding.weight.device).to(self.embedding.weight.device) self.feat_drop = nn.Dropout(feat_drop) self.fc_sr = nn.Linear(input_dim + embedding_dim, embedding_dim, bias=False) From 6a0ebc978516ac383abef3cf3065b61101b4deae Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:48:55 +0800 Subject: [PATCH 33/57] ode --- src/models/niser_ode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index f55c632..02b87e5 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -73,7 +73,7 @@ def forward(self, t, h): # graph = dgl.add_self_loop(graph) x = self.x - print(x.shape) + print(x.device, graph.device, self.device) if self.gnn != 'Linear': # x = self.lin_xx(torch.cat((self.x.to(self.device), h), dim=1), edge_index).to(self.device) xr, xz, xh = self.lin_xr(graph, x), self.lin_xz(graph, x), self.lin_xh(graph, x) From cabfb446002093674f899fea398d538b588c5c74 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:51:20 +0800 Subject: [PATCH 34/57] ode --- src/models/niser_ode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 02b87e5..dfed798 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -241,7 +241,7 @@ def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=Tru activation=None, ) - self.ODEFunc = GraphGRUODE(self.embedding_dim, self.embedding_dim // 2, device=self.embedding.weight.device).to(self.embedding.weight.device) + self.ODEFunc = GraphGRUODE(self.embedding_dim, self.embedding_dim // 2, device=th.device('cuda:0')) self.feat_drop = nn.Dropout(feat_drop) self.fc_sr = nn.Linear(input_dim + embedding_dim, embedding_dim, bias=False) From f4b38eded6eae86c2d4f9dd43db215e4b2857d66 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:53:40 +0800 Subject: [PATCH 35/57] ode --- src/models/niser_ode.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index dfed798..502e27e 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -73,13 +73,13 @@ def forward(self, t, h): # graph = dgl.add_self_loop(graph) x = self.x - print(x.device, graph.device, self.device) + if self.gnn != 'Linear': # x = self.lin_xx(torch.cat((self.x.to(self.device), h), dim=1), edge_index).to(self.device) xr, xz, xh = self.lin_xr(graph, x), self.lin_xz(graph, x), self.lin_xh(graph, x) - r = th.sigmoid(xr + self.lin_hr(h, edge_index)) - z = th.sigmoid(xz + self.lin_hz(h, edge_index)) - u = th.tanh(xh + self.lin_hh(r * h, edge_index)) + r = th.sigmoid(xr + self.lin_hr(graph, h)) + z = th.sigmoid(xz + self.lin_hz(graph, h)) + u = th.tanh(xh + self.lin_hh(graph, r * h)) else: # print(h.shape) h = self.lin_hx(h)+self.lin_xx(x) From bc0d45152beb139b2b2c69e63474cd265d6f1ecf Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 18:55:17 +0800 Subject: [PATCH 36/57] ode --- src/models/niser_ode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 502e27e..8c1b723 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -241,7 +241,7 @@ def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=Tru activation=None, ) - self.ODEFunc = GraphGRUODE(self.embedding_dim, self.embedding_dim // 2, device=th.device('cuda:0')) + self.ODEFunc = GraphGRUODE(self.embedding_dim, self.embedding_dim, device=th.device('cuda:0')) self.feat_drop = nn.Dropout(feat_drop) self.fc_sr = nn.Linear(input_dim + embedding_dim, embedding_dim, bias=False) From 20ded6d43955c0e63620953de78b0e4718494f27 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 19:05:47 +0800 Subject: [PATCH 37/57] ode --- src/models/niser_ode.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 8c1b723..d65b948 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -269,6 +269,7 @@ def forward(self, mg, embeds_ids, times, num_nodes): self.ODEFunc.set_x(feat) t_end = mg.edata['t'].max() t = th.tensor([0., t_end], device=mg.device) + print(t) feat = odeint_adjoint(self.ODEFunc, feat, t=t)[-1] last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) From 9388fd5d990db3a7d73c3c41030813bf2baa2219 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 19:07:23 +0800 Subject: [PATCH 38/57] ode --- src/utils/data/dataset.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/dataset.py b/src/utils/data/dataset.py index 1224e10..f3c4e46 100644 --- a/src/utils/data/dataset.py +++ b/src/utils/data/dataset.py @@ -53,7 +53,7 @@ def __getitem__(self, idx): label = self.sessions[sid][lidx] times = self.timestamps[sid][:lidx]# - self.sessions[sid][0] temp = times[0] - times = [(t - temp) // 10000 for t in times] + times = [(t - temp) // 100000 for t in times] return seq, times, label #,seq From 79c3788eb2889fa71ab83191f76d962f6129c201 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 19:14:49 +0800 Subject: [PATCH 39/57] ode --- src/models/niser_ode.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index d65b948..1b5b56f 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -67,10 +67,11 @@ def forward(self, t, h): # edge_index = self.edge_index_batchs[0] - edge_idx = self.graph.filter_edges(lambda edges: edges.data['t'] <= t) - edge_index = self.graph.edges() - graph = dgl.graph((edge_index[0][edge_idx], edge_index[1][edge_idx]), num_nodes=self.graph.number_of_nodes(), device=self.device) + # edge_idx = self.graph.filter_edges(lambda edges: edges.data['t'] <= t) + # edge_index = self.graph.edges() + # graph = dgl.graph((edge_index[0][edge_idx], edge_index[1][edge_idx]), num_nodes=self.graph.number_of_nodes(), device=self.device) # graph = dgl.add_self_loop(graph) + graph = self.graph x = self.x From deb5010f46db6093eb51a98e9ee5de6219d75544 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 21:16:36 +0800 Subject: [PATCH 40/57] ode --- src/utils/data/dataset.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/data/dataset.py b/src/utils/data/dataset.py index f3c4e46..6aa046a 100644 --- a/src/utils/data/dataset.py +++ b/src/utils/data/dataset.py @@ -53,7 +53,7 @@ def __getitem__(self, idx): label = self.sessions[sid][lidx] times = self.timestamps[sid][:lidx]# - self.sessions[sid][0] temp = times[0] - times = [(t - temp) // 100000 for t in times] + times = [(t - temp) / 1000000 for t in times] return seq, times, label #,seq From 1fd3d924061ca22b9590bef06d9af459e5198d7f Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 21:18:03 +0800 Subject: [PATCH 41/57] ode --- src/models/niser_ode.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 1b5b56f..7ba933a 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -67,11 +67,11 @@ def forward(self, t, h): # edge_index = self.edge_index_batchs[0] - # edge_idx = self.graph.filter_edges(lambda edges: edges.data['t'] <= t) - # edge_index = self.graph.edges() - # graph = dgl.graph((edge_index[0][edge_idx], edge_index[1][edge_idx]), num_nodes=self.graph.number_of_nodes(), device=self.device) + edge_idx = self.graph.filter_edges(lambda edges: edges.data['t'] <= t) + edge_index = self.graph.edges() + graph = dgl.graph((edge_index[0][edge_idx], edge_index[1][edge_idx]), num_nodes=self.graph.number_of_nodes(), device=self.device) # graph = dgl.add_self_loop(graph) - graph = self.graph + # graph = self.graph x = self.x From 86d970a2df039d073becb9b9807b6fb297eb56e6 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 23:36:08 +0800 Subject: [PATCH 42/57] ode --- src/models/niser_ode.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 7ba933a..74cef3b 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -269,9 +269,11 @@ def forward(self, mg, embeds_ids, times, num_nodes): self.ODEFunc.set_graph(mg) self.ODEFunc.set_x(feat) t_end = mg.edata['t'].max() - t = th.tensor([0., t_end], device=mg.device) + t = th.tensor([0., t_end / 10], device=mg.device) print(t) - feat = odeint_adjoint(self.ODEFunc, feat, t=t)[-1] + feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='eulr')[-1] + + print(feat.shape) last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) if self.norm: From f2a8bae5f0692c473ef288d482e8b9b9dabeee4b Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 23:36:59 +0800 Subject: [PATCH 43/57] ode --- src/models/niser_ode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 74cef3b..848618c 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -271,7 +271,7 @@ def forward(self, mg, embeds_ids, times, num_nodes): t_end = mg.edata['t'].max() t = th.tensor([0., t_end / 10], device=mg.device) print(t) - feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='eulr')[-1] + feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] print(feat.shape) From 1b30ca3aff0dc7c8320feea77f2585f21beeb00d Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 23:38:13 +0800 Subject: [PATCH 44/57] ode --- src/models/niser_ode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 848618c..dfb0b02 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -269,7 +269,7 @@ def forward(self, mg, embeds_ids, times, num_nodes): self.ODEFunc.set_graph(mg) self.ODEFunc.set_x(feat) t_end = mg.edata['t'].max() - t = th.tensor([0., t_end / 10], device=mg.device) + t = th.tensor([0., t_end], device=mg.device) print(t) feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] From f3cd51736c46a8c33054a6e6a10e810c51fdda28 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 23:39:23 +0800 Subject: [PATCH 45/57] ode --- src/models/niser_ode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index dfb0b02..f2eb01c 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -270,7 +270,7 @@ def forward(self, mg, embeds_ids, times, num_nodes): self.ODEFunc.set_x(feat) t_end = mg.edata['t'].max() t = th.tensor([0., t_end], device=mg.device) - print(t) + # print(t) feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] print(feat.shape) From 63696aa86f42f7908fa56b31c4c4f5a3c123dc27 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Fri, 14 Jan 2022 23:39:27 +0800 Subject: [PATCH 46/57] ode --- src/models/niser_ode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index f2eb01c..9362642 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -273,7 +273,7 @@ def forward(self, mg, embeds_ids, times, num_nodes): # print(t) feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] - print(feat.shape) + # print(feat.shape) last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) if self.norm: From 6387f6ed9c7dfb154237e6c26850467ce9fba24e Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Sat, 15 Jan 2022 00:02:38 +0800 Subject: [PATCH 47/57] ode --- src/models/niser_ode.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 9362642..7bd6305 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -266,12 +266,12 @@ def forward(self, mg, embeds_ids, times, num_nodes): # out = layer(mg, out) # print(X.interval) - self.ODEFunc.set_graph(mg) - self.ODEFunc.set_x(feat) - t_end = mg.edata['t'].max() - t = th.tensor([0., t_end], device=mg.device) - # print(t) - feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] + # self.ODEFunc.set_graph(mg) + # self.ODEFunc.set_x(feat) + # t_end = mg.edata['t'].max() + # t = th.tensor([0., t_end], device=mg.device) + # # print(t) + # feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] + feat # print(feat.shape) From 881db424dfba14d3c8cfb96eace06a4c8de2cae2 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Sat, 15 Jan 2022 00:15:57 +0800 Subject: [PATCH 48/57] ode --- src/models/niser_ode.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 7bd6305..9c6fdd7 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -266,12 +266,12 @@ def forward(self, mg, embeds_ids, times, num_nodes): # out = layer(mg, out) # print(X.interval) - # self.ODEFunc.set_graph(mg) - # self.ODEFunc.set_x(feat) - # t_end = mg.edata['t'].max() - # t = th.tensor([0., t_end], device=mg.device) - # # print(t) - # feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] + feat + self.ODEFunc.set_graph(mg) + self.ODEFunc.set_x(feat) + t_end = mg.edata['t'].max() + t = th.tensor([0., t_end / 10], device=mg.device) + # print(t) + feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] + feat # print(feat.shape) From f866d6b9493adae7f76547e0d2e683950ad0010e Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Sat, 15 Jan 2022 00:36:18 +0800 Subject: [PATCH 49/57] ode --- src/models/niser_ode.py | 19 ++++++++++--------- 1 file changed, 10 insertions(+), 9 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 9c6fdd7..1fe3ba8 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -261,17 +261,18 @@ def forward(self, mg, embeds_ids, times, num_nodes): feat = self.feat_drop(self.embedding(iid)) if self.norm: feat = feat.div(th.norm(feat, p=2, dim=-1, keepdim=True) + 1e-12) - # out = feat - # for i, layer in enumerate(self.layers): - # out = layer(mg, out) + out = feat + for i, layer in enumerate(self.layers): + out = layer(mg, out) + feat = out # print(X.interval) - self.ODEFunc.set_graph(mg) - self.ODEFunc.set_x(feat) - t_end = mg.edata['t'].max() - t = th.tensor([0., t_end / 10], device=mg.device) - # print(t) - feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] + feat + # self.ODEFunc.set_graph(mg) + # self.ODEFunc.set_x(feat) + # t_end = mg.edata['t'].max() + # t = th.tensor([0., t_end / 10], device=mg.device) + # # print(t) + # feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] + feat # print(feat.shape) From 65ffea76bafdd6a5aae734f7136a0863dce8b435 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Sat, 15 Jan 2022 01:36:29 +0800 Subject: [PATCH 50/57] ode --- src/models/niser_ode.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 1fe3ba8..a21d768 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -266,13 +266,13 @@ def forward(self, mg, embeds_ids, times, num_nodes): out = layer(mg, out) feat = out - # print(X.interval) - # self.ODEFunc.set_graph(mg) - # self.ODEFunc.set_x(feat) - # t_end = mg.edata['t'].max() - # t = th.tensor([0., t_end / 10], device=mg.device) - # # print(t) - # feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='euler')[-1] + feat + print(X.interval) + self.ODEFunc.set_graph(mg) + self.ODEFunc.set_x(feat) + t_end = mg.edata['t'].max() + t = th.tensor([0., t_end / 10], device=mg.device) + # print(t) + feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='rk4')[-1] + feat # print(feat.shape) From d38a3cabe84ab42cb2c49c87adf9e01733a238cc Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Sat, 15 Jan 2022 01:37:05 +0800 Subject: [PATCH 51/57] ode --- src/models/niser_ode.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index a21d768..740e320 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -264,15 +264,14 @@ def forward(self, mg, embeds_ids, times, num_nodes): out = feat for i, layer in enumerate(self.layers): out = layer(mg, out) - feat = out + # feat = out - print(X.interval) self.ODEFunc.set_graph(mg) self.ODEFunc.set_x(feat) t_end = mg.edata['t'].max() t = th.tensor([0., t_end / 10], device=mg.device) # print(t) - feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='rk4')[-1] + feat + feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='rk4')[-1] + out # print(feat.shape) From cd55b687d9321ec0fad66f734e7fb184b782fef0 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Mon, 17 Jan 2022 10:31:32 +0800 Subject: [PATCH 52/57] newdata --- src/models/niser_ode.py | 38 +++++++++++++----- src/scripts/main_niser_ode.py | 2 +- .../data/__pycache__/collate.cpython-39.pyc | Bin 9000 -> 11292 bytes .../data/__pycache__/dataset.cpython-39.pyc | Bin 2212 -> 3199 bytes src/utils/data/dataset.py | 26 +++++++++--- 5 files changed, 50 insertions(+), 16 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 740e320..3cd671c 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -10,7 +10,9 @@ from dgl.nn.pytorch import GraphConv -from torchdiffeq import odeint_adjoint +from torchdiffeq import odeint_adjoint, odeint + +from torch.autograd import Variable class GraphGRUODE(nn.Module): @@ -23,6 +25,7 @@ def __init__(self, in_dim, hid_dim, device=th.device('cpu'), gnn='GCNConv', bias self.device = device self.gnn = gnn self.bias = bias + self.dropout = nn.Dropout(0.1) if self.gnn == 'GCNConv': # self.lin_xx = GCNConv(self.in_dim+self.hid_dim, self.hid_dim, bias=self.bias) @@ -68,10 +71,15 @@ def forward(self, t, h): # edge_index = self.edge_index_batchs[0] edge_idx = self.graph.filter_edges(lambda edges: edges.data['t'] <= t) + # print(sum(edge_idx.long())) edge_index = self.graph.edges() graph = dgl.graph((edge_index[0][edge_idx], edge_index[1][edge_idx]), num_nodes=self.graph.number_of_nodes(), device=self.device) + graph = dgl.remove_self_loop(graph) + graph = dgl.add_reverse_edges(graph) # graph = dgl.add_self_loop(graph) # graph = self.graph + # x = self.dropout(self.x) + # h = self.dropout(h) x = self.x @@ -243,6 +251,9 @@ def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=Tru ) self.ODEFunc = GraphGRUODE(self.embedding_dim, self.embedding_dim, device=th.device('cuda:0')) + + # self.initial = GraphConv(self.embedding_dim, 2*self.embedding_dim, allow_zero_in_degree=True) + self.initial = nn.Linear(self.embedding_dim, self.embedding_dim) self.feat_drop = nn.Dropout(feat_drop) self.fc_sr = nn.Linear(input_dim + embedding_dim, embedding_dim, bias=False) @@ -253,6 +264,11 @@ def reset_parameters(self): stdv = 1.0 / math.sqrt(self.embedding_dim) for weight in self.parameters(): weight.data.uniform_(-stdv, stdv) + + def _reparameterized_sample(self, mean, std): + eps1 = th.FloatTensor(std.size()).normal_() + eps1 = Variable(eps1).to(mean.device) + return eps1.mul(std).add_(mean) def forward(self, mg, embeds_ids, times, num_nodes): @@ -261,19 +277,21 @@ def forward(self, mg, embeds_ids, times, num_nodes): feat = self.feat_drop(self.embedding(iid)) if self.norm: feat = feat.div(th.norm(feat, p=2, dim=-1, keepdim=True) + 1e-12) - out = feat - for i, layer in enumerate(self.layers): - out = layer(mg, out) + # out = feat + # for i, layer in enumerate(self.layers): + # out = layer(mg, out) # feat = out - + # mgs = dgl.add_reverse_edges(mg) + # feat0 = feat + # feat0_mean = nn.functional.tanh(feat0[:, :self.embedding_dim]) + # feat0_var = nn.functional.softplus(feat0[:, self.embedding_dim:]) + # feat0 = self._reparameterized_sample(feat0_mean, feat0_var) + self.ODEFunc.set_graph(mg) self.ODEFunc.set_x(feat) t_end = mg.edata['t'].max() - t = th.tensor([0., t_end / 10], device=mg.device) - # print(t) - feat = odeint_adjoint(self.ODEFunc, feat, t=t, method='rk4')[-1] + out - - # print(feat.shape) + t = th.tensor([0., t_end], device=mg.device) + feat = odeint(self.ODEFunc, feat, t=t, method='rk4')[-1] # + feat last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) if self.norm: diff --git a/src/scripts/main_niser_ode.py b/src/scripts/main_niser_ode.py index 255cc81..934f861 100644 --- a/src/scripts/main_niser_ode.py +++ b/src/scripts/main_niser_ode.py @@ -35,7 +35,7 @@ parser.add_argument( '--num-workers', type=int, - default=0, + default=10, help='the number of processes to load the input graphs', ) parser.add_argument( diff --git a/src/utils/data/__pycache__/collate.cpython-39.pyc b/src/utils/data/__pycache__/collate.cpython-39.pyc index a7596a4b5d61660316761361754743eeebf6e048..c980daaddfef0f60e7f4d86139b7f9697ffee742 100644 GIT binary patch delta 4075 zcmbtXU2Gd!6`ngY9*=+X2!zx}UKfEl=T7W2Rf|AG znQ!jB=iGbGJ#)_auD^f!+ZUsyNF*r1-;dA#x%AgP526|JuLpZ~wd#IZB2ry11c*c> zD%XOApm>K0A@L3u!thpV;d-PH0Uee4R;5CeYSa(!E~?W2yt`?ThTz>p!!!c#UK*uc z@Q%@L+5_)69i}lFUnPYEP0%EaBxxT_trn$1A5GJK7)TjuIir@a-~7@)a|T#*Y|BWGlc>w|P?gRBrRRK6hPM|hb0 zjJ(BvLypN46I_;elhDJw+(k$aKdjs(E`Ow4h^4?SmX6LLaiWuM!mH3}Njw%D2IVO+ zi-+{Z+$c81D&5%XLD&uOIc$u8$}awizIak>nT9c!EaYW2iP8fIJ7*mdvk*Q5kPosv zwqHQl%PWE0nQ?4RpiK7`5vE`qxc+iu*|ZJj`V3kzEcRJc6=p=Rg&XA+whP;lWmnTIAM}chi6)yQ*M%Ho$d=_27-n_5b&HX zxxt38vj-p_fAY{cnm~{d;s`Gyyo7*X*wt*qv>MD+YgPWwa9_3KA@oDZF0uWn8b;^^ z$OoSC8#|1m&m(j^=4v;LYGvML9nWpbp6A3j!;kF7`MaSwndTov##^tzKp8D;|AnC8 zgGYhKx&#Ncr3KRoA$P#@J+bw?5s zVykxvjQMfQ_JJPgQ#}QU@|6~8`W)Z7wjnh&5NaR{EUN57+wW-iCD_j2)NOq|U{q%w7fBQx3)_2^w(c7*5RbI+|u zt8JB`P_1g&}kk|M>68p$~{`2G{znI(~dkxlr zIP%Xm*o_jSP^sVKA0qqHJhj(gis_Y8Nt|F8WL}l!p zMR4G+aeRb-vHxmo=WI5MlAX$b6v_bLf1hdJS~9$R$hb+_7Vk-eq9nX{Zdl4XIz$7pCP zoVToKn&bC~YuztMO9d_6(*lOkZDniRcYKFH7Byc(5^~gaAeIfd-d#(A zddN{=RB=>Cadby^{Bu6I^l4}L`BdTnrj~`@yaKjg4mXKI)};;UfqYH9u;qQiu)Kj| z#7$%ZFKk^H$1!k;SNu=%siBzODg7@GRDZ1`Gx+Qasc0s)%IozJphsO^ER8rP{Kwsk-EJ%T<<2)VraGtIlnq zQt;PHH;YicpURU&`zZgH$zWjHlC77+q`^uJEU0+6-5NkU<)%5!~CYJUAJ0xsl4cMR~>~S+A%}n`fK7q1K|w962dhE zhF~FJWETPkkQ}v&RmyA}Tt;{YK^TS^5(UF7{9z!GC{e|)f_}gh3Bo@K!*hE-#>R{7 z*qAt4Ym`ehYkH#GsMSigQJgcq!`{S&y@bOSZs88#+`jf)_+8aW!6;jK_}knuGAar< zDFR8H>nzam+aBRJzHjJDEpfL!fhxjn81Y{D+lwvX9>Mg&TV7cJ0r-2{TicEjK3c$u z!ewGk5xTBCJ1>4n2!}z&2%E5;^ms?>IBt+||M1&`qSOeb?Rs+;*A>T4!ts+Y5dXGOEPPM)X{l}Pn-XM1KfybZX!23Jly&X#I{?6_l z=(@6X0O%TP+}Zq)8^Uj9lx?_gSb`dWQNCg zj$Pl=mi7C+8vXrSUCR3VuC~S4ovXvEtEP)GMt+e{NkNatD<88d1)^^^@J9zLCGXuPnhse^w!@P^upw-E{PwC=z zUgN92U}lKc2(43LD`*9Hl=r-^n_X6j2lBcV=Fz;K*UfHgmDOc+Uu0a1X!!pu!egJa zi;O-@Kcc1Ma*W+#)ABDiCCBw7^Da;8?TiKFtm8dq%0C@PI=ji^vH`zYh#4%vL~Qv` z%SPr(km(;yj=EOMubc;5aZ*J>I$Q&;6vQO$Y-I^qa9w68`L`>^~z55c7$}BolNWE8SG2|non$0 zCje^+(q1tLISm++S3L3H2*gqB8MRBs<#(P)avjX-IX@&^ns$f*$hD{BeNTdo$*4EM zwk&V)9$-)PZ#xcAXnX_z>)#6xsTZns|8B4+`L#MpEd zfDQ-&o&{_IAltU9Vl9?S!fq)RrXCy9Y4L0&E5<XfNI0S&L zMma2(VqG2cq?eIVgBcFPA&$!1vHdC3g6$~qbK<2szups zUdK)&D0pf~kqIc|&c(|4)0QgXbj)@v6c%eW`)J#Kkmw3DSFYGDZzfi?s%#A7wh{Sb zB1!FeFEJ2P8AnP?cdnG(RmfFrS4nV7Oo(wg**9HPGh#0Y-n_`3p3kUSe?lz_)pdJv4nak!XC2{^I+AOyaXw2soC2Ee)@+WK4tX6F-Lj{bWJ(jDxFeT-0gZrv#ZtQ@b z1E_(b=F1RN^^y2$!2y^r0;x&+gavsHkO!dV#YsRGPynbiuaaczv-4^^i&LcC&}-iW zKRX`zGPt~=$?wx|gjERChd1I-mJ>sL?SByIfx~Uxq9SL8qSd!>5dF4-p}68Cs14L} zP!XMmn)a;RhvDsJz!eRWxX`SgA=TE)>nJ&R=3qK%w;~r-u2Lv1%6mg?!9h@*loCp< z3e$+(Iov-P!$lhaTL8NOvw#wy0+MygM8YdaIh# O@EcJ>$6WEzFZzF&SF4Hu diff --git a/src/utils/data/__pycache__/dataset.cpython-39.pyc b/src/utils/data/__pycache__/dataset.cpython-39.pyc index 73770adef43ec478541b9f7fa7eccb8ca966874e..2aeae27ebb6f9b93117891a2a42f38e6744415ec 100644 GIT binary patch literal 3199 zcmbVO&u<*J6(%`9c6PP9791n7W21@FqQ$h0);aX%x=L%fXb}i$oTA8}?O-&VU5`9F zJ0`h~C5%pMqwOImdhbCx>VN62=!Mtj;&X4gwBO@MidHs?ATuEO_((oJzW4Y&W~0%7 z;rrJ=zszoLF!pciT>LzAwo&vSQAs9w#S&^o#pfasyyJO^@5X{8bZsw*q$mC7Ea}TY zhG++}CnK~QvM&c{hreWULk^!aIb4ZkB>6)Yk6vOAwi9#ft0I#*r|#(&pbb%efubu^ zhV5}9jMuW20PRcug*fDf9R(<%K`UBsDlqD|J~RtRsX&_jrg%wOt*J%~GSfp8P1ZCp6RCxTVn_Ml&}x;{vjPOK zPMGR}`PFBW-|9l?$!u>wt0(st`Ep*=Mo&Igi@hQTQ59ON$wSoTqW+}F-*5KKLgkMp zTIG|iE!QkfS?P()Oy(+Gn5|~tMt7@hzALkP@G1Kb3Oo2t4x^ePEo)gk+qzpV@~qPL zKy(x14J+!#MpLzzm!?qGn`e#ns-m`jwWwz{R7C>;Hq6{fY@cSC#iG(S$RA~8ZAY{= zC3d!36?T{}s^z>+waN0wabP3e9YjN5QR`&mw4avpS<!Ss*gz|PfU~gar~dYaRc}Mqr|RZq ziQmYM>K!bLeRUg6XT@MzEi$9NON{SPwSF#VMGA|~inrY|e>q`)b33kkmIK_n&Hy8z*^qw;H5ygx2Z?Y99oB|y< z+9MQ`p`(3A_Z%J6A7f>Nm47!r@Ghj)Z!y_l5n_esyx^SIYePA}j-u_k)l zI~l{9OpaFKueS~((`%!n=$JI>{|*@M<6n84)tetM{f|40yZ64fL#TWCWH(Pw3RaJJ-Hwpl;WuBRX$iVoHdc!f(b9XrART^IB#UcsJ9CDIDz4*MT zvi%ZpB-E~=pmF$mFshdhdFGm=SgzteE|*bd)5Qiwup$>(uX$V%qGi(hi>9d6d&G4Y zrz-0r9@q^6Q{ZSSOI-TXjvaT}o?SsZa*CdHlPi}FlZ@97k_@^Dvdm4=cVk3aKDGUG z-7iJ29lDr1)!-KN)bFCAOBn(4=pTc+hCk|lvB`hP)z2^|r->Io%23-Vnj*Yl2{14b zNXEdy=R85OCJaOx4`mN+AR{<)U^hQp&L}Ar@_uKO2RY@lvyU)>GP8^EUxR3o_)>}q zmJ$kb#E;*4YTI~6$fST>_CiQLMfBqipeMx(kK&Ia$i<_Nm4O7D#$J~?)dvt8hYq#u z;5=BgLwiN(&Gs@i(+dSSb-67Lt*?t}+5z2%n30TWINjY(k$7f+K+=~GSYEbxlp;Tt zCQWJA(;O4>P(;EP@506hH!d4w>ug0%pm_EX43WjDxaUgC_F$YF?Du$z9O7Ez$pIrs z)_N}p3UKU>({|=ydtz-o#D&G`zc9c+xZ<4siA@0_@Lkk~_vxC3Fk#pRQxB%|cih!j zgCo+8T=xcPbRE?Ufb@q=0--;%>(JyT`yG7=qI^@#n+qvp zaF>!^h=CYYz`MYug%3N$BYeC;b6GzzIP?y394CruMim$C!X~R_rqvHY!8)X$#T_3UO>34%aKD5@w@r>ar{3Q|#vP>a;&vQ^6gF+#G2*c(QgOt!YE zLNqjoaIYd?d!P~*uKWXf>8(BWGRNLH_r`_3=g?MKTYg`De&2q3?)lQV9E}DHSNrG7 za;tYdUYmS@oMswoR-9lcc;rQqXs*RmR-`)7DdL{)=|1BAS4>NtJ!LvO;)T?3hvkFU zWOj1=^KYvHi5$Armd%bq0_qpc_E7ldyYM4?9)5+3kcGb>KoNf6mp17t-rGJ>-NLpt zr8lZ-wArUSH%7o<8P3BP7GM#fJsW=K6Uf5_v2`=&lQ-~jo^(lFId2!ph^QLhrEXSN zo+8*ku^&B=6Dtv%A|eaV#Rq+S%xAg?e~5F3Iq&4Os?DtQ5A7M!nG0})@oD9ZbJbqs z7E!~I27^-=o4t(GHW2@ArHIMhmw4QL1fI26?^SjLPax7f(%b{e#o;Q`$r1l)?TPTb z70<*kq?p{r<|Md}1h8gtgSoG-Fo^pe!WLTICR2fpNGAuGw`J9=`+a{Py>UKLqvp}B zs=V2C#CBlwB*xk$do$D6_b{+dVrHjxT{R|8yP+<9>5NypvLWwZyxqaH%l@-6UzLSe zm3MEJDc3k?*f46@M_|`cwtc*`68`Sr+N2rzcTs`p!~)$$0<&R?VL1RZ7cI9pJbxxy z5gNI&p0uJ(Rv5|!Y9*Rs&SGN8;On=z^7N!>XHj4W;gvkSm3E0U^)$+hvD>p=J$S5Y z)9he}uc97}i=IzNDjdPeD@CPoG*v8I&aREA*W@jRdp%6Es}GYlo$@>E1x#`t4Mtaa z&t5}4OT{RJt8{?_s`oqL%=cYdmk&%m?{lcsj`7qbrE+}cXmA)G<7kcf!CLrsuojkw zA8y#|C`Ml+MA2amC%)!JwXY~mr8dIu@SDR;(*Bgd4Fdn`Mk@Y`Pw?Ubo;cAG216+= z$@0uDA?i}>|Ln1))^@#F+33>Q_mQ!&+%)^K&+*N?gG=#y%g=WA^rPC`q(2hNDgn2U J1KE@Pe*l`J_|^ab diff --git a/src/utils/data/dataset.py b/src/utils/data/dataset.py index 6aa046a..dc0b773 100644 --- a/src/utils/data/dataset.py +++ b/src/utils/data/dataset.py @@ -2,6 +2,7 @@ from os import read import numpy as np import pandas as pd +import pickle as pkl def create_index(sessions): @@ -25,13 +26,24 @@ def read_timestamps(filepath): return sessions def read_dataset(dataset_dir): - train_sessions = read_sessions(dataset_dir / 'train.txt') - test_sessions = read_sessions(dataset_dir / 'test.txt') - train_timestamp = read_timestamps(dataset_dir / 'train_timestamp.txt') - test_timestamp = read_timestamps(dataset_dir / 'test_timestamp.txt') + dataset = dataset_dir.strip().split('/')[-1] + if dataset in ["gowalla"]: + train_sessions = read_sessions(dataset_dir / 'train.txt') + test_sessions = read_sessions(dataset_dir / 'test.txt') + train_timestamp = read_timestamps(dataset_dir / 'train_timestamp.txt') + test_timestamp = read_timestamps(dataset_dir / 'test_timestamp.txt') + elif dataset in ["tmall", "nowplaying"]: + train_dict = pkl.load(open(dataset_dir + "train.txt", "rb")) + test_dict = pkl.load(open(dataset_dir + "test.txt", "rb")) + train_sessions = train_dict[0] + test_sessions = test_dict[0] + train_timestamp = train_dict[1] + test_timestamp = test_dict[1] + with open(dataset_dir / 'num_items.txt', 'r') as f: - num_items = int(f.readline()) + num_items = int(f.readline()) return train_sessions, test_sessions, train_timestamp, test_timestamp, num_items + class AugmentedDataset: def __init__(self, sessions, timestamps, sort_by_length=False): @@ -53,8 +65,12 @@ def __getitem__(self, idx): label = self.sessions[sid][lidx] times = self.timestamps[sid][:lidx]# - self.sessions[sid][0] temp = times[0] + print(times) + # times = [(t - temp) / 100000 for t in times] times = [(t - temp) / 1000000 for t in times] + # print(times) + return seq, times, label #,seq def __len__(self): From c3771c4684af88f2a64f92ab721314122d3e447b Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Sun, 23 Jan 2022 10:42:45 +0800 Subject: [PATCH 53/57] new --- src/models/niser_ode.py | 106 +++++++++++++++++++++++++++------- src/scripts/main_niser_ode.py | 13 +++-- src/utils/data/collate.py | 1 + src/utils/data/dataset.py | 24 +++++--- src/utils/train.py | 19 +++--- 5 files changed, 122 insertions(+), 41 deletions(-) diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 3cd671c..9e8a687 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -8,7 +8,7 @@ import dgl.ops as F import dgl.function as fn -from dgl.nn.pytorch import GraphConv +from dgl.nn.pytorch import GraphConv, GATConv from torchdiffeq import odeint_adjoint, odeint @@ -30,14 +30,21 @@ def __init__(self, in_dim, hid_dim, device=th.device('cpu'), gnn='GCNConv', bias if self.gnn == 'GCNConv': # self.lin_xx = GCNConv(self.in_dim+self.hid_dim, self.hid_dim, bias=self.bias) # self.lin_hx = nn.Linear(self.hid_dim, self.in_dim, bias=self.bias) - self.lin_xz = GraphConv(self.in_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) - self.lin_xr = GraphConv(self.in_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) - self.lin_xh = GraphConv(self.in_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) + self.lin_xz = GraphConv(self.in_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) + self.lin_xr = GraphConv(self.in_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) + self.lin_xh = GraphConv(self.in_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) self.lin_hz = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) self.lin_hr = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) self.lin_hh = GraphConv(self.hid_dim, self.hid_dim, bias=self.bias, allow_zero_in_degree=True) + elif self.gnn == 'GATConv': + self.lin_xz = GATConv(self.in_dim, self.hid_dim, bias=self.bias, num_heads=8, allow_zero_in_degree=True) + self.lin_xr = GATConv(self.in_dim, self.hid_dim, bias=self.bias, num_heads=8, allow_zero_in_degree=True) + self.lin_xh = GATConv(self.in_dim, self.hid_dim, bias=self.bias, num_heads=8, allow_zero_in_degree=True) + self.lin_hz = GATConv(self.hid_dim, self.hid_dim, bias=self.bias, num_heads=8, allow_zero_in_degree=True) + self.lin_hr = GATConv(self.hid_dim, self.hid_dim, bias=self.bias, num_heads=8, allow_zero_in_degree=True) + self.lin_hh = GATConv(self.hid_dim, self.hid_dim, bias=self.bias, num_heads=8, allow_zero_in_degree=True) elif self.gnn == 'Linear': - self.lin_xx = nn.Linear(self.in_dim, self.hid_dim, bias=self.bias) + self.lin_xx = nn.Linear(self.in_dim, self.hid_dim, bias=self.bias) self.lin_hx = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) self.lin_xz = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) self.lin_xr = nn.Linear(self.hid_dim, self.hid_dim, bias=self.bias) @@ -83,8 +90,13 @@ def forward(self, t, h): x = self.x - if self.gnn != 'Linear': + if self.gnn == 'GATConv': # x = self.lin_xx(torch.cat((self.x.to(self.device), h), dim=1), edge_index).to(self.device) + xr, xz, xh = self.lin_xr(graph, x).max(1)[0], self.lin_xz(graph, x).max(1)[0], self.lin_xh(graph, x).max(1)[0] + r = th.sigmoid(xr + self.lin_hr(graph, h).max(1)[0]) + z = th.sigmoid(xz + self.lin_hz(graph, h).max(1)[0]) + u = th.tanh(xh + self.lin_hh(graph, r * h).max(1)[0]) + elif self.gnn == 'GCNConv': xr, xz, xh = self.lin_xr(graph, x), self.lin_xz(graph, x), self.lin_xh(graph, x) r = th.sigmoid(xr + self.lin_hr(graph, h)) z = th.sigmoid(xz + self.lin_hz(graph, h)) @@ -100,6 +112,8 @@ def forward(self, t, h): dh = (1 - z) * (u - h) + + dh = nn.functional.normalize(dh) # self.x = self.hx(dh, edge_index) return dh @@ -141,11 +155,11 @@ def forward(self, t, z): z = z.view(z.size(0), self.hidden_channels, self.input_channels) return z -class SRGNNLayer(nn.Module): - def __init__(self, input_dim, output_dim, batch_norm=False, feat_drop=0.0, activation=None): +class GGNNLayer(nn.Module): + def __init__(self, input_dim, output_dim, feat_drop=0.0, activation=None): super().__init__() self.dropout = nn.Dropout(feat_drop) - self.gru = nn.GRUCell(2 * input_dim, output_dim) + self.gru = nn.GRUCell(2 * output_dim, input_dim) self.W1 = nn.Linear(input_dim, output_dim, bias=False) self.W2 = nn.Linear(input_dim, output_dim, bias=False) self.activation = activation @@ -181,6 +195,35 @@ def forward(self, mg, feat): rst = self.activation(rst) return rst +class GGATLayer(nn.Module): + + def __init__(self, input_dim, output_dim, feat_drop=0.0, activation=None): + super().__init__() + self.dropout = nn.Dropout(feat_drop) + self.gru = nn.GRUCell(2 * input_dim, output_dim) + self.W1 = GATConv(input_dim, output_dim, 8, feat_drop=feat_drop, attn_drop=feat_drop, residual=False, negative_slope=0.1, allow_zero_in_degree=True) + self.W2 = GATConv(input_dim, output_dim, 8, feat_drop=feat_drop, attn_drop=feat_drop, residual=False, negative_slope=0.1, allow_zero_in_degree=True) + self.activation = activation + + def forward(self, mg, feat): + with mg.local_scope(): + mg = dgl.remove_self_loop(mg) + # mg = dgl.add_self_loop(mg) + if mg.number_of_nodes() > 0: + neigh1 = self.W1(mg, feat).max(1)[0] + mg1 = mg.reverse(copy_edata=True) + neigh2 = self.W2(mg1, feat).max(1)[0] + hn = th.cat((neigh1, neigh2), dim=1) + rst = self.gru(hn, feat) + else: + rst = feat + + if self.activation is not None: + rst = self.activation(rst) + + return rst + + class AttnReadout(nn.Module): def __init__( self, @@ -225,6 +268,7 @@ class NISER_ODE(nn.Module): def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=True, scale=12): super().__init__() + self.num_items = num_items self.embedding = nn.Embedding(num_items, embedding_dim) self.register_buffer('indices', th.arange(num_items, dtype=th.long)) self.embedding_dim = embedding_dim @@ -234,10 +278,9 @@ def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=Tru self.scale = scale input_dim = embedding_dim for i in range(num_layers): - layer = SRGNNLayer( + layer = GGNNLayer( input_dim, - embedding_dim, - batch_norm=None, + embedding_dim * 2, feat_drop=feat_drop ) self.layers.append(layer) @@ -253,7 +296,11 @@ def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=Tru self.ODEFunc = GraphGRUODE(self.embedding_dim, self.embedding_dim, device=th.device('cuda:0')) # self.initial = GraphConv(self.embedding_dim, 2*self.embedding_dim, allow_zero_in_degree=True) - self.initial = nn.Linear(self.embedding_dim, self.embedding_dim) + # self.initial = nn.Linear(self.embedding_dim, self.embedding_dim) + # self.enc_mean = nn.Sequential(nn.Linear(self.embedding_dim, self.embedding_dim), nn.ReLU(), nn.Linear(self.embedding_dim, self.embedding_dim, bias=False)) + # self.enc_var = nn.Sequential(nn.Linear(self.embedding_dim, self.embedding_dim, bias=False), nn.ReLU()) + self.enc_mean = GGNNLayer(input_dim, embedding_dim, feat_drop=feat_drop) + self.enc_var = GGNNLayer(input_dim, embedding_dim, feat_drop=feat_drop) self.feat_drop = nn.Dropout(feat_drop) self.fc_sr = nn.Linear(input_dim + embedding_dim, embedding_dim, bias=False) @@ -266,31 +313,50 @@ def reset_parameters(self): weight.data.uniform_(-stdv, stdv) def _reparameterized_sample(self, mean, std): - eps1 = th.FloatTensor(std.size()).normal_() - eps1 = Variable(eps1).to(mean.device) - return eps1.mul(std).add_(mean) + eps1 = th.FloatTensor(std.size()).normal_().to(mean.device) + # eps1 = Variable(eps1).to(mean.device) + # return eps1.mul(std).add_(mean) + return mean + eps1 * std def forward(self, mg, embeds_ids, times, num_nodes): iid = mg.ndata['iid'] + # print(iid.max(), self.num_items) feat = self.feat_drop(self.embedding(iid)) if self.norm: feat = feat.div(th.norm(feat, p=2, dim=-1, keepdim=True) + 1e-12) - # out = feat + + # feat0 = feat + # out = feat # for i, layer in enumerate(self.layers): # out = layer(mg, out) - # feat = out - # mgs = dgl.add_reverse_edges(mg) - # feat0 = feat + feat_mean = self.enc_mean(mg, feat) + feat_var = nn.functional.relu(self.enc_var(mg, feat)) + feat = self._reparameterized_sample(feat_mean, feat_var) + # # mgs = dgl.add_reverse_edges(mg) + # # feat0 = feat # feat0_mean = nn.functional.tanh(feat0[:, :self.embedding_dim]) # feat0_var = nn.functional.softplus(feat0[:, self.embedding_dim:]) # feat0 = self._reparameterized_sample(feat0_mean, feat0_var) + # feat_mean = self.enc_mean(feat) + # feat_var = self.enc_var(feat) + # feat = self._reparameterized_sample(feat_mean, feat_var) + + # feat = self.enc_mean(feat) + + # if self.norm: + # # feat = feat.div(th.norm(feat, p=2, dim=-1, keepdim=True)) + # feat = nn.functional.normalize(feat) + self.ODEFunc.set_graph(mg) self.ODEFunc.set_x(feat) + # print(mg.edata) t_end = mg.edata['t'].max() + t = th.tensor([0., t_end], device=mg.device) + # print(t) feat = odeint(self.ODEFunc, feat, t=t, method='rk4')[-1] # + feat last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) diff --git a/src/scripts/main_niser_ode.py b/src/scripts/main_niser_ode.py index 934f861..0dcf3cd 100644 --- a/src/scripts/main_niser_ode.py +++ b/src/scripts/main_niser_ode.py @@ -66,6 +66,7 @@ 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) @@ -76,8 +77,10 @@ test_timestamps = train_timestamps[-num_valid:] train_timestamps = train_timestamps[:-num_valid] -train_set = AugmentedDataset(train_sessions, train_timestamps) -test_set = AugmentedDataset(test_sessions, test_timestamps) +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) @@ -101,7 +104,7 @@ ) model = NISER_ODE(num_items, args.embedding_dim, args.num_layers, feat_drop=args.feat_drop) -device = th.device('cuda' if th.cuda.is_available() else 'cpu') +device = th.device('cuda:0' if th.cuda.is_available() else 'cpu') model = model.to(device) print(model) @@ -117,6 +120,6 @@ ) print('start training') -mrr, hit = runner.train(args.epochs, args.log_interval) +mrr10, mrr20, hit10, hit20 = runner.train(args.epochs, args.log_interval) print('MRR@20\tHR@20') -print(f'{mrr * 100:.3f}%\t{hit * 100:.3f}%') +print(f'{mrr10 * 100:.3f}%\t{mrr20 * 100:.3f}%\t{hit10 * 100:.3f}%\t{hit20 * 100:.3f}%') diff --git a/src/utils/data/collate.py b/src/utils/data/collate.py index 27b7431..6c17176 100644 --- a/src/utils/data/collate.py +++ b/src/utils/data/collate.py @@ -279,6 +279,7 @@ def collate_fn(samples): bg = dgl.batch(graphs) inputs.append(bg) labels = th.LongTensor(labels) + # print(inputs[0].edata) return inputs, labels, embeds_id, times, num_nodes return collate_fn diff --git a/src/utils/data/dataset.py b/src/utils/data/dataset.py index dc0b773..24942e7 100644 --- a/src/utils/data/dataset.py +++ b/src/utils/data/dataset.py @@ -26,15 +26,15 @@ def read_timestamps(filepath): return sessions def read_dataset(dataset_dir): - dataset = dataset_dir.strip().split('/')[-1] + dataset = dataset_dir.name.strip().split('/')[-1] if dataset in ["gowalla"]: train_sessions = read_sessions(dataset_dir / 'train.txt') test_sessions = read_sessions(dataset_dir / 'test.txt') train_timestamp = read_timestamps(dataset_dir / 'train_timestamp.txt') test_timestamp = read_timestamps(dataset_dir / 'test_timestamp.txt') elif dataset in ["tmall", "nowplaying"]: - train_dict = pkl.load(open(dataset_dir + "train.txt", "rb")) - test_dict = pkl.load(open(dataset_dir + "test.txt", "rb")) + train_dict = pkl.load(open(dataset_dir / "train.txt", "rb")) + test_dict = pkl.load(open(dataset_dir / "test.txt", "rb")) train_sessions = train_dict[0] test_sessions = test_dict[0] train_timestamp = train_dict[1] @@ -46,7 +46,8 @@ def read_dataset(dataset_dir): class AugmentedDataset: - def __init__(self, sessions, timestamps, sort_by_length=False): + def __init__(self, dataset, sessions, timestamps, sort_by_length=False): + self.dataset = dataset self.sessions = sessions self.timestamps = timestamps # self.graphs = graphs @@ -64,10 +65,17 @@ def __getitem__(self, idx): seq = self.sessions[sid][:lidx] label = self.sessions[sid][lidx] times = self.timestamps[sid][:lidx]# - self.sessions[sid][0] - temp = times[0] - print(times) - # times = [(t - temp) / 100000 for t in times] - times = [(t - temp) / 1000000 for t in times] + temp0 = times[0] + # temp = times[-1] + + if self.dataset in ["tmall"]: + scale = 100 + elif self.dataset in ["nowplaying"]: + scale = 100000 + else: + scale = 1000000 + + times = [(t - temp0) / scale for t in times] # print(times) diff --git a/src/utils/train.py b/src/utils/train.py index 0b6c7a9..68c9d5e 100644 --- a/src/utils/train.py +++ b/src/utils/train.py @@ -85,8 +85,8 @@ def __init__( self.patience = patience def train(self, epochs, log_interval=100): - max_mrr = 0 - max_hit = 0 + max_mrr10 = max_mrr20 = 0 + max_hit10 = max_hit20 = 0 bad_counter = 0 t = time.time() mean_loss = 0 @@ -112,19 +112,22 @@ def train(self, epochs, log_interval=100): self.batch += 1 self.scheduler.step() - mrr, hit = evaluate(self.model, self.test_loader, self.device) + mrr10, hit10 = evaluate(self.model, self.test_loader, self.device, cutoff=10) + mrr20, hit20 = evaluate(self.model, self.test_loader, self.device) # wandb.log({"hit": hit, "mrr": mrr}) - print(f'Epoch {self.epoch}: MRR = {mrr * 100:.3f}%, Hit = {hit * 100:.3f}%') + print(f'Epoch {self.epoch}: MRR@10 = {mrr10 * 100:.3f}%, Hit@10 = {hit10 * 100:.3f}%, MRR@20 = {mrr20 * 100:.3f}%, Hit@20 = {hit20 * 100:.3f}%') - if mrr < max_mrr and hit < max_hit: + if mrr20 < max_mrr20 and hit20 < max_hit20: bad_counter += 1 if bad_counter == self.patience: break else: bad_counter = 0 - max_mrr = max(max_mrr, mrr) - max_hit = max(max_hit, hit) + max_mrr10 = max(max_mrr10, mrr10) + max_hit10 = max(max_hit10, hit10) + max_mrr20 = max(max_mrr20, mrr20) + max_hit20 = max(max_hit20, hit20) self.epoch += 1 - return max_mrr, max_hit + return max_mrr10, max_mrr20, max_hit10, max_hit20 From 24f6621f36ef197c202d1d09cd419953874b6f5a Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Mon, 24 Jan 2022 18:17:32 +0800 Subject: [PATCH 54/57] src --- src/models/niser_ode.py | 20 +++++++++++--------- src/scripts/main_niser_ode.py | 32 ++++++++++++++++++++++++++++++++ src/sweep.yaml | 11 +++++++++++ src/utils/train.py | 2 +- 4 files changed, 55 insertions(+), 10 deletions(-) create mode 100644 src/sweep.yaml diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index 9e8a687..e097c7b 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -328,12 +328,14 @@ def forward(self, mg, embeds_ids, times, num_nodes): feat = feat.div(th.norm(feat, p=2, dim=-1, keepdim=True) + 1e-12) # feat0 = feat - # out = feat - # for i, layer in enumerate(self.layers): - # out = layer(mg, out) - feat_mean = self.enc_mean(mg, feat) - feat_var = nn.functional.relu(self.enc_var(mg, feat)) - feat = self._reparameterized_sample(feat_mean, feat_var) + out = feat + for i, layer in enumerate(self.layers): + out = layer(mg, out) + + feat = out + # feat_mean = self.enc_mean(mg, feat) + # feat_var = nn.functional.relu(self.enc_var(mg, feat)) + # feat = self._reparameterized_sample(feat_mean, feat_var) # # mgs = dgl.add_reverse_edges(mg) # # feat0 = feat # feat0_mean = nn.functional.tanh(feat0[:, :self.embedding_dim]) @@ -346,9 +348,9 @@ def forward(self, mg, embeds_ids, times, num_nodes): # feat = self.enc_mean(feat) - # if self.norm: - # # feat = feat.div(th.norm(feat, p=2, dim=-1, keepdim=True)) - # feat = nn.functional.normalize(feat) + if self.norm: + # feat = feat.div(th.norm(feat, p=2, dim=-1, keepdim=True)) + feat = nn.functional.normalize(feat) self.ODEFunc.set_graph(mg) self.ODEFunc.set_x(feat) diff --git a/src/scripts/main_niser_ode.py b/src/scripts/main_niser_ode.py index 0dcf3cd..cce07f6 100644 --- a/src/scripts/main_niser_ode.py +++ b/src/scripts/main_niser_ode.py @@ -1,5 +1,35 @@ 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('../..') @@ -53,6 +83,8 @@ args = parser.parse_args() print(args) +wandb.init(config=vars(args)) + from pathlib import Path import torch as th diff --git a/src/sweep.yaml b/src/sweep.yaml new file mode 100644 index 0000000..c80e4b9 --- /dev/null +++ b/src/sweep.yaml @@ -0,0 +1,11 @@ +project: gng-ode-hyperparameters +program: scripts/main_niser_ode.py +method: grid + +parameters: + embedding-dim: + values: [64, 128, 256, 512] + num-layers: + values: [1, 2, 3, 4, 5] + dataset-dir: + values: ["../datasets/gowalla", "../datasets/tmall", "../datasets/nowplaying"] \ No newline at end of file diff --git a/src/utils/train.py b/src/utils/train.py index 68c9d5e..46da2f8 100644 --- a/src/utils/train.py +++ b/src/utils/train.py @@ -115,7 +115,7 @@ def train(self, epochs, log_interval=100): mrr10, hit10 = evaluate(self.model, self.test_loader, self.device, cutoff=10) mrr20, hit20 = evaluate(self.model, self.test_loader, self.device) - # wandb.log({"hit": hit, "mrr": mrr}) + wandb.log({"hit": hit, "mrr": mrr}) print(f'Epoch {self.epoch}: MRR@10 = {mrr10 * 100:.3f}%, Hit@10 = {hit10 * 100:.3f}%, MRR@20 = {mrr20 * 100:.3f}%, Hit@20 = {hit20 * 100:.3f}%') From 7238c472f7f7a3d23bf49ce07f59782e0e307912 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Mon, 24 Jan 2022 18:32:33 +0800 Subject: [PATCH 55/57] src --- src/utils/train.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/train.py b/src/utils/train.py index 46da2f8..93554af 100644 --- a/src/utils/train.py +++ b/src/utils/train.py @@ -115,7 +115,7 @@ def train(self, epochs, log_interval=100): mrr10, hit10 = evaluate(self.model, self.test_loader, self.device, cutoff=10) mrr20, hit20 = evaluate(self.model, self.test_loader, self.device) - wandb.log({"hit": hit, "mrr": mrr}) + wandb.log({"hit@20": hit20, "mrr@20": mrr20}) print(f'Epoch {self.epoch}: MRR@10 = {mrr10 * 100:.3f}%, Hit@10 = {hit10 * 100:.3f}%, MRR@20 = {mrr20 * 100:.3f}%, Hit@20 = {hit20 * 100:.3f}%') From c11d1b3198ddd928b3d6e12d1a9774d589fbec5c Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Mon, 24 Jan 2022 18:34:13 +0800 Subject: [PATCH 56/57] src --- src/utils/train.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/utils/train.py b/src/utils/train.py index 93554af..8902b3b 100644 --- a/src/utils/train.py +++ b/src/utils/train.py @@ -115,7 +115,7 @@ def train(self, epochs, log_interval=100): mrr10, hit10 = evaluate(self.model, self.test_loader, self.device, cutoff=10) mrr20, hit20 = evaluate(self.model, self.test_loader, self.device) - wandb.log({"hit@20": hit20, "mrr@20": mrr20}) + print(f'Epoch {self.epoch}: MRR@10 = {mrr10 * 100:.3f}%, Hit@10 = {hit10 * 100:.3f}%, MRR@20 = {mrr20 * 100:.3f}%, Hit@20 = {hit20 * 100:.3f}%') @@ -129,5 +129,6 @@ def train(self, epochs, log_interval=100): max_hit10 = max(max_hit10, hit10) max_mrr20 = max(max_mrr20, mrr20) max_hit20 = max(max_hit20, hit20) + wandb.log({"hit@20": max_hit20, "mrr@20": max_mrr20}) self.epoch += 1 return max_mrr10, max_mrr20, max_hit10, max_hit20 From 28ed106caa3f6a92668b9732d79370f71f3f88d9 Mon Sep 17 00:00:00 2001 From: SpaceLearner Date: Thu, 27 Jan 2022 11:45:41 +0800 Subject: [PATCH 57/57] new --- src/.DS_Store | Bin 0 -> 6148 bytes src/attn_sparsevd.pdf | Bin 0 -> 27361 bytes src/models/niser_ode.py | 9 +++++---- src/scripts/main_niser_ode.py | 5 ++++- src/sweep.yaml | 8 +++----- src/tmall_hit_solver.pdf | Bin 0 -> 20382 bytes src/utils/train.py | 4 ++-- 7 files changed, 14 insertions(+), 12 deletions(-) create mode 100644 src/.DS_Store create mode 100644 src/attn_sparsevd.pdf create mode 100644 src/tmall_hit_solver.pdf diff --git a/src/.DS_Store b/src/.DS_Store new file mode 100644 index 0000000000000000000000000000000000000000..5008ddfcf53c02e82d7eee2e57c38e5672ef89f6 GIT binary patch literal 6148 zcmeH~Jr2S!425mzP>H1@V-^m;4Wg<&0T*E43hX&L&p$$qDprKhvt+--jT7}7np#A3 zem<@ulZcFPQ@L2!n>{z**++&mCkOWA81W14cNZlEfg7;MkzE(HCqgga^y>{tEnwC%0;vJ&^%eQ zLs35+`xjp>T0h54l z#3*VAG_N{{~6e+X0DK|GQQ+1OiPR?TFa_ zW5O^hI@%knm;#Bkzw3xeeCIcH0}?Sx*?!j)`Jbcc{~V==bcq;MEDTL7?aYZd{u{2L zU}R-#3?$%*nvS!Oro$)ctPN(D6I${}Q2W>SXWiX#CC0e@6d*%>0L}|3%Ax zO8j3QB`j^eCnsW*u=%D&%+%Q4HQwS3Fv5QXba<((Wxn8kH-m@ zb@vW~zZPH`>!H+4A$LY11k1%dBRV5`#%1U496lrz@VNXm&&OWRx+({|5|LTKg*G@$ z8LrBpzy9&?`FOdw==5-L_UK5kA<$GH(2zX+YeSGGUts$UdyBTwge8Z+;ddMbhEdE; zE}mBSDmdZw-H*dXCf2XbCc?F)_~@ZSYzE6PZ)#u;J(*K@8qY4DYXryS1xB0aP^NHN zN1Hkm^49dnwN*}&PaiO!0ZtEMX(+XV0qSBv^H-afMfO(7_?1+x66X~UtKDXk{en8Y zSEaWSC*#_c-RhY#oaQV*@ePl=gEVB+TJqDz>pUxCAPq7(lXUoEq!z@b@a|*p=^c!l z;b2tWRZGp}ec^rfSYqb|VOsRg`p(s9+0%;P$;NcDv$WST8~l=WsOZa*%umRM#9_+VSP5&XK5+lKn$HEvagX$QCuiq9n^Uz99SCOG?bAL>;tTs z=%Y|7k>SM$Zz{-?Gs$D8RQ+B*O4Z+u)={bqhVYaAq#FTd<;TRvEfiMz|Ao*si8X1w zZ;!~BjlL-Tco4lSdx2Ea9TLVT7quoG5SGGhRLj0xm_8-N@G5PQ3C&DlSM1mtBmpI7@*c9!Tl=}O>)b{%}XvQy}F8|J2eb96Ehka8I1;7EITdb zDF*v)+3CuB^l33{BFbz@m|Fza=>?D3v97L7zwil+Vb;0Ee$=Cf#H;>&AWD-g(!O2%{BJ>9jj-Ixx|6PdppDBRuHF(UP&QHRV2!ZJS{Pmei&H^BiW z*Tzspeh-In4Pb5-qRR)?6D(C4_~(!Pa+I_LN1^6gE{uSX?LdGGR}L?|6~`Z>Bq3s7ctIhJaYwXmq3E z?*w5SgE3|9a@oyhzD$?jHj7~A?EbSKko;V;LGJ%FLE+TP4YQKGib**9v7Z+!MxN_+ z8KU&1+!@+#Nieg3mnNz?pes<{YA?{}+*-3>P;eIS$Xl zxHaw=wzYv}Pce?d>Bt$f7|;EVV+)sWtA%r=leNc?@K5?+-i}Vv`9QhF7V&EL&I?BN z)8z1BSE~iV%~|$gRpi&`uFhh_9aFm?myj5S@xWZgPX!e<#`_UVUO#@RfzhntFh_UF z+&|^B(>SWWL@@aQS79jRM~xUCT*^^9Gaxeu^~_Jl@ecbQ@!EqBvk$08%#QpBEd|n+ zdVd&a>6gSercJmN7*!|25TRI?s47Cs8jy?4X!GL3vP8aq&xmr-!h%>3%z}k#RwZgn zOHWkINcRQ5!Wv~^ao{OQv6?S6=e%Yh|MF0b#C8nfk$xuZWtGV85U_v04Yn&jQapOZyU_bdD&1F=N+rsL zj7dyaqVSJ@AD`U{iNZ9Vn}k3*Z$hjJWza@;i{e%+Lv58}YW^syTZ&GVW+Ef(3*7{u zTK?#R%`Kil)Nc0$@vGEB{XgRUKYH^&GX8&b=s$wR%)!d}f3^Mp2;C|4rin|iFTGRxC@Xv_Y&Nh}}eJ@QJG2x*VGU;$7fQLL*hCBecE|DL=v zkQR_1cNce68T0*_JA#nUgXSq;*B$8N`?JSy#UQs=kbgIIXX>xV8Yu~pLGYeHN?aH+ z>KmXka(ZuXyqV}bc26Pd)sugNCGf+~xY^9U(^j|TId{Mib;9}THyqD#-^FaaDlsk* zioU?48N8GiLGO(-5$FI9>dLElD+GKm07macv`G``cj7N=>PGHt-8W<0>JOpq1C&Yx zYhbt2hdS=1?~g0vKdM^s=Lya=>wwsL%4#3fMzn&DW+fb-X~~K zPQn#OSD>x**aYY_+nC?}yAAxX;_p`(}6E+(RHmF?1+51$>B%Pl0UiB^esbnc%4BPa^xe8x{HX5 zm>2vxm_d1Fh)=O`gkmhj_#(V4#MoFHe4%H+AKtYeqhW2 zxwsi&KK@WFNUYdqez>!+1Y+1Xyw}1Md4wIIAb*hqLce^Go{{>uB`3^MzKBFp6bs7)}$MkD`3I6xAwFwlS=3^mRn zyJ1%qrZ!?vAdNP~3V;~^g!%;BD28c(86%;WNK}aP!J`_`8Q~8gof+a8f$i$&B3mJ) z3khe8IG`+0e1gUsf*P?d5H|_&@1j}!%J7HHL9RtlLt4Yq`-zIG-z(094WQ|_oQdSZ z*Fw?Hi&hZM7_GtALgvkZ2Vm%=JcF9Cc3?4xoClVMCigT99L+$?Bexqu?&76{abe!D zra)bU-Wg)-7Bb-iKz563kQgL%z%xUWnRVcF!xXk1k0m-l4~g_fug#e_ZaP}^So^9g{Jiu}G zG=@Fxm}vyL!M=mu{0_#=6}%y^@jpZIfmkNp9J*e10NcHAFy5VZfT|GgK=lrD!>kx= zzYI=idLjt!_Hzz>>Vq73+cryR&E7FJb_EuZ_|ljQ9Fg=g#5SyX`$vEAHu6C6j?%zZuW^`Q-xan1*cx)k z4)lQe&gek;HD3Rh{}zb=B47X$Y0DSWkaCyM0CQKS2KR=jKiN=W*O?3a53;_WS7?ag z&l|OsY&ZDm9*}{2IQ@WPB>j+LRLLFqXa3!ExLWj3`(B6LpB?biBm~CK2=AD??bn{W zYR{Z`^l^zkpY$ZL58kEAzMa8)}^4{oU#RGd?-??wh}% zlUyIRLnrpE@;(uOSMA98t?wo<_U;8)A=~$&ySpsVUZX7W`t_*p7h&bk{9gb z`uMqbul`;8dt2>0=(|)R;MH*GB;M;!2zxhr2lQ8e`CH!G0N}` z1-uduoiKjaO9i}+f7dh1@cyPsc<2P`JIQy;b3UZhh(c*HAB$PEq9iJRhwZv!Cz72y zTMEX`Dq$j|7^h!X> zO30Z`Nbtwc7BJ)^Xos>`A!}EuTZ>82YfefimWnNMW7q3dFPF~GE#yfc!KJY&RVuAj zP$a~vt}d9e8^N+$YUrm?rOnlFL=9Os`%{&P>z>x7Pf`S6(5hOC&*sgSvbV7{k3+^4 z#KM)1n>K36@&pJMFG-yO!Ip}aP22DVV*$l$6&MM#*7sJ8_@)gWof#*fjXZL;(z>Y# zA(K{2sc@DOQA37omW>htjI;rxDJtqo^OnunrIAwkfK*wn%u%FUbAO5aMGckYtqT`! zSqz(a@%+((#WXqhe*!aB@LJlZV^b53<{I<17@cy(cmL+2D9Rg|7=Q9)DjogObr;Cy zoIkA5j&jMIvTshOd%TbHx0yMc^iV&g`3dQz*FJB`*#>O;%i3?_yh1B|+hb6JwO(3I z#>;!`)V3in6Fi%8=Avz5XYT9{f)(}K*zA;MdU71W#Kg=@gHOVEAl}m1jpcMG+3%kX z=bOs%vqb|D<@nWf!&^lXOZRjZu{j5)YQW-YU9gI}sJOV4A0yr)wI$fDYhU3M^l@Xtf{!W zrll01p`?aLFfKVg9v=Y$%scyha_qi^RYyLAf5VD7Jrc4L!8};24|~W7{1z89YQvFl z9UrsLFANpdf{zIb_3)i4mzTj~eiutY!`F#T^=e^hu6$3aN&@*Lv-3!CeRyhFz(({@b1oXLnbxd#Y^y*>ZK_9vEiR7uT|6 zZ9dp|s8jCH5=pg!cR_0i|_~NU8voSJD2wE{;VG^+v=&3k4H#jm5{V z9bQW01EA&DYRf<(xQ@^j6&wEc)$#&FzhemhI7Q{^Un;^GW1B~^OsuZpB=yLaB@1P{ zb8&A9%<;V-f{uqkOi&6JJi9)yOCQ=+$pEEHR40qlSY{RMS=N|JhX-Q84W&b-|HXA6x9%*a_c2;NsuX|;D63uPinu#Rka2-o|-m1br1k_Ln)^ze@Osvb}-N?szXt zmTq^eoR>b$3&;M~OVuq7g*Pd>RKY_}JsPn%?3En=aUefCwQrg0Y=ujj@0V)skWjzr z1}dP;5O;5>&1G}anM^bz-CR8PQ_r)@1jllDwl47R$i2_!-vS8;Giour$n-W!idxdi zqc=gfdBt6}&IotG^(2oXH|PE0n_9uJuLAeB2R;sspH7qWW|jKckMSK6WaSTMxL}TZ zua^%_!6+Y2V0F|r-hx|G&DB?y#{!)1qG4dT>o-H`yw}7Su~w?K2o&dVf72P(OaG=9 zO<8T(gZ{NoTM`&SuB!SY0_7!AWvBHPGwe(4A_)?G;+-<&(-<**x3sYO96vMZgP{E! zv%|#$BAa5=rwHFh07p4fiqiiM)sBNBoj`m~8Z<#TF-reWffbUKkGy+Na63y#W+4K4 z80x|<>F#qo5V}-oFvY|`T~B!QW|i&J>zS4(+9z8uI*GR3T3z07^hArF0_hXefD9ad z2jr7KG=EQMXKMo;F4@LC*azZH{}3^cSi+u{tAsPzz|UvVhwgcEy&%RoMeOtBhJWusA?8gX2man3)JItc&%<+~wzpfHaI{Fh?gqYCzT0X;C! z7HG#ZbdNP}EE}UN*;>CMf|?gu1&=>HC!Zr)Da@CQDRL zzm+UY2GVh~NGET=%1Q=J(0|bA1O^z?=GfK*BLgZUBIB53gRTWunY)y+R4F2tL!lb4 zQm;k@gBBF-7~TZ4>C3r4d1Hsp(R6O%UUcv)VrH?fz|Td~gwV zP6s%tPZzQ-D3iIeWK1wsSaiN(cD|u`mT-md&|Yv( zYspx!#vc#>J$!Lkk&AhpJXqEu_9jeC*quJ`L*5r&Y47nhMAc#IZ+$UL9#|CPsNw*_Dntu zOCDDSpNkr2_C(o>9?J%wQnDWp2A@3lcDNJWpq)_x#H&4n0kwSgm8o=rM3SaQ!lC)} zmQssc$TD4^A4U6<@BuIPRUTTsNTs-NX!*%<+AnO-Lb&BqyjTW^sn8(51ED}GX6=45 zZwWhoi@AjlxQhRAz-A^!;%;PMoW3`p-J1j{FklL}bttiCATU_a)M z6BHBq(^(Cfb=!r6N5Tr(+5aMHHQ`3#6H5B0gSR$*k7_)&a})mswo?I8<8tZ8;jW^6 z7`odP{KV}hYTt9tQ`>v7CF{1d=Pq{2My%lG8p_(GEgU@Z#BU>w)VRY+gnRdLCie-X z)CTMJW-czVbd*7srmKfT!>FAf)?@c$U^?8tOj_d@>EjQjSqoqbYn4k?iAI%+_G9N3 zF&d?870(td&lM#Vovaj#5y|DLkY9;a+eW&YP$0rVVbi(7b#8g7qf59mcmcFq6YL6^ zLc2+37Fis`df1i(aVVIB!(wTBr%nGZc%fSx!HfZyMZ}=Z19_gnt?y+IjvO&d(zEe+0g6Pj}7X`cL#1#!LE$d z8{YKdreVVA^U&lOA(Hxh?oBcJ;SU|g=7L;r!DZ3T=*JHe^5r!LWI0OOGv28Pzw9$% z=1`@W@C#5ubTGkxg9BPASY@T%x!%QF*=QcovUwsU0@KX7L8Ep48W2f~68va;bwnd$ zFpiq4a1X?qDL2DGC%#ykh7$NSeSE4G=2U?sFT^x;^(=v)yrPM8f8G%P`6;MOiIF`JRnL(SG>u{>f+6?Fxm%Aij*hRbv#(x+$oh zE-F2qqoUzzUecPC0Tz2e(KuS5o~$EdqkbQggB*i z$OeA6;CjB=)K>LE>0HUHufB2mrMP0HrGZ~*t$pV+VpsHU%|dswp~_5anZp#fy!<5( zx}devs~Cp@m;*BMuA9LDe3))bE`COR(tqN6$C~M^cd3Wj-h3@iGfLW$j9*)om^(8_5`N7d*B!{5Kxp(T14mCu z#z~zo8og6O+I9VT??7wvU~_22Ings=K4EX*frPnlygVRwLhtNtAu|_w)Se*vzDr*G zm9?ZIQp2Ok=iy1ok=P&o0-8?-i8O7AAw5gMhhdD z>iC1W1+_^CNj~1l4Iwyow0Y-uF}z&}1s5o9&kyx@(Uo+Nbwj>z&mHkidvGYFUjAzY z9XPLR<_w|oR_skm47y?vj~T5BY0u~tI8O4%d{vxBybA*x3!caHO$O# z2qRK>@_79#AJ@##6fS*Ya9*Hh>zUQT`nlMhSbFj^)ujPUm;}Vjank_%MR?pm7DKOH zM=aPs+c5>$qAJ@~i_DV*>-g_51SAM78C+SassWrJ1+)$n}e_20+BVyGxe8h%mBV4Tc5%NdR ba@ALD&%dOsw4QqjqvRQ z_c?lh=L+l+bNzH+Vjw>(rk37noj1}C))0fv5B=K74N3BW+l}B;qU(r_%3aJ&xKY`e z2Y~x%i$!6urmf^blhg?!X^9kk{e&S9Mz|}$SA20>COG>WyJrh#8hqtaeN1G9@LG(+eDm9M~J8;#)rhHy^ff0EC=Pm9K>q= zV5N)T{h_zf0ly+dTS?z80A&3Np_kPG>GS7RuR0f~kmt(bc+DfN(%I5q?HgA_cYi(^ zJtC-4sD16%z4zzw9hw%Ls%Yzd3`gs0S6&E5?ff=INfluj*aE&!lGR^xhD?x0 zTsNoP$DfVD?@Ktswg!A$4w6OjJ7X^z3`*&-YT+(N(jdrsz}IPqCIg^nMFx6pXb z@l=IbA29T7iP%AZdb9}GK}64xt~1yPK4I7&h>kQlJ?niEd^W#3y{EqmLdcI`h6koF ziE2=-6X7B2p+-Y*m3NduE|M&oEjs0C)1ekeL=3hM;_sU8YDy)nI<$LK-WMe(hWxGk z^quZJqTG7N43ivDR%>%Wcf~sjioa5-B%iMx7PG_gf$*XIYO@02o*duTt;_0;?B=^c z+8Xg1W+Y~9F%QjtDl2<|a9|}Hak&J$gnKGAAf0cSmWsZnt&7p_Th0}ibbo;FrC85- zg>!AQ7iDh_hjU`Q=Dg-4^M+b1o7``)Vtv9H>4Tnv83z2WKDg{l*E?X?;}B|BEP5qi z==b)Kbw+=Da=WC68}6pI366RB@#NZ2bt$&QB})M-e2G#OZX-E_Pps(>elLR93Wj!z z<{ksLll83hi9FYL;?S_0k)^uZ^u}~WYJY1v)BDLsGV$*hu4`KQlpsC9!T9VsH%xWE zbYDEx=~SyVns+ekwhIF;cO}TbeRyS1`4-rh&xzytgj%&>Rb0hsg6+EX9+1~baIx0*omDZx)R+!Ih?jVU_8IWbvia8R%g;LyPhr29xSHkq%c1j6GDnU~5&_ zsxYXXx&d9q5Mx)R-QWSpV`;v>Em-u6rD2k>KI&wUxu8pBJT2H|p^CwHfBy2toXuvB zK<;0JT{GM&g>vn{zJFk>3Fh&ln#`JvVO^lQE^Huo#? zcz`ss&W=27!#LwH$nwQq$FXT7w1B2QS+3BO{{8bxEYpbY>>5zx8)5XvsaHOn>A`3* zDf`$HB>H*TL}0lu<{%q$sXT1}w(w|SnBBaQe%0rA8tbBHF^POBvk6(Uy`I^)y_ohc zx`NV6;%gdIok<`=`nA{9>LgObi(^Ywp8b~ps3july?j8b4qQ?Z7I_dXyyO5qN?Er` z$!e^f%_PDsBzV4JmR?V4{PZG1OkQ}hbl0L}z)l7e?9;(r6T*+Au&Jh|>5)4)NF3%v zC2cfp`q`^Y+@KG?$FG;$3R~_@IB`sZzf-_C6bn6q0<2FuMWt_J$L~6<28mus3n`r-7iGx0`kJ*!k()E zS#s{>qS3SUpird5IvY}=MCG2_dfx}wASE?MKyOLPc>d7hklYYRz5$RM!9m!OWwAAO z1#P}GzJ!jwHZA3!viOAB^<-)Ne)_KhEX|;M7cJ4kLty)c;ay_P(U>LF%#olR#al&f zgkKpu1O7@I^Ay)~{vD@fN z^d9~Z^h$r|wm2-4d-98YILE2yF6y~uxSSnTw z)~H?w^Ga+JERkPuX7$Q0#yXHp!xK+qoF8EWxNrbefq7E9tkbvuEFQ-pX{Pc(60~V_ z(M;A**M>MlibX5;M@$5%BY7J4Y~~|!sym~;Rmdpz$fG*ZZ5gRL z?SOI+b=1pN7>dK$;69K&qDC!qK417{ZeL+m;9C4kzw2VHgm~H_+Jfu&{30}iO&d?m z#)M8yErXi6?oZ+w2IiqFmXWnwnc6gN*0B5L+Hx4F#ANKx9j}8Q2c}%Ad1xg=dKmm9Boz}qIDXiBH*lz%D0s`*=+a zo9$nNoSgJJ&)ZqE-+!VEE|jH**Ly-nn?qavL$nM-4QQi*_mV=MS6j5gnZs$E^xg># z-am{URU2y?k50&>dXr28H!b*G&r?%(eY5?!?bjC&Gz>N@>O=G1_vQYJQ!0~J%x*I= z-q@e`a=$8H2?c+5{PCJ5K652xB#A#&x}b#4DRkk2<;2M;izI>wpKV<-ya|5TcWq$k zhoqyrjwDdj1hokB9@8Nm1`?(PRCa5gymH8#a97~>TmlP=B8SdU;*KJ22+d;F?$93?BDfQ0ZENm zlq(~6UzrSTf?QJq?xo?MlT=!fV&oLLo)?8Ft_q#0z@cskX1cN03~rbPE!NRC*S~)T zFIcm4WW2`Njzj!ALcUX?JauzRIm=%W#z1U&5CTnNztNH7%3ZSK1sN zIL3BZrj|aq+AMQIniKKbv#@f7tFhb%f87EHD2R7Z^zu={a^_EkXHFoLiib51l<4Ak zbmJDR^!|$BuX1?g`Gl{Bym`?7%R5#RTPP7%szd?u>`=2nA4R0mDK_9l#ZqcN7!W9)?za4myAuUZc+cL{OAdi-lh1ou{7vr9Yu# zYKCfNYbd5+>X8K*B|*xfrB8WsenHyg%Wd$O7_WT)v)kutDfX!4KyZx*hs@j#{{vog zW?J9p(d^PMmSIfaNolKELVH89E}<<1!sukAjn2B!sIbdrc^6JRqK_=qQ%&{TCSdCr`!9CLi*x!G z=-W;Hne+R~?(SE@JW{_PwIF&4tmokq?e)r~J>JWJsm6?QmJEaBfwzI_{Gc2Wd`fEp zA-tItL5r$Lft1Y{JrMLA>N~0!mUkhUOy{w@Vp zh!e{{e)~nZOsTI>bySrsZh|MxHp1$R*-X1S9z>q>qRC&DC97r*bwU=WfKmb(k(aPa zNpw??fa%Pceyy(0)K(H;ajDCCHw2rk*Xbx_NJQFR;y!kpna-*6!iKC$A9FRi={Y}h z@aL_bwm8X+ghKz{)Z+4cZJW?;q$=+kfX*e)g$MBE@_SUfN;CG+ssAC(RU|t^ES9qT zIk8Xr4JQOi6jAOFR@OuQ8x}v)5iVJbOH5TCn!4BXx2<9 zD{h$DeZw8ho$@x2PrYu78McC=+$Ks3{fy$wSFNti(Pt$u<06B_k)=`6P}MImx^vaN z46!}-eI~iR(Si$50UIlL%TF9dJ|BOuBj)sr5Wy031c~D5`xaRyG+Eeba)Jmx%y%Ae zu4c*Hf?leC6}S#y^zsS8ppF^z?*x?DB=4@_*!LG+d(Zvc@ql#i#+}5V<*quzsVjTZ zp(_SL_rbXZ4IazGw)EdY=4BW`_mzmBh8rEIH^cVc@4LEWPK>Ess{+d1E>=%asNn=% zf8M@tTqs{Qnywzsauvps9+f^2f+9{+wzt%KjGIgR^tSD;63pN7@+uMR=(=YU{ASAG(ivY!BjHY)-#wShwE7 z@Md`%og+4^es#LQY@3G9Huuqsqdj5&x2PySlu&ZIm@__tft7-1qfS!>Q)0eozGU@x zhphx`L7dq}oE+w=?nRgBwG&=%%ImZP}$`;Jyci{8dw#%-fman2&- zKfO9sESO|a{OqT|!kl@8uZLU&N=l3=*3rvfm!tKTS|udxr>amDBoE4GkeW|IN3=2g z?{nyDdusGom!6N|*y{1Saw=&2V=BVh@5~md%P6D9&3OXeL7?Ji9bXq zE}|~K`AJ2ly>4yULtS>Y1--T^+pFc?(&gEuF1M4>qESYalLT&|=VB!jNKx!l!r+>L z_Jw*QUs>qzex3WO&OATuISRAzl=-#zMHu!L$IlDfHJi0_N%NnM>7KcsO35nMezFb^ zcIX(rh3kD2jYK-|XFG3N-og=glo}@2ps| z5;l~YcMmtPS<8m7C?6DdDfq1_lRlKXckS!}si~{gF@t86 ztrlR^FQ$j5jbDB^@Loe<0y|FZM1?gaOj1I~X^G9IEcT3Qi@oh&t6^L?SF~(RqAipP4>2CtN!9FGg&VaNBh6rK5XX7rB9$58~_r zPBJpr*HrN3F9(I&cv?orQippSR>-H^K*3`9q+1pStZUc|`VEqVF$4O8o{%Pu)=|(x za{s|08w%B7Y337F?Xk;}gc++i~}Z+#U0!B~wNlt7=QzE!Uf*7gUCxM#2la#gCRp&cE4L z`zwbXQnP2RCrCKbp6$a{2Sp)c!6@eFu<%KyGZifOI7aj16;wGImUW5G&I?ax7EL}> zH`rQu6gb8HpOIRqU&JPPXNX*} zyw?Ldn6em@?XyGR>y>zLL_Tu!+}b?b46xxuEZ)F&^<0T&m1#t?!>)4r>Zy%iz$bNC z_xO&}ZDY^Nd+JR74x&P@c@6WY`vydpQ0p@P6lM>>Pw;@to4p~+MoS$-VZY(M-`W-4 z@;C8&-XcdbF+M#$3qE~}uOVETqfG41oba4Cj;NMaKVQbzWRDL^fkliQ zretQP7yVJ%`jK{P$4T@_VxAPTJjMR!+y2{4eaS3sUPg zu}Jck^jjmyG8B?^t)OREXKG9Bvl6lnaf?^YDaTmPUO>3_j^)Z5h`G*7IQ^D)`-`;M(kK&`;Zqlnkh8p4Oflj5+GBgNf@xFNRXMAjqW zriGpB?V&#~m+i&d{>zvzNfZXps7=TU2uGzV)++5%4Z7tk?OtI7lEUWty6tiKwJmh5 zEsNE44Sgmum=kq$euGJgNqt z*5dUBsSjzM{m|p>A4y^nj`zI+trpz9>Qcc1%hqGi7-SzN70YGd3HicXF> zk$VZfIq?Xhp|OKng0&UU-AHb^zS(G=1-?!m_(8D#9l;hNJBp|c@#O1q+wA)|cBQVR z9eqA(7W|sDTe>xwHi#Fng1Ucyu8O@#I67Sqoo|6R$Az@o_TocKqv1t-HWDf2H2Y?n zI3#66T1Wt#4V{6)*L8UfeC$UzNa@SA_b1KzHmz!EWr@VLl zox;NdMQ0bJxGDQK$uVyrvfEQ#%l4okd4eAQ7M9glQ~AT8_d`4`Xo=>sicCt%s6o=y zoS-qoORl)f02ac|T8neIXt&|f!NaAc02m^Pa5=d71$3jJ%M1U|C*7C8NT1T2t=G=zx*SQN6 zHFx(_vTB6H6&^if;)5oILQA3Ne(4-GaXRH6Y3i5p+J0QZRr!pKv*ZkX^9PHpKZfJ* zi+3%Fc*Oa%ubnTsGgwzFpD;LO?lv!Fx!1?+-8RdB=6K0!{mwz@i=CSvA{kxX!QIr0 z^&@n4zYi-a#zzm<4^orOOCvnlyMAwNl(nzNlr^=IN1K@C)}?rQ!!PFg?!);evtloH z5!-3hTKxh?3&|Za8rWSmfLd)6)rpSFU8&D zKKqdj_Bfy(o^LBAnn7%l0Y!YKZfA4@8Ygj)pWIV!6z30W6Q^WHu zB&u*|!lbLx=TFPQ{#N()Ib2V`1VWynW9F;*)q5-NqIE}afsivrXD6&%Q#rz~*(^{s zFxuFaG8HUYQDe?5ARS9nb6V45*m;y@h8w(pY&4)UIY`vbcF~&F6{?}!75a;;`qV+% zQ_}^}5pxP^6{)akl*o@x6j%(=Z7A-_kw?Is0Q;;c0d=aB@eXMsn#y$gz9bC}S(2N> zrs-B^J9)?+A}?n9OqUV}^#vo>*MDM1?us}O_q4)!cuPB4 z>#98FT2fboR@u{V=j|47ivQWyIBlj?;Z$jVKeAHWx9(q89Gd=t{duh?m&hQgU6?&*LcDOk*hsCKB1k93S_3#rE+XktDIsuLZQ^(*SXlw=)O+->9uPi8wPk@;lDSi z&B1%K4YgOD@NaTx@kIlMm6@1Kd53XI?Kg<6l94P%RljM@O`Tty zMJU3O8R3M@EMhvR0I)v0E+Um8_DzaURP85AyO2%zoX``)#KPl3kODdomKPALx@LBa zT%=T9G!%qPEAdpAgHhQ}+1M#x-|Mnb#~amHQdlbqYBZwsIzyhpjMdZ#u<{w=rQB1O zcA4Endt(&ohxmHjlHv$J8>w#N%)|&!g7oIrpt!~%t?7UY%h-MbCjA?O&jls1l#p^0 zYt|)Y2{->vaH{e*p~cH99miH_pm-1$AK`4(a1cyQ5gz7e6h!p{x(?r0y)OHL+DU=HgXN-s0Gh_mps>NI4X=ANtRY6d!L+LT)A zWnSVsv-@p$<9Tu~n~BVaZi-B1b#K?j0fvm1uq#x`mk$f^HR4}-3RCZynlwBCazKXi zIp=K?Tqa7IYs&T$v1KLpd?`Aiqt@uLv(gUvBcN}OWre5(SIRs(e0bqn5B{8m7k&%O ze0*5i6Upd2G>5{h5mRSqd=uBHeCBIV?N5S;p0{KE`%St;8azhNiuler)57<&}^#YyT!W_rNqNPc)RUjU%2n2X?^>uW7vNw7l_ zbShePT#!amR+<)S(@o}|VQEfdx;&-|m;l1ZrY6#`TFQRN8%9*07r2gOfi3(hR1L3{ ztq!6|9xbzgv)C#Ur7DQMy>oA(P~4yqkyN(s`O^XlH_xOCF-5&1j?n>YLO(sP&Cj1I zHhMm_&(f(yCt+=_WrkTIZ3sJG$_Oh#N+RiG(ptilJX#a-ef7*piFm%fKepgi$(Tz7 z>b`Sos-L}B@ktRv>BC_;BuP|7MH5*UmO>0>ye*;%y-KYT^EJW+4IWcUTFD0)O8k*e z+oa8^%|mNnUMh_WxwV9)lG;{cb)%2e9ZY=AVfN%FE1-_0Ek+bk?2ptZRB_qhl3MyS7@yTSFd`=?b@ZogGWCTG{eN zGQNXLwpYzpHFtV%6ujnNEX{HSOY-uD*bh(ai4o30&gCj=Yt+ROuRMTw5yn@NT0}x@ z^pycNELlM+M+q*Hl6p}?=PEd45MoJ~Q8IBw^}EOHdAp(@lRR&T_Uxl zubg(NK41yx!J$YrfF^%jaN$9r>u$jH%h{Rs)ImZ_VFE%TdqXz|pI?`1=n%b;b0S)w z`H$gV%ay*HADf)}xp&dnR`-)o{4m}z|F7u9C)Gv#IYSQ^m z3oa&PnaFBLa1*RGwo$zpYQ7^#h5Bvj$C3#0FUQFI+p|F&^q`qUP&5}*MQ*G)o z;Z#V@Ty{wQWx)4Xp-iEe_%cyeg2-@rJ+5Qxc@;2-Gh~>RhBC{Y+`fo-`?{e|AUMrCE&9u0v1^QX1N6*EX ztz{|8-jqSJ_h!PA>F+7`WX7UpG#JOg{Hkwtx8C}b*%Vi5U0@X{sX7%ZFBsPsYr!U9vJ(?xFIu%;6oe|$nhkt zIX)h}l42Z>8+Yt73irzW7s1{`v_+W(E2D3$>sGj@PJ5)R6CFI$`I=;(462PR%~V!X zhi*}|>c;i!0fQL|UGAg=OXh$h-=yK$egI1KSiQPm9Yf!94$Y{ml1(e3nzcV-Oi@W4(cc@VR$vyIrNW(T zr|}L)uN(=aI)LZX)v29Fspvlp+c;32$;oei`Fa!o`w8ahN{uS{lB_c>}sBJ!14_;9_OWd6Q^u^5)I|r$`mxUBdAwi2$txacxBvi8)X!OBkba_!^%l=VCmJ1Zd<{SmlYI)^y8SYIfnX0D- zc&=lewbT%)mzV)hJ@6KNxv5WLLaexu{c@a$K_QDgTQ|#RRBn8{ARHEA^;m;uUMx&8 zdoZtS@?}7RN6N6y7hekwB?T(bSIK0OgOx~%yuk3#$}oW(-h3uY{#C2+NYmHR0;{lU z;NHDX@jG;u^}<4ex=AAbzGcZ!__39nc@+-%t{r_?|9}H~{C0=cZUyzjN@{tZ(!@!* zsBj1Dt&*@4njMdq$LaS~)hL#yJyf3EbfCl`m{ER z`R4`t!y#00GRArg?$rB~r%rXSqcnk_lu8Wf9ZEc7dFnLJ!qd>z76Kv5&*gUMI969xsq(T2n2Or7Q=m9I${C0lc74wa=cfOO0NlbG}rJ zk!z~@B#~!1coazn1{cpV?@(cJbjFwu3)jCBU)V77y3gDpy!12+ON5*ERGjQSg^SV^0ix5E#$8E`DxyZ;^B{#^n0vf&8t~+k83}&7di?`+~qlZ*E!e9Ff&EW z#y6NcRO{SVa>ufq%^puG@T-;*d)4%9f?6^l3QXJ<8^baG&|Af)v2=|KNdc)w+FqhV`b}He<}Nb0$<<}OLB;4i&<{OY3X97hYe4g ze>z8pS2$<0SIvVFCZUfLw@~MI^2m71sfFG+qM;UQkZrf;gn5!qU(GEfzwh}9oa)Ee zN#MIJ=kbQAj)fmzm{Pr`tp<> zsfhcnQva=Ue7V8uunylsw)%~4Zqc8kYvd$&G-B1zHhT;6No{-2`rHJY6sTQUqga`Q z3QGz&r$l0r1Rm2@Q;13!Y+e~>)fl$%7#j#DnY^mXwhX2aTXm1 z2Oq-ErI*Mw(v7(3AZh*ebcXI^fg%MoUzKgc3f@8vid_UZTWL%MQ_Aj!#7@1xqy8$m zV1#BPGHfu-%C{ngaebgLH=LDAXsBJo1W)Lz>LsqnV;fgO?R8zJ&9f1X5L}C#r3(@+ zfq}sh{kaGjWo4=2;qYX@O_7bFz#(Y{Vz~r8J8eLrtYF&vD>|w6N0aq8 zKd1SoP`nQx!JTV^sE3)4Y84fKh$rZJ?dKX2D@v8S<1gy#)(rz^zMtX3R?t8Wl- z5XaT4U+Ijhc;#7m5XbQ-=H#&QK~u_sFPDp>SF5C? zq>h#_bzIF74Or4F4F>4Xv z>jQR>yQ9wKayXc1CQN%m+u{T4!YQbA6mzY|#t_uWV)>8F;<**@2t& zqmZkNbZA(oN9^MXx3POO&)&whz5C?Xr%F)zL@lAK2>UlPxkpRgj6{ z#k65a&}K%ZH72^6q8g$kl%M8J_vN|PdSE+Mxu4aPy5mNjZ7Eby!}11)2F+lMO%u0q z62?YqqmPP=uk(iBo|Cwd4wnO$%hL4GpoZbghiNL*n9l+QGhUE}`bh*Q>lcD;hSAdq z+OfH2xGEy5qpd6rCng4)W`$2{Cb{IJzoJtKWUz+}KQVyH%T4zgp!Bpbe&UJKNa$pK3nFXuq<|<;BY7$F7i1mC0ec56$P@_4#PAV6$HFAG<$U1j`Z;4IFH0It8IIh0 zpWLENg49Z~h>B^KJ1Q`lo&`}-2=^OB1fr{?;l{W`*7j(()5vj(1&2!Bcm6~qrhey& zlu^=^X7l+_pyh#l^CDQG^lf@TKzs(N?tbC8Q0lWuMv{FzVloWrSKBl@D>u@t9TjEn zn{aQoDi`}Pa8*vKnWc)G=}C+xg&0`sGz9B37?3lvJWnzqV{r#{a>nZH3vIPOw0KGy zu14-qg*7%Z*k6yjeVakuRk=g1u6Arj99pvQjk+Z4^9S3zZwr(24PSExDz$%7{3xW> zsQ1;v#+3qOP`b~e;xXLIUS7B0&?UF3@?;yZub}t?`??pU=T$!r5 zR_fp~5d(eK(L1<&4wzP>0c)h+H z_m{MNW73VagrS1~iMtCXhWr#m$LGhHGnMHbICE@Aq(1C&B6D?JTiB+gone=I5eiVN zl59S_FJlin3F`R|pW`d%+BRPv^LlfLEKxVOEYxt;vSKNWQ+f4Mw+E2cd&~B?ept9r z8W(o_u0;}V*G&qznr!M7^jiCN#YCctBXXP0zSF(oN{9_gVM)Qg&2}Ny%C9E-602I~ z5n)H`7$L*m)x>Sl*I5b@twfGnU*vAhaP||CCgX1sT-xASe+&)vn`ZrhgZl2-D#2h! z`<)yXwP5a-uWqR`wiPP++)bqLe+|e%4Tz^ngG$Njao>vyFLbDP2g$b-8xCD3HFJK< zor`%l1vhU{eN+Yy7M070hq3K4_>k4A-c7vWD=6`V#CLXS23w!O-7EXfx1m?BN6JuD ziPp?RyPi22izgCV+$!HwK?wsn7zwuCE}IT_a-%M7>*NLdo|Y$(HF}JkJ@EQ`7W!tR z8H#5h`2@|UkgC}SFTb^s>Vg1xhDjj%tlX!-FHzeSqw4(!@2K8Oj z`-U@y(=Mw=cXgGv4B9fh_ebkZGtw^x)Sh@Rw%~D4`tGxpb6l+GUb`NzJH^5R zUq8OUw@^9utdyr$So?#PWJIap%%(xQYO;=Iw|0!r!Ij|@yq_HnYsbWyE~%zim{kl zDtZGN>G(bl0S+uf*cdX9m}kw&Rra>l@F@dnvnMnZ73!=VPNru1I76+{=}e-LZZ)1P z1xs_zc@MDYQwPmJBqG*O9Yn84UxT-tx>U9^XT_UTtQ_m+2fbJECPf2E1AOfqEV(ce zmV%hpdVl(??Mv*}sWrMvm2 z-1QnB=&3d%B_EUG!03gUYuP9NPa-HR^f-En(u6m{-PgAjHl|B84mL+sdRS|iiH@o@ zPaK1242$ne=zh)g7VZ*G&Cma$8=p5lfA;a_YEGuOYvs9Hdsep2%JGb!&C6HuO4L<8GiE4^?CZj zz#da}o!Xk&f$MdWvvMXT)1hi>N~}V;TTA-gn0Hj4qP=<@33Kn|+*_&|1J_#~_?41w zv6XU`7CRJDx#K1@-#_1+so&WBh8KN9hC)8(m0q&N8p`q-J4dr3XS`+M`g5%PYA0Y6 z1k9)b#oD1r=i%J&tkfMjsDEv!Y!~l2e6*S)tV~LK3tsfNoqKU@f%p^p-UsKFA?D9R z_}{&&s}2dq#R7zgm||Iiv@F(xZ71?n27=Yuq0_E2>2H*#-kIY~rCvsi(h}RG8IF)q z*sd~_S6W}+FS`M$s;;OIMWHxcX1tA$)2{VJ_>DagMp?n(!hk5~OD#wiaQ0e^5c^yy zsJC;-$<4tuwSZbb4X&u_cD3PyWntt!tuq<@^nK36lzaS^3n>2n$vMba`OdrW#Ah_A zFhau&COO@7lzz@WVcP4p?~<4F7hR&66$G(Cmus!d_eWflQ&#dX=doc@`Td0*>IXqP ztIi{rH;gW%s+Y2YrjGHJ&<>R1237{+g>ihcW65`_LK{XcM+@1;1u#O!oO60g{G6Mx zbr?u;ydFSpvuUvCTw|f5IwqrsrBbtcY|KeY>H`VAIqYu=N`2|v+-s-{QUbM6}gTsPVpmu zvD&>Uu%W^16NL;lgNkjgh?dBBXv5;VSW9iQ&*5n4G0yjl2$oF#4Iq}sq1KtDT+pV& z;hSK3Vk}_+J2j!)$8@1Dq#kJ0;aeDyJ`4|_cuI_>i6X4@&MweppY9tuW!=sI1eKdZS?-shqyv#l$EQ=fqcV+>%?!pT;_dFr8wO{xZcDNUWaYr&cd zB}49VjZ3p168h~!^mE@Y5}XA)bqh!G>j`pXK+>586$X>$nVZEq@va;BnB+_$7+;J$Dz_rb6VQ&N^H*jiIgBO^5JH#2%%kr%xuYr(F9{saZ7X0h zNGXk#fNuxriR5AjF2cWT=q!dR&L@b)*E>BHziR65>ypWJ8zu=!2!B<@vy}I8mzhMP z#5yr>(2nI(z_3*a8`#J-sI_BA7Ez@m+o2eg7o>2EOmmgSCxzUo6>jeSO8=jm|XeI)!nWrdK{RQR)fG zGlkQ(=xe80mzR^LTtU7j9#JG-=Zn!5d!ya;-tb1t9;2E$rhKiX1kMZPp8GZ>y<*d* z`CH3a#OLVwN--C0=~Oi7IAMP0-7$u`nyl(>%i&=oGxO%X_5Sx!HryUy&A);-<#E-c zCNp;gSsM%TQ+mlI68nq*?}t)ig>NQtc*ASrfJnh zGNFtt0P!dfL*Iw;lkRb(t3eokRh@%(FmaxlJhYP{4NUUo#=z_+NvuL+VPY$H^oYbKs1uqY zRO%aX)LE_oP+~WAa%-?OFWv6C#HK#SXCX0N#v@KtZUr_Q#DT9=u z#B@(wB ztS8*N>4uMWwfxlJ-Z?GC)rNQur=k@GOn+!YT=?aeT~5gx>ZG0qg-3$DjutuF4R75HdP#!6{M^Q$WF?{c)!^e%JGkdx`b6X|`#q=HVfd@Dgl&T|t(w-r;6ZkGo z`3mK6EakSpE4q;w+P-_M9^UG8?DC@}tVk*o?eBZG`pS0~cj?jhu?&<~4!7seuVZUO zwXQS|P;9?{au?%J4Wz=3rt#lPgEjU*_#5#CcIgbCYNlS_ZBZ;*+RGZ#wV$_o96Ez{ zTnc|eMDF*IU2psPWRtP9w)a4i!FX7TmR{bK5H+Y9a3`8h%Boa4KGNJwK=DSO8)<~|c_P$TBxSWhY5UEAp>IZv1n zlgAd#SQ#ar5~1zx7^-I}pI<4+LtH`leR%a7GJ`vY?RUn<&9wQ98Ze$0Pgj{95D@B6 z)NUf6{k&j0!BwiFsz$M^H-li2cO!QL+1hIfZ5%C!Z;ki9T`32)oS`z8@@CL$9)2H^ zp>hsd0WDhZ(CUA@lzd^tV=0F@%#%YV%YhNwff9p0<8k8@JEUq#_@n7I(+7#*>Rppt z=R8t-wyPwYX}3M#&t`5m>E0Q<6!q-ILk-)uaCJd^lm?gda_G8~RCi&07mYQQUMBhD zN(=A2ee=kL`QgWv!4xt3tDLj1D}5giKPIi5raT$B*n!S_wyd9uIp4;iLBgm06>;?k zJ}MGi6$Iqt0U_W2af3u6M>dcE?Eej3l^^^UEY;jF8O&C0>|5xnsT)Ug*VjE5F@eBu zL;;#r!%x3h-kwvCrN5soE8Ek}(vbj@7izAoyRx90^}2RF%0i{PTQy`-!$vOjVj-rl zcV3vwlk_InzbK3Za{s?k4b02?7v1z#8yy0toWxg0zNZs$4rwo=V|!qV z0cW*qO|~Ul`H&7Ng_?%3-W^mrtDYT(Te#FTkEUp5=Fa@u2q$6 zOr{IilXAu!PRG_x^O#7!BR;-VIuAYk5HCVFU&R>jmHwImD0|0-yegRlP?E8wrN{Nlzg=0A}5xg=l@<+UHNsaxBdyTHsnRN?l<4$5j! zbxbY^IZ3t0F0SVGat;=705bWMoXyOgk%;@ua%Kqldu!LnEC4Q5b4zPP8b1awKSZRf zIYOf5=ICf=ZjS`W=SD^yIg>hEMotoH><9q;j=YatUDDje)Y;n6748fG{)XC*7;EDC z8?AtxB(e*VV*rHtJJrB%8fyVSe<0-lgev|OTmOHy4Zww588Oz_!4)~D3j*%{hkI0! z*#1%u2snPEm0UlqmRuNpx@aHexpP9q2n*)14z4m?vRJb zMO^{>0w5%9f}^atwWXCS00II3RLB*;f!O*#XaOu;5b(&DT*&FfeT$To6PWSbB1`%^iD9EhA2 zIR-%mf$RqY{xX;TpBUJGDTF`(@Hak*KfQxQzWyT(fIoJV#6zU#f6JObX)?G>jRC+P zyXGP2XE#Jh{!2IT&u-kD{Quq&`G%&x41g_lDk8`I$bFxcrv_{A&*U|9Aqz^Z(!o zuyXezNO^B@ahP`$tTDu^*?rr&ZHt1lKo&OvSQ&gn8>J1Mh1>tY&G=@+hi+Ole1ZZR zGCe(*9twQ`@xZ@W@k{P-{j`FYlb;U=262OVzz_h(U+(_%4WiWkw|Mg)s^T9-_kRla zFIEnNc|bf6RsT=*7&QEJFR&u;#pwfI^@U)e>{R1N{)<7sx@Lmmx?65Ow}>-SOWwB z^87<*2$=UDw($sX|HB77NDu#Q4IVzk6a3#~U;zYs!(U|xGycWTK*V1DYd;VV5b^Z# zHyKy}jCeRiUVp{|0^$B8mafjmh$kKAAO8vccq;vozX-x5RX7}(_CLs%)U3VC5lMqg ZBo|j>XV)KT1cLBEfS3#nQcBX8{|6;C)Mfwx literal 0 HcmV?d00001 diff --git a/src/models/niser_ode.py b/src/models/niser_ode.py index e097c7b..cbf2954 100644 --- a/src/models/niser_ode.py +++ b/src/models/niser_ode.py @@ -77,10 +77,10 @@ def forward(self, t, h): # edge_index = self.edge_index_batchs[0] - edge_idx = self.graph.filter_edges(lambda edges: edges.data['t'] <= t) + # edge_idx = self.graph.filter_edges(lambda edges: edges.data['t'] <= t) # print(sum(edge_idx.long())) edge_index = self.graph.edges() - graph = dgl.graph((edge_index[0][edge_idx], edge_index[1][edge_idx]), num_nodes=self.graph.number_of_nodes(), device=self.device) + graph = dgl.graph((edge_index[0], edge_index[1]), num_nodes=self.graph.number_of_nodes(), device=self.device) graph = dgl.remove_self_loop(graph) graph = dgl.add_reverse_edges(graph) # graph = dgl.add_self_loop(graph) @@ -266,7 +266,7 @@ def forward(self, g, feat, last_nodes): class NISER_ODE(nn.Module): - def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=True, scale=12): + def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=True, scale=12, solver=None): super().__init__() self.num_items = num_items self.embedding = nn.Embedding(num_items, embedding_dim) @@ -276,6 +276,7 @@ def __init__(self, num_items, embedding_dim, num_layers, feat_drop=0.0, norm=Tru self.layers = nn.ModuleList() self.norm = norm self.scale = scale + self.solver = solver input_dim = embedding_dim for i in range(num_layers): layer = GGNNLayer( @@ -359,7 +360,7 @@ def forward(self, mg, embeds_ids, times, num_nodes): t = th.tensor([0., t_end], device=mg.device) # print(t) - feat = odeint(self.ODEFunc, feat, t=t, method='rk4')[-1] # + feat + feat = odeint(self.ODEFunc, feat, t=t, method=self.solver)[-1] # + feat last_nodes = mg.filter_nodes(lambda nodes: nodes.data['last'] == 1) if self.norm: diff --git a/src/scripts/main_niser_ode.py b/src/scripts/main_niser_ode.py index cce07f6..91d9230 100644 --- a/src/scripts/main_niser_ode.py +++ b/src/scripts/main_niser_ode.py @@ -38,6 +38,9 @@ def get_freer_gpu(): 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( @@ -135,7 +138,7 @@ def get_freer_gpu(): collate_fn=collate_fn, ) -model = NISER_ODE(num_items, args.embedding_dim, args.num_layers, feat_drop=args.feat_drop) +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) diff --git a/src/sweep.yaml b/src/sweep.yaml index c80e4b9..54f15fc 100644 --- a/src/sweep.yaml +++ b/src/sweep.yaml @@ -1,11 +1,9 @@ -project: gng-ode-hyperparameters +project: gng-ode-solvers program: scripts/main_niser_ode.py method: grid parameters: - embedding-dim: - values: [64, 128, 256, 512] - num-layers: - values: [1, 2, 3, 4, 5] + solver: + values: ["euler", "rk4", "implicit_adams"] dataset-dir: values: ["../datasets/gowalla", "../datasets/tmall", "../datasets/nowplaying"] \ No newline at end of file diff --git a/src/tmall_hit_solver.pdf b/src/tmall_hit_solver.pdf new file mode 100644 index 0000000000000000000000000000000000000000..74de46a1e0a3e5583159255ad3a2dfa65d6d6579 GIT binary patch literal 20382 zcma&Mb9861_byyaZQHhO8&kXe)V6Kgwr$&;+IBm&?RUOE+`e$6L;%hJ zYdbSSetu{ZTjT#UrBY<-+4uVv$Kh#Eg{=~iqH&- zj&?@MCeDP~|Hcs&|JT2XyE7q!l+C{}h5tv2{Ew6()Fot4HU}76*qRZt|F^k@f}y2} zkuxF3e|7$U-C_EFCm>{NYv=4l$i(>H9*F$&#>DoYg6qG_awf(W0AV}#e?I)D1uG*v zJtGqb8zBn|Gd(8E7a%?X+Q8{B{HJ2^X=0BoS$vwJjU>`=K7ubccr z;_So6RDB5m01~1r;PB?KaL~-)Rq}dLt=Er_4t-YzZoV2RW+BkUlygU)C;BE2m$!?H z=lkt5nYS^o1}cS3ZK3;joK4<#Kw zN;C+A`8;H0MwKvLMEI8abUD8_hRqGe7Hd`Mj2*z>-A=2)Y9h0vK!N>hI`IP=ff5S^o7(N0)ZQkY{1ouKL ztkajQ_I9H#E3F{+wnNoyl4s>~EO6XCZF-#m6Qj$}SPCW+Y_8?EQ2;!>r zSCsr}E+}}D_IcXK{6E;99q#aoB)VH3X91IMIxngPm0bP zO_+bREQ0qbtX+C8?V~d!LUiupxfJV$DTKQj_?tj7Zx;WBTVPjGw4g(HfZX% zhQ$kH^XgDY925-@H^Q6uz1~X2O-(CaUmYaCReIa~#1j4!Ex;w$6?|;$v~uFiR$s7= zqNKrQE$p%mM>QNF;;&LN5icbm?n_>y(@C5nI23GI1G8P<#wosz^$ zVd&M8W;(8GlTAcm#+yYs(o#!#u-2wid6o@3DwoB~G;4=G0wX$wf}PVq27Xm!TSPKu zJa<%HQ&fK6!C+%)3^N zu05rF#&{<--*&GJ7ISjhFBZm5H|DL=8FSPyZr^bDC9`JoK(K%)rBSvwB}AgkO_MlN0lCz3sW1vVs_J(&|#T{^c2>-BUk%|*!IU@J4<(b+Q ze9O0mB1e5KePaQ8-G#>BP%a^cesEj*X7-*3%p%gvYaO0DY9J}O;#yv@ty%8!Yh~#$ zcDPU_QL@9FY?GQ-No3sRm$nQ6m+n(e89EdM4^oqqT?80`uc+0}n$FFqThJ>(+m0Om zf7tqeDCU27`@g8?KODux&dmD1Wd2`F`5&eBKj!;hF{g? zwVDY(&2MM*L7i)X=3+Wsub0;Sx-Y-Ifo8XX7TT0A&QAg{jH~_N?yHo|=sMVqPUo(G z-flphx!7#(FGu98a^LX#Gav|Mv)D^)@AH1#JrVf(VBfszVZ93WWn>xpomZOIVbB@! zY|lfl?hHEtE7?*maj-p3WGiFut9Gv22OUZg&d6_?+}^s|msH;`k#3+%-Toe{GdaqA z{wdA~ZW!xG zfO%$Ui$qD!XGbBu3kqG2q}Q+~W{fj(1*`>ZqCg4WOMlr2oPsfQ#QZLziiT55{XCpzrct2ywmL3$}fNU~qu3ml_gw>?6j zpqS0Jc?0n5q7Z9HIeUW6MBl$`?TJ16 zLr%p$k-YZ?+6lZd5&Q|5#+l$i!p{(weEku}7>kKa1tFPG2E&ntZ-`)+5u1omm?B^C zjs$--pcoUmhEf{Rk6|8}V-O;66S?lMp*bQw5QVP6brQNtd_tfKst*f~OV9?|Zg^1s zT?6eTNfjhARG?h=hA+YO5@Is6i#3@8yADFchE^xr}>N zFZZ8Bq$xl+0HR!-84wI7fe-8rPZJRsrr}f*@*Acp0C8Z1otqg%j1>V89x!4CbK=sk z_J$zgN(`cf7v)RyfEEE54=`iZ2M~uF7*#&ND*>1XoWm^kY_Mam?wAo8phbvv`<%ey zR{A1z(L&WhRfrVCzF-6xq)dA9BF5Htr<143h2xY{RA(dgzSQvxMXDr&_v+yVb@PnuDyg|eCi1uHmBs)N> zL_6aSR8vAa;7&+(fgT$c8Z!;R=Sb@PAj4UJf@i2LkSIiTM;)+LD}te#dQ1m4S|J@! z4#L_HoR+wJq*US^@JU89yaNe)o~>Ya2wDjp1R6vR!rVZfebeEtJ@nz-e(5XHt?4^_ zjSzQ+YVmG_6@#;{aEg6kji^^bR!f3GyM6ZIFax}UAzXq09z>2o&tbWG;Je?$2zyNJ zD0dF+Fb*uA;CF8JZO{C?z%S-E14;G_2g>#Y2aC9E{a@j;{h;BqgLX)5j^6P2!du`7 zB;|d+@Dj}e0~ok{!EK0rfo&}!w`{m#pP`3XoT0CDxT`zCO}A#lPPjkNenOu( z-BGV@oB{5+hXQ*(S8(k>1cbg21jN2^?ufqrN#O#6RKt0DB*S@Jt%DRB7jxicNNlgD zcYf?&KYm%?t3Qre-%ozli+!IQ>|c}m>yHfIF>*hj>g-=lKmDb7z2=K2`i6@FZ~tmh zZ2!p=xu5jS$2gv~$GDFFoAotWRB*`hjof@x{t03K+WwJY|62R8e*c&Ja^!yOUKW46 zY~K4G7JtnCE49e|9R7f^NiftisN^|di2DD0n0HM z7F@a@18VvW{2O)C<4A?_l;>7sT71+0AcrV3M`GAdSoNu@s?YYBbSgfZ(u zUU4{esOQm`kmCk0kAb{y<%o4-Ivx5%v6G=Wqh>tCo-2rJe`e=paJ8_N?hz^(Wx2X;YOeHUDV;I+D1} z5LkAv$T8byrrb+KL{u?{68gNcGkY#w`W_L2!QzSQM&KjU5a5qXDP5tyu@T{s5@UdY zgC{SLZV^&H|A(?VWQIn$dom@sS{#7~G`HA*VZ$$Fb#cm5qk;*u^WcDKj7Q+d(*-Ph z!Ubh@HE~a%AXcaTEiyX3ZG1@>pvUuE2u{TPrIaKzd7}cZ?JYb|2`x^zd_N0dPTbhb zrR5SQS8|w|j1?u=t&DXnAk||&V4jSuPUy}faZU56jXzAcVopuK;_#QCYZJ8Y{HiPKg0fFi`X<)H1(ZV;w65?Kd2=#_}1#tm~4i$ z`w>&y^1MEfaf%8ibyujcrq(Kh1Qq47uf%B8U{FL|jkb=4Yqt$4XaUjD=ky*>3w}z* zh9JaW$Sas;0pFiYP`W?CxdP<8IZxoWFm|(>yY01NArUQ1i$~)}op(kk!n$l;p>(0M zHInYOoL66D5hwQ!S(y7WkB-!yKU&M>pW=n|wzzEUi3yYlqLmlaftxvVoQ4zd`n;E2&BZc1uFI*oYgg!2)P~G%(Xt&Hqej^| zBbDMhxSE4ns2#WzqQtyZ(R5fi1N=|Mj3pGT^5+rO?y4x=G)!1z%zQ*Ibv7zq4c@MP z!hh6A_DB7IX8gx9?C7T}WR6c8m;HdE8a*5xDgho|yhu@Gj!x^S^Dhz^hBzO$&JJna zVsxy?8C9W7{rf=x&Q_P@o_mK^7Upz+JFAExQCU*#&^Wo$Ax5UTUw5AxE`ckJiC-NS5cD7d4 zB6=yZNiRKGnGxjJZ=(<(xe;^0y%J8UF@0mxh=$VO^cX*G{tAk=Q#ei<4|pg9$?}$< zM2#9_cU!Fa5;cVvkaBU>sU;i{)~Mn3z9gCY;NO2;IDB zd(mOPjaUZ1P#Ceai`h~ZSpwjHgY#XB7iR{9jZ$c^vY3D-)-XJ*@oK2 zCNdKaN3Yc!Ec93`tx?FD{cQRy+2fmk${X-X##9TjSA(}1rB9nl&f)Mpgic}0;>+hc zs8Z24=#C&QV z+fj+@x*udi+1vOEThDgLI`M*qf%GXT7?(Va zgP%}5pm3xMJ#Np4SQKZ?5-qA$eiN6>KzYKct$|Ekj@<@*BqU2?)zxk%gZ8)k4qDA6+U zwAcdRl5;dx1atwsSWXg(G}|PcS}7f0)Dz{5(^({v1Wtr`jzsP?9o@5?YJ*{ERqjVr z9~2?>mFQgpzgJ5Ccpv8x4Zk-Qz#R&4RwZ{vbZRy1*zctRs|x$jpi}6#mSBQl;bY@W z5UUHxL{NciLAmg&_(_txGtQJ<4t~8MvB!Q6Xng&b{{1bK`xMX}R)Kr7fP6n=ka^uC zV`x27HDsGPtK= zB5CYKDG_uF$fQJfm}&FQ?s$k4-nIEOD!kCNC!q_Y5-iR_EY1^_Ir2*p-ss{x72J%$ z3@KB6ZWtD~qg0ATT`#n$GoJ9?XpY`!Lw78B;t0z8GzaXS=#=hgG!3N51aiVF3GS$3 zqO*7&6rA4NBDDVk5M;seKRakAyGauDZ*gXO33cxyHTzxhX=1^hDbBQha3xD|1n29dpDLb);;}$W^19F5%+=D`69IOe#6**W6 z`snCpM6ez(s)abxOI1a-O8r8!>PoS?7|G<}O5}Hddfc}BINox4R=$-MB>fYesF7wB z$P~jaT=1IOHogI%Otnp2J!rM3d`sK*#)t+_@s1&%J)}XO{)X(XdALWFWQ@GESd{`= zLsX`!R{ArP>$+=Vcf8|Ip^|S ztUn~%nA_`2+>(a8oZ&RkMMQChEa`P;u7doNqst z7osVFf+y>t9cPd5gzMDSfH#6an+#Jl^^wqw$QRNvmGI6fm8q}(s7zx;&scLZk2GH2 zJx~h&p~4maf$AvYt86vdw*hl~zRPpLz4fdjMnztzQf{l&Jd8lHMHzWHfeI}`Dk1XrdFSv)?9N?ni$|#eZ&W=NyUWMhTUt!{#{K4xq^{UgY`{>Y zG`+{lM${}d_A>)D&L}mF;VibsJwv37#+~4EHS&L@fV+TC&Hj_u`i4!oxjJkgYy9hc zaj+9B8^WT_Jnfv#Fom5Rteq?nuQ_cG=;-!i#{K8z=;6rD$E$@jv;z4CVwlbK(#UR( zry0mgB88Gjm8n*-cOmP6!S|58nryv<&#Ez)CUjb3(A;5-ufDGOYGjTIv?my#O5_fb zD&*tSY1`EcrnXcUs`)CuS@UYqUG9u^G#0N0S)umv@%uj-xwK`Nesn(}74-s(f#98P z5D{fbV2b*5qzhHzzgfYGl(fV7PGuU73wp76{#|?(_x8l9G#N2TmIgd^*PbbL#7UH*JG!(IHWhBi|;86es)DDScQ7Kqhhp5 z>?7XmfpHAr+oRMylhKB$4*u@8LQ5w?c^_m<2W?BDX$fJ(myc`{gwlap9ne|Gk6FH+ zyIx=)x-n?GmH5`1K}foFll)M7k9y{J9^knSfY>ALklBH9@RkO3iL*fjFyaz_SioB4f#jhuXi`6pY2mLTZvY=b3&Y6!e_$U!YaS@m$L@jy8E zjhI&6fmjFVKU?~v>;uz_oiWuk_{(tU-`Ph zB_xAC17;Frl>$)HZ0BMFP&$Xqt*?^YPt1;ZJW;+=Vd^45n|%!i96FL>*t>b%M#8VTOkL2%x zxI|;)8O-&$BWH@uLM#S(_r)3rzQrS-4(JwY7Hm2hHbz_V%)P3vcxnW@&_WNrH3Dk- zAn*rruR@1z>A&R%$3%(64F~kojpOH!#6P!1Z7Jds%ElLqvGN>)y=8lc$<7FlC?aX- zmgj-E1*X6v=iW{MJ8}etVRvr9z9N2zyuzi3D3Kw|1zPDrHClAm{k!t}B=ZdT=xn$r zAZL0OMUTmk(T|DF@@5OW(N>>X0s=dko9_Ue<{FExwwj6>%@0T&N)a7C`ytiPn!2G~ z*yNWU&coJ$--d$?4@o);VLA~cECEx5595Zo6DdbR_dGk#C*v@j(thM-4b!IBM>f~e zQ5FK+sD7(+RL=O^MsD+_p4s2ydRUNZ17t_gpdr|R$a+By;w>7|V2%)k;J6T&!P(bE zQ8Czv3`3RZ2+$Fz1K~TQJEuD*i|)tJ#|ruR>6PtDtNAaqo=pCC!2A9?_Lt`u{TGu6 z-ONlHfP#1Tjz-H-T7yUhnr9zgi}Kwb=7A!67R8{EG5tWcL?MCA8Ox$nbHPy8ZX*je zz{wDG>|M1pK`*~2WS}@QIvA_jtC2lNc`HyTW< zt(({x#b^-6UYJwr@u287w;U0HOC%39bf4D!_a#Y82<$$QyL`Gt>VUi!RQ{*2M%aTp zKDi#I@PUb9m^t;#83 z^~|{;fm@PA@y_Czj&XuG@F+N?SlnIFEwvX0X1~6?KWwmR?S1Hx3sCPak1bSM08SgE z4ajW;NtHAjX`4_%9LF3-U+0-F5e`TX*mrEEKAbPT6296Aa?P(}#B}QBbgCWpV8W&w z4tnh(xC#zuG3P#iRW%>^{x#kCyBVxh8a+@ALeRU&=2!G$u0KR`g?)}IhO>}Yy_R+) z8$Vdypr?CewPhF5-vhOK-K8s<65z*rR>2EloU<-x(w-Nkw*zxC6bA2*z7xiIi4B;T zgFFCfveiYio4u=>Fw6&#tZ3L1>+`*~_aGra0JRmaV6z2yQ0MYS3bU54A|xDSv00N5 zRnLF-jI4ZdOmP(Ip8fenJY87A9445`X3+t6<&bjNp&7mKOLMb;=K;YsQCU*SndnW}T8zP28#-HX%1#x8!xyhC}m4|ZR+4Rhee zLp?2UXu2%_Q710yf_A7$j)=IOo0tj=OUoI-w~Xs;=qMOnFL*reVBygDjgvyT4Fu}q zg!xS>p592BZ|C2RrCHPX9bzAa>y{-&nxxk(Sxly&D3ic#6T#+&9`SKiM?}FQ8UkLz z23^;yY&m0@&ef%}nbW5DZ|v7B*e=!5DaZbn%uL1}7*(GoGp4$0W#HP)PI9c)pFOV7 zTavD%G`9TYk)6xVF%B_LE>WZxF?w2O?F%j;b$$lPM%`B3cA^HQzNHR{MyXb%b?o)Q zdUe{y-22mqVV= zI1Bn+F$k8Iu!G%e87eGN)4V#+Z(am}D!NySSA5VgTee*?C7qI{@=FDsFfg8LOXj4? zbvS$vr=?$*k6DG6rYxAW!=2li-9v5G;NnxAU zJYMjtX;{1ReJK-Pq|mn%$vGgK2LK}5R6_jzyJ z14cN&;9{HkHUE@?=IQY~a3HX8!}onG zuQX`ut#y~y^2{VNZI~6}m`O*gl>IV$^f>%?RD=kvD@SG133Esy*DV`MX{626Jr!kd zaFe;SNKBJfCHnEz){hV_iew|B`7&p#XA8WaLA@+IIy}Ysg6?@9D&H%6RQM2S42F+Z zrI@>>Q~qjC8}OxNFK?ad;JoH&CJch)afm0a6Q4uzr!h7R zH-LIUIk`{&+lR)E5j%X&EU%}{?(nTjgQl5RPbh;apb#acHsAPz$MX1`| z-yUkn+jAqE!k+AJv59=e6t*(YV=;XiRo&>0tIw`MBPHY3D{mTd3$`t}rb!Rce~vRy zDJrR<5nEEGs$7%XGo}v$7t?g2TPNe8e$OHjXUsEA7kHV{?>?_J4$|bf8u?7tu$0z0 zv7326Z5;%$)a7_Tgtjuo{s@t@xf5vQw0oQt%}AUcm?YWibuf27v{c9?QSdIYeU0!u zw|yP5ST_^c_(y(FgF#SW-v9`#GEF__9-GfxlAKHIOWtgM*B|0SB+1XoME;%&f=ciQ+Rc5T5?GqMAYp;l;ecCcBlZdPUgGP0gswhP-%X?3!UyzKh zgv@bqn2|v2BO>158^u^%lk=_+y+ronA;#JZwXeAD3@L{Su zxgxJ(3%>GXl0JD|dYv#Wwe}ht<+n+-8Vx+hE^t;0Le^!8xnt}|Oh}2aX_}IdtPG8j zs!9=KJjFHZNK8p-D4?X&y?H4aYq(y9s;zYyt2NqC$(+Vn%sMb;&P!W;AcjOZYE3>CQ>CW)@{YbEKi? zoHbGY%Yo+r9|-=w)~Gf1A~JiNJ#q`a;Iu*H{Q8VdU5ZhPpOlu-Mo%mS9z#}a)(#A& zaoB3&1;&xB62~Eavk{18*tw#PYuI7_$aynoHhLA7D(STWT!c2eX8DL=OqXQk?^$gz z$K(nGtrqumwpn(eoFaPl)vCfs>4RhG1F}8Uo~6xN(j7R!2TU<$^$~uy3SSm2&reEY zFuiK3{!$4z*tj+Rjd=_l)b!xla%e3RR0-tPtwWZ?-{pK(?n=1xNRpuxqWoU}$qESt zLhg65RnR6|+nvtqN|&76-k{#i_s_W$w`+<}u(rG(BN_`u96O%dwIDf>$?eCw*;YF2 zUDy4Z)oW|+k?z(3N-m6zZq{)!5oI+7V=|&GV<^%&Dn46sQx)#o___qFfv_qUob!dmh;wkA{q!~|RJqf}k8K-6ND;30bQG289J3Rm1&loiR zG9H5cIJyNH8r-8Esu?&8HL}g(p9xUwNP|8wBFmZ4kcC{!8j118#lagu9>AJIvq07#UfC~Fbgz`HS3Ug zY}d>fo5nw}DA43Fl%jxRyPFatDy7A5b6y)u9!`H}v{ILQoX5S+JbRpKUtb47oA{Nb zH7JYc;EGsry9jT$bIdaKh{u=k>;tfLx_FXM^IdX%-*3^k{>pDLKE)IHYv+48iqjfS zLdVhxcfE45`9iVW>ht1yn2}{~)Sk@y6;`JA{uY!;w@eqe)n|hPFNn9n)AmT34zS&F z?3We&ECT;6u07gZk3&{y&AbW&Vj{wv%rFn~qz(6UzsiVeA`k$5j5$R#Iq&hPsnp~i z?%8Xqx4!pa7Fv#x&QW>*RKMJ{q?V~u>Cng>$)dxJ==?5;(A85czt&NEC98^XVY~mg zuMW*`tMfXS-D^{U?#6~pKVcQ$@O#~8P0jVMNwC)eeNM71c*6sfG@DszzctzXt*J)Xu1D~y zNAT3Z;gDg~MWwt|s5(N|>LZ~=ZU~7#aRCTFfDs#bQ%Jqd=nuzqLxie4{H*Mp$8CPx zIe9&7w$lSu&DUQ5--yk&gZEGtuPn0VRhU?6Dv@q?F@er#$>G^p7#e8{GTUvxAMR$% z=ok0|cn$HlQ_vOY9(Xir&OBv&b>vKUbd>zXuT2!J+nRK;nquoSD{ZyZX%k zX!|aku9cTsO|JTjF+sPA(o2`)`u6O~y~ne4X0r;B1B+&xI71s8!?tN_Mwx+$&SHW$ zZAhdrjx#w)M~x=)K#u`&9K)nCR?8+EN5@<8qAdMmThqGv^_1 zw!xKdN7Yu{X5DO#Zc${CbzCi0e~VWfQbTBsNq0Fw|94FLRSPiRkP|k7jH(#cc6q-O zgkOV?nUsfMX%ef1h#H-9vBW)hcQ-RLE?k$TK1w`wZXjzcJn(37oo3cx1A4(DzGHSj zgz2ZK*45mob-L^zBcJ`Ix^Lcg+FGs>+0?#nDCAc!#?u=m(b?49&V?7FboQ;(S&^LV zsj>Oqh|RaYGxb^R(>cVAy^wNIJ~QZ9dC0DIVsbT~p<#r1E`C8<-NoHIvfy&F4p@ToZ)A18Z}^&R%^IqIa1Wio6{StGt9FJ zmnBbM89y3c^POtj=PX#$Jmw&I3{y!MIj3p3zZ9}O$GU&ylO}9Bmmi_fXfsf%ovJ5m z0@3iE=5Y7)3YH)+DO7;gq}s0;4JGob{BA?JXSW-lw0SMSDFKDLF7h> z`jFN!`KkKdQ2GUeSR$#HLcpbz}xTPaZ7Cy?NRxrynS*_FyRp;trIv>u0`Z; zWCJ=9a(+W6S~qGst+q_UUbu14s-+zy2RpK!4SDyKH$4Lz7?Z_T8p@8BCB)9aAW{@l z&3V8MgX4hX^uC^NboAjaK#sQw;wzkYwsq2F2nU{QuFZVwRL*OH-5zW^3xSf8Q8}wOF3^qS2S_pOESlP zI=9=-&2QZ`m2=lGOk;djpHi-hyvB-8C@jk}Rv(c8jfw48-O$?|f)dyuoA zOI%=ubNaRD9wDv7ij<*Dp;%!ig$w7N{rJqw`0-RHm!qGPu&D8&JoAe0Rh3g|gP2`qzqv@BMYKq_GrITM`q3 zErs#U4;`&b?tnOAIr0yagTK#ncxHz2JPyvxNaIs-jd!3v z*W?qvqhTEwiNpyxvZTx`q24eQG!;W1FkNeFe?+!q-+0CNB9X71Y`Ir^CUWepZ0F^2 z)`JPytJXrDE|c;Y>-|v49P{{(Uw$xO>cxTlSTWa%gpS6(tEZ+Ba;aXPC|`2J+QH#J zMqK1`soPSSAkMsLSAM5oV&<-^;B0O&(p~kD+Rke}uA}Er2~L4q`3pu58`1!!XLW>M4Iy$2cTdhmbd#p==fs$AoXDegi_lKiBqxzg6KdiV;Z%%zoO_=f zgj$V8X%QXr!$Bhm#T9kq~-d2@^Rdr>57Z{UL#l;Zp^Q zweDelnB68kf|@7@=ZPnI4si~7^y$u$>+(WXVbhWDB78Rh?}khns7?ly{xw! zRzEoQ3XP}0Gv&s?36B|3%^)L*c5-GNQf?GqHDT-{0&-@ZuMdFTFb2$+{gRn(H{eqv zu}TLnr?AqBU>dMlwY@EyYKS9irv6gt(vi)=GdokTl`&B)NCz*$(BP}U=4El4wQuAu%Mc#dAOAoe6VO%goc#La2#qbfoT zY44ku4n)XWQqpDKqe4Mw#Yy5AtNcB^7US>%LCnOt{mNR_6b-uLQZO-T`)%q)I3aIY zndF5YJ$>jY1q~H}$TuU}cs$vPWxN|^6Gxx+{P_tY)l6rZBP+tutULYpm=2Om2SZ?; z%t6ffrpz_@@#v<)6%~zQ$aXskVCs7RrOKYET7{a(E=@t^Ksm@I1u5lJ0xp##+N=eQ zmKYknJPemvRo`w-690Ajz8I@^^`T#_r3F!yn#{H;&&arZX^)w2aWw`?g%?6}P{REu zIoaklJ@bCq+?P2kW{PflA(sH6cm-M))h7tjBPIto?_?QAO9^Qyzi$`h{xQ~>NR}_uUi2A!u#r%7A&478oBK%=$Tr3dagBcs7)DL8}uUY?$VAm^% zr)A-DD;s(8tv!i>OX{~qB6UW}wlMn5IOf{CpCp#TBl0plml-#$Fk6iXb(xAcJA#|B ziHV81VF&W&D^26VGuXtDvuF;vw*T66?K2qFv3z5b*SGx0kbdl)NyN<>H@#vUIxmya zyiY+EJb$B!gt2#ChZ{#+j`wk-p5bFI5ej+ic9V^|6}*r8O0eFGwQWxvJDncyO2!1;5LA;W*-Mo~cUgC)IxJ2$df8@d zFW7v?bmY2Wqp?Sa!2WS>NP-ud__Zn~gD0AJ?XOYUnSp<2=06&P_poBYn5`;C=DMIP z5fV3@=wCn4BHglQpHoA(xep5+J90b0Zhh*H(9Ltle>*)2)=g#b(7$EIgx@H}23H0e z5S26A1#ExJBuv7?`sVxG3(#{U^#DAzI?#l~ew1Wdxn)p%%T^RAuwD8)#bb^*nX1``mao zJ%P36RrP=Az)R1=@-sW7MfG7g*R$|ne=w(1n=uhE)L)?GS>)Ts1EOm;7*#p84Jeho z*1bBSdQvP?dj9C1;W5OdNtClaW4}P&SX;F6YUbGFUFBZ(OEKC!c6%~S)u+<|&ypECGL=AYnPF73l=FTb@$i7D9u&)gvhxd$`k%R>nTg*nKer_)~6FV8Gz6hBA&yaJ zaOKZ!&$fD8+RTVe?33616_%5O&}{YV8+hKcUQd*eCOM8+5Y;rcLNRiX^Ju3mn)q!6QD z)>0v!1;%kCWO}?U6oo|FTTKG}JEdKzij~qZJ&F-mq38RP*48*srMu!W?h1apO(#|Q zqm36j%L<8o?GMNSWpT)}wg1;CSb`j>&nGWs@`ZO=0U56Jr8)IXi}oI$RAN8BUU3|h4TMi%t6+! zYsM^);D$#*2$5SBK>?V|L+}Ur8UL14z4KJno1o2w(Sn=*m_f5Eag$6J!p$+kC zCQZJI?c^kH^bSy@yzvc!y|k*Xor-5*#Fn?zT*r^kCqO8k)O_B+zxeNKt5l)Sx2Mjb zG-_Vr1n&IFc&7h;_FjWXFvA9B+HwG%4$HO8=%r76ujL??lb+)4Vc(i^;o$qGuW z*}Wd}RO8sqlXKM{p6dh5>K`TPCrlIoZbz+tf5%8J;iJ7r z`pm|k9wbbTj?kCUFB$TW!QjA}{c|=<0yxez@2D{3W0|p$&;f9e4;VAmNc*%j>Mk^P zxolH$R%@QPrnAjjuK>CGJ@NQaLnNciNpD9G>m;d<@KYb35E(b)FA6wFgxMe9mCo zcm(PlcYTaX`#LXgyN(+7rrwmP>7(46VCeOkjQ0}d3)b>+$~+~a34S9(m2xFl5BAE8 z#cjGTH+iQ8dwNEGb%8NEbG)uwP+$9*ukF(}GMClOUDHN#Nv00$vm=FVhKK$^M~zlI zoq6_5SN?l~pBMJwiO(Z2-1S;GJtUWOQkZt)P#XK7WF~=s-RhN|W@g&B!fV_Y<{tX` zRvXFX{_W(;z1laCcWdQ4Z}%sa0bA@D%v6J)I4@DGW1I31MvUS4rSv|9TD7k&jLG%> z#OJ9b{y~?q(MW8#@713I7vo7k4^4~n6i#`>zH+T8HJ0TwIPM^0K{t6;uzu@$Fn=zm zjt&}oj{z^`lP9iu!>!B$#*BKe`gXguWuT2a)L@Qv|x&GPjoN z_MVUaeUjC>f)aro!L!%(3<7BhBK?w`U=LDIN^rb+D#8iy2t}$Ws;FdwkG=T!0U~Vx zQPWX%MZifE8AwK}Qph@eQ;mwaSIz;aKE}8*2~bIc>hww zt6WYex)A1o^A0o?jB+o!D4R|CKK#~yY?+Y<&KbF zuA?R35z?4ZlK^arjM2Is)26fN-Bqnct;3gGF5M(=B1O+T=h4ysuaq-^YU)Y@Fjzng zuC*dqEe{k0CFCu6F9`}5*$WAZD4>Fd5FjcfgkS{+a0gkGQK!nF)zM0E2hkC!)*>p{ zaf*V1`yh2jsgB|72m-~QW!xKXY8dE1Qq2cwoev)Y-HQnlN8 zSHV;JF3rXGlzyzzjVUbK$f7(O<)B;J^76+l9TynT;$+|AnEmOL!pHvR>e~533JO=s zam(5oy#FDeIcC8N+#A|A?CrEmh|KA15(GrJwcAv7j9fcxD{sV;`DQPt(3Z8X^{1>; zNY$%zZ`3Ehu;rAnY_4k_uT9r<9f%pjuhpErX>-2A=yM~<|@#tmS;LuS8nf<$mggIH<{ZsP zeEUv|i3`jIt`;gS!)yM8Ta2uLFLM2gf9x*}K0l2H z&x-8HxtUG;7EVf|<9g2$fvh{N^-=EL410%jUk(?Gv$j|iUwd-sTIC`<^vs>|s1f%( z%M+heWLq2%)!GiQ8+nXd?z9)a9I>UjAb?^Ml z8{~73B^*mQ%D&Nb`WNp`PK>61i%alO(-XMb^Rnn*EkP4FsXuN@{9bkhk1{%2yQeUU zwS21GO7Fz&ul~HS^lkU?i#9J#k1NYv*)e?GSs&JmwAIXX@4SGc-p zpIY)D=kv*b8MwgP>FZcGAFJ-nq#cf7?3(t+x9uOjJ`>zI=ZCELyWf14?ngb$aOH}A zicXDX-_Ef2C5C(M@Jv;BaVJ`fs^*jyuMJxwExPa4MV^y4+uS^rTe0Vx<%DfkqOf`k z@2|g}2rt_AZP;mU(YEW3F0Aoo)Xh`$&!_Bj*SMeLtl9he`4L`4&QHr`oQTR*1zFr`L?Kws-ja?f6NahbD&DKe%2pB{OE$);R%G_JFTS zM`LDHe>^((pnLJ%i%~1@t=g3*jo*7X`NhJL$|W_gE)2fPSacma)_%HVO3Tul&te*x zZ41L0*$(gPZ@;rMO$vdyC;FsJlcdVDq!ms;sDIS-s7XSVT(Vdc&f@q4`OZvFRm+ls zlH(N^V!DK?;$$ksQE>^11FVT$o$iWp!ej|@IHaA9xd?&jA_K|HG^KK}ED3>6aFkTk zQ@A1^$X6s$V%ib~^a&c%SC%SO$(3q_3L~^=7SuLY-2;mW@;BQO$i zTEbA*`t#_9AerCt2;-pfpoJt^je1H2bc;3=VTjb?mke+flMw3ZpzgS^p4ScnEyV*;E(2pnagOrDSk>>2P_XOJ3WlW?LQL0%G4F$&y8 z-A(I(%di=S#@G}|VC(ucjQuGa>8O6 zsB%P-r4jNt0Lbtu3Os}a&&F`W5F4b06R9(ohP?Iqr%eFzZ!VY%VI+BcFpA6cUh?6R zrkM}^L}QEq)EI^b34qi7a0vvSXiUJ%llibLAo9{(t1YRGNwrSGO>&FJ3PA>Q--n{fkz#g+Fmnu;2$ zLkjHvb@hl#-Od^8<{#7eWet~g$zG}^r^aPSu5NUnLLSxtXeeEG99;2E8-cd{Bftjy7V-aPRT?$z{TpUp`^&$RC?y zvtj78ajWwGw$8@v@zJ40rVbgtUHB-{V3mbA7A%j`F6}*=1XO0v-0T68^-c-WFr_{K z&{E44$-am`=Hlx?;v|U^B#jeXhM*_m{0TTd!4-BgqzBuTi%6nYE>@@)%VV)6lsoNC zxM7KEwKCO%!%2Gg$X%gIfJQ(oi%XLN#>=pkGA0WdslGBuW?vbJGbS>e#7+7maW16KXbgfcK-d{$nSg+ur>~6SnZ$(;+3c&w<8w`Y z!)X&4kAw}$cnoB+i7k9CY|_Sh1dnSH1A<4w(%DB(z;F#@YLx`Gc9oW+r+xX-)+)R- mg((!M#I@|Wneq%76bLGsRJBB<)|Lz*fOm72v$J2QKkGjM))kNd literal 0 HcmV?d00001 diff --git a/src/utils/train.py b/src/utils/train.py index 8902b3b..a65b84d 100644 --- a/src/utils/train.py +++ b/src/utils/train.py @@ -115,7 +115,7 @@ def train(self, epochs, log_interval=100): mrr10, hit10 = evaluate(self.model, self.test_loader, self.device, cutoff=10) mrr20, hit20 = evaluate(self.model, self.test_loader, self.device) - + wandb.log({"hit@20": hit20, "mrr@20": mrr20}) print(f'Epoch {self.epoch}: MRR@10 = {mrr10 * 100:.3f}%, Hit@10 = {hit10 * 100:.3f}%, MRR@20 = {mrr20 * 100:.3f}%, Hit@20 = {hit20 * 100:.3f}%') @@ -129,6 +129,6 @@ def train(self, epochs, log_interval=100): max_hit10 = max(max_hit10, hit10) max_mrr20 = max(max_mrr20, mrr20) max_hit20 = max(max_hit20, hit20) - wandb.log({"hit@20": max_hit20, "mrr@20": max_mrr20}) + self.epoch += 1 return max_mrr10, max_mrr20, max_hit10, max_hit20