cvdpkgfl/DGM_Z_Version
import os import torch import numpy as np
import torchgeometric from torch import nn from torch.nn import Module, ModuleList, Sequential from torchgeometric.nn import EdgeConv, DenseGCNConv, DenseGraphConv, GCNConv, GATConv from torch.utils.data import DataLoader import torch_scatter
import pytorch_lightning as pl from argparse import Namespace
import torch import numpy import torch import pykeops from pykeops.torch import LazyTensor from torch.nn import Module, ModuleList, Sequential from torch import nn
#Euclidean distance def pairwiseeuclideandistances(x, dim=-1): dist = torch.cdist(x,x)**2 return dist, x
#Poincarè disk distance r=1 (Hyperbolic)
def pairwisepoincaredistances(x, dim=-1): xnorm = (x**2).sum(dim,keepdim=True) xnorm = (xnorm.sqrt()-1).relu() + 1 x = x/(xnorm(1+1e-2)) x_norm = (x*2).sum(dim,keepdim=True)
pq = torch.cdist(x,x)*2 dist = torch.arccosh(1e-6+1+2pq/((1-xnorm)*(1-xnorm.transpose(-1,-2))))**2 return dist, x
class DGMd(nn.Module): def init(self, embedf, k=5, distance="euclidean", sparse=True): super(DGMd, self).init_()
self.sparse=sparse
self.temperature = nn.Parameter(torch.tensor(1. if distance=="hyperbolic" else 4.).float()) self.embedf = embedf self.centroid=None self.scale=None self.k = k self.distance = distance
self.debug=False
def forward(self, x, A, notused=None, fixedges=None): if x.shape[0]==1: x = x[0] x = self.embedf(x,A) if x.dim()==2: x = x[None,...]
if self.training: if fixedges is not None: return x, fixedges, torch.zeros(fixedges.shape[0],fixedges.shape[-1]//self.k,self.k,dtype=torch.float,device=x.device) #sampling here edgeshat, logprobs = self.samplewithout_replacement(x)
else: with torch.nograd(): if fixedges is not None: return x, fixedges, torch.zeros(fixedges.shape[0],fixedges.shape[-1]//self.k,self.k,dtype=torch.float,device=x.device) #sampling here edgeshat, logprobs = self.samplewithoutreplacement(x)
if self.debug: if self.distance=="euclidean": D, x = pairwiseeuclideandistances(x) if self.distance=="hyperbolic": D, x = pairwisepoincaredistances(x)
self.D = (D * torch.exp(torch.clamp(self.temperature,-5,5))).detach().cpu() self.edgeshat=edgeshat.detach().cpu() self.logprobs=logprobs.detach().cpu()
self.x=x
return x, edges_hat, logprobs
def samplewithoutreplacement(self, x):
b,n,_ = x.shape
if self.distance=="euclidean": Gi = LazyTensor(x[:, :, None, :]) # (M**2, 1, 2) Xj = LazyTensor(x[:, None, :, :]) # (1, N, 2)
mD = ((Gi - Xj) ** 2).sum(-1)
#argKmin already add gumbel noise lq = mD * torch.exp(torch.clamp(self.temperature,-5,5)) indices = lq.argKmin(self.k, dim=1)
x1 = torch.gather(x, -2, indices.view(indices.shape[0],-1)[...,None].repeat(1,1,x.shape[-1])) x2 = x[:,:,None,:].repeat(1,1,self.k,1).view(x.shape[0],-1,x.shape[-1]) logprobs = (-(x1-x2).pow(2).sum(-1) * torch.exp(torch.clamp(self.temperature,-5,5))).reshape(x.shape[0],-1,self.k)
if self.distance=="hyperbolic": pass xnorm = (x**2).sum(-1,keepdim=True) xnorm = (xnorm.sqrt()-1).relu() + 1 x = x/(xnorm(1+1e-2)) #safe distance to the margin x_norm = (x*2).sum(-1,keepdim=True)
Gi = LazyTensor(x[:, :, None, :]) # (M**2, 1, 2) Xj = LazyTensor(x[:, None, :, :]) # (1, N, 2)
Gi2 = LazyTensor(1-xnorm[:, :, None, :]) # (M**2, 1, 2) Xj2 = LazyTensor(1-xnorm[:, None, :, :]) # (1, N, 2)
pq = ((Gi - Xj) * 2).sum(-1) N = (G_i2X_j2) XX = (1e-6+1+2pq/N) mD = (XX+(XX2-1).sqrt()).log()*2
lq = mD * torch.exp(torch.clamp(self.temperature,-5,5)) indices = lq.argKmin(self.k, dim=1)
x1 = torch.gather(x, -2, indices.view(indices.shape[0],-1)[...,None].repeat(1,1,x.shape[-1])) x2 = x[:,:,None,:].repeat(1,1,self.k,1).view(x.shape[0],-1,x.shape[-1])
x1n = torch.gather(xnorm, -2, indices.view(indices.shape[0],-1)[...,None].repeat(1,1,xnorm.shape[-1])) x2n = xnorm[:,:,None,:].repeat(1,1,self.k,1).view(x.shape[0],-1,xnorm.shape[-1])
pq = (x1-x2).pow(2).sum(-1) pqn = ((1-x1n)*(1-x2n)).sum(-1) XX = 1e-6+1+2pq/pqn dist = torch.log(XX+(XX2-1).sqrt())2 logprobs = (-dist torch.exp(torch.clamp(self.temperature,-5,5))).reshape(x.shape[0],-1,self.k)
if self.debug: self._x=x.detach().cpu()+0
rows = torch.arange(n).view(1,n,1).to(x.device).repeat(b,1,self.k) edges = torch.stack((indices.view(b,-1),rows.view(b,-1)),-2)
if self.sparse: return (edges+(torch.arange(b).to(x.device)*n)[:,None,None]).transpose(0,1).reshape(2,-1), logprobs return edges, logprobs
class DGMc(nn.Module): inputdim = 4 debug=False
def _init(self, embedf, k=None, distance="euclidean"): super(DGMc, self).init() self.temperature = nn.Parameter(torch.tensor(1).float()) self.threshold = nn.Parameter(torch.tensor(0.5).float()) self.embedf = embed_f self.centroid=None self.scale=None self.distance = distance
self.scale = nn.Parameter(torch.tensor(-1).float(),requiresgrad=False) self.centroid = nn.Parameter(torch.zeros((1,1,DGMc.inputdim)).float(),requiresgrad=False)
def forward(self, x, A, not_used=None, fixedges=None):
x = self.embed_f(x,A)
# estimate normalization parameters if self.scale <0: self.centroid.data = x.mean(-2,keepdim=True).detach() self.scale.data = (0.9/(x-self.centroid).abs().max()).detach()
if self.distance=="hyperbolic": D, x = pairwisepoincaredistances((x-self.centroid)*self.scale) else: D, x = pairwiseeuclideandistances((x-self.centroid)*self.scale)
A = torch.sigmoid(self.temperature*(self.threshold.abs()-D))
if DGMc.debug: self.A = A.data.cpu() self.x = _x.data.cpu()
self.A=A
A = A/A.sum(-1,keepdim=True)
return x, A, None
class MLP(nn.Module): def _init(self, layerssize,finalactivation=False, dropout=0): super(MLP, self).init() layers = [] for li in range(1,len(layerssize)): if dropout>0: layers.append(nn.Dropout(dropout)) layers.append(nn.Linear(layerssize[li-1],layerssize[li])) if li==len(layerssize)-1 and not finalactivation: continue layers.append(nn.LeakyReLU(0.1))
self.MLP = nn.Sequential(*layers)
def forward(self, x, e=None): x = self.MLP(x) return x
class Identity(nn.Module): def _init(self,retparam=None): self.retparam=retparam super(Identity, self).init_()
def forward(self, *params): if self.retparam is not None: return params[self.retparam] return params
from torch.nn import Module, ModuleList, Sequential from torch import nn
class DGMd(nn.Module): def init(self, embedf, k=5, distance=pairwiseeuclideandistances, sparse=True): super(DGMd, self).init_()
self.sparse=sparse
self.temperature = nn.Parameter(torch.tensor(1. if distance=="hyperbolic" else 4.).float()) self.embedf = embedf self.centroid=None self.scale=None self.k = k
self.debug=False if distance == 'euclidean': self.distance = pairwiseeuclideandistances else: self.distance = pairwisepoincaredistances
def forward(self, x, A, notused=None, fixedges=None): x = self.embedf(x,A)
if self.training: if fixedges is not None: return x, fixedges, torch.zeros(fixedges.shape[0],fixedges.shape[-1]//self.k,self.k,dtype=torch.float,device=x.device)
D, _x = self.distance(x)
#sampling here edgeshat, logprobs = self.samplewithout_replacement(D)
else: with torch.nograd(): if fixedges is not None: return x, fixedges, torch.zeros(fixedges.shape[0],fixedges.shape[-1]//self.k,self.k,dtype=torch.float,device=x.device) D, x = self.distance(x)
#sampling here edgeshat, logprobs = self.samplewithout_replacement(D)
if self.debug: self.D = D self.edgeshat=edgeshat self.logprobs=logprobs self.x=x
return x, edges_hat, logprobs
def samplewithoutreplacement(self, logits): b,n,_ = logits.shape
logits = logits torch.exp(self.temperature10)
logits = logits * torch.exp(torch.clamp(self.temperature,-5,5))
q = torch.rand_like(logits) + 1e-8 lq = (logits-torch.log(-torch.log(q))) logprobs, indices = torch.topk(-lq,self.k)
rows = torch.arange(n).view(1,n,1).to(logits.device).repeat(b,1,self.k) edges = torch.stack((indices.view(b,-1),rows.view(b,-1)),-2)
if self.sparse: return (edges+(torch.arange(b).to(logits.device)*n)[:,None,None]).transpose(0,1).reshape(2,-1), logprobs return edges, logprobs
class DGMModel(pl.LightningModule): def init(self, hparams): super(DGMModel,self)._init_()
if type(hparams) is not Namespace: hparams = Namespace(**hparams)
self.hparams=hparams
self.savehyperparameters(hparams) convlayers = hparams.convlayers fclayers = hparams.fclayers dgmlayers = hparams.dgm_layers k = hparams.k
self.graphf = ModuleList() self.nodeg = ModuleList() for i,(dgml,convl) in enumerate(zip(dgmlayers,convlayers)): if len(dgml)>0: if 'ffun' not in hparams or hparams.ffun == 'gcn': self.graphf.append(DGMd(GCNConv(dgml[0],dgml[-1]),k=hparams.k,distance=hparams.distance)) if hparams.ffun == 'gat': self.graphf.append(DGMd(GATConv(dgml[0],dgml[-1]),k=hparams.k,distance=hparams.distance)) if hparams.ffun == 'mlp': self.graphf.append(DGMd(MLP(dgml),k=hparams.k,distance=hparams.distance)) if hparams.ffun == 'knn': self.graphf.append(DGMd(Identity(retparam=0),k=hparams.k,distance=hparams.distance))
self.graphf.append(DGMd(GCNConv(dgml[0],dgml[-1]),k=hparams.k,distance=hparams.distance))
else: self.graph_f.append(Identity())
if hparams.gfun == 'edgeconv': convl=convl.copy() convl[0]=convl[0]*2 self.nodeg.append(EdgeConv(MLP(convl), hparams.pooling)) if hparams.gfun == 'gcn': self.nodeg.append(GCNConv(convl[0],convl[1])) if hparams.gfun == 'gat': self.nodeg.append(GATConv(convl[0],convl[1]))
self.fc = MLP(fclayers, finalactivation=False) if hparams.prefc is not None and len(hparams.prefc)>0: self.prefc = MLP(hparams.prefc, finalactivation=True) self.avgaccuracy = None
#torch lightning specific self.automatic_optimization = False self.debug=False
def forward(self,x, edges=None): if self.hparams.prefc is not None and len(self.hparams.prefc)>0: x = self.pre_fc(x)
graphx = x.detach() lprobslist = [] for f,g in zip(self.graphf, self.nodeg): graphx,edges,lprobs = f(graph_x,edges,None) b,n,d = x.shape
edges, = torchgeometric.utils.removeselfloops(edges)
edges, = torchgeometric.utils.addselfloops(edges)
self.edges=edges x = torch.nn.functional.relu(g(torch.dropout(x.view(-1,d), self.hparams.dropout, train=self.training), edges)).view(b,n,-1) graphx = torch.cat([graphx,x.detach()],-1) if lprobs is not None: lprobslist.append(lprobs)
return self.fc(x),torch.stack(lprobslist,-1) if len(lprobslist)>0 else None
def configure_optimizers(self): optimizer = torch.optim.Adam(self.parameters(), lr=self.hparams.lr) return optimizer
def trainingstep(self, trainbatch, batch_idx):
optimizer = self.optimizers(useploptimizer=True) optimizer.zero_grad()
X, y, mask, edges = train_batch edges = edges[0]
assert(X.shape[0]==1) #only works in transductive setting mask=mask[0]
pred,logprobs = self(X,edges)
trainpred = pred[:,mask.to(torch.bool),:] trainlab = y[:,mask.to(torch.bool),:]
train_w = weight[None,mask.to(torch.bool)]
#loss = torch.nn.functional.crossentropy(trainpred.view(-1,trainpred.shape[-1]),trainlab.argmax(-1).flatten()) loss = torch.nn.functional.binarycrossentropywithlogits(trainpred,trainlab) loss.backward()
correctt = (trainpred.argmax(-1) == train_lab.argmax(-1)).float().mean().item()
#GRAPH LOSS if logprobs is not None: corrpred = (trainpred.argmax(-1)==trainlab.argmax(-1)).float().detach() wronpred = (1-corr_pred)
if self.avgaccuracy is None: self.avgaccuracy = torch.oneslike(corrpred)*0.5
pointw = (self.avgaccuracy-corrpred)#*(1*corrpred + self.k(1-corr_pred)) graph_loss = point_w logprobs[:,mask.to(torch.bool),:].exp().mean([-1,-2])
graphloss = graphloss.mean()# + self.graphf[0].Pr.abs().sum()*1e-3 graphloss.backward()
self.log('traingraphloss', graphloss.detach().cpu()) if(self.debug): self.pointw = point_w.detach().cpu()
self.avgaccuracy = self.avgaccuracy.to(corrpred.device)*0.95 + 0.05*corrpred
optimizer.step()
self.log('trainacc', correctt) self.log('train_loss', loss.detach().cpu())
def teststep(self, trainbatch, batchidx): X, y, mask, edges = trainbatch edges = edges[0]
assert(X.shape[0]==1) #only works in transductive setting mask=mask[0] pred,logprobs = self(X,edges) pred = pred.softmax(-1) for i in range(1,self.hparams.testeval): pred,logprobs = self(X,edges) pred+=pred.softmax(-1) testpred = pred[:,mask.to(torch.bool),:] testlab = y[:,mask.to(torch.bool),:] correctt = (testpred.argmax(-1) == testlab.argmax(-1)).float().mean().item() loss = torch.nn.functional.binarycrossentropywithlogits(testpred,testlab) self.log('test_loss', loss.detach().cpu())
self.log('testgraphloss', loss.detach().cpu())
self.log('testacc', 100*correctt)
def validationstep(self, trainbatch, batchidx): X, y, mask, edges = trainbatch edges = edges[0]
assert(X.shape[0]==1) #only works in transductive setting mask=mask[0]
pred,logprobs = self(X,edges) pred = pred.softmax(-1) for i in range(1,self.hparams.testeval): pred,logprobs = self(X,edges) pred+=pred_.softmax(-1)
testpred = pred[:,mask.to(torch.bool),:] testlab = y[:,mask.to(torch.bool),:] correctt = (testpred.argmax(-1) == testlab.argmax(-1)).float().mean().item() loss = torch.nn.functional.binarycrossentropywithlogits(testpred,test_lab)
self.log('valloss', loss.detach()) self.log('valacc', 100*correct_t)
####### visualizations ###########
try:
self.graph_f[0].debug=True
pred,logprobs = self(X)
self.graph_f[0].debug=False
x = self.graph_f[0].x[0].detach()
c = torch.argmax(y,-1)
D = self.graph_f[0].distance(x)[0]
D.diagonal().fill_(0)
sidx = torch.argsort( (c[0]+1)10 + (mask+1)1)
P = torch.exp(-D[sidx,:][:,sidx]*torch.clamp(self.graph_f[0].temperature.detach().cpu(),-5,5).exp())#>0.001
img = PIL.Image.fromarray((P*255).byte().detach().cpu().numpy())
img = img.resize((512,512), PIL.Image.ANTIALIAS)
I = wandb.Image(img, caption="adj")
self.logger.experiment.log({'adj': [I]})
except:
pass
import sys import torch import pickle import numpy as np import os.path as osp import torch from torchgeometric.datasets import Planetoid import torchgeometric.transforms as T
class UKBBAgeDataset(torch.utils.data.Dataset): """Face Landmarks dataset."""
def _init(self, fold=0, train=True, samplesperepoch=100, device='cpu'): with open('data/UKBB.pickle', 'rb') as f: X,y,trainmask,testmask, weight = pickle.load(f) # Load the data
self.X = torch.fromnumpy(X[:,:,fold]).float().to(device) self.y = torch.fromnumpy(y[:,:,fold]).float().to(device) self.weight = torch.fromnumpy(np.squeeze(weight[:1,fold])).float().to(device) if train: self.mask = torch.fromnumpy(trainmask[:,fold]).to(device) else: self.mask = torch.fromnumpy(testmask[:,fold]).to(device)
self.samplesperepoch = samplesperepoch
def _len(self): return self.samplesper_epoch
def _getitem_(self, idx): return self.X,self.y,self.mask
class TadpoleDataset(torch.utils.data.Dataset): """Face Landmarks dataset."""
def _init(self, fold=0, train=True, samplesperepoch=100, device='cpu',full=False): with open('data/tadpoledata.pickle', 'rb') as f: X,y,trainmask,testmask, weight_ = pickle.load(f) # Load the data
if not full: X = X[...,:30,:] # For DGM we use modality 1 (M1) for both node representation and graph learning.
self.nfeatures = X.shape[-2] self.numclasses = y.shape[-2]
self.X = torch.fromnumpy(X[:,:,fold]).float().to(device) self.y = torch.fromnumpy(y[:,:,fold]).float().to(device) self.weight = torch.fromnumpy(np.squeeze(weight[:1,fold])).float().to(device) if train: self.mask = torch.fromnumpy(trainmask[:,fold]).to(device) else: self.mask = torch.fromnumpy(testmask[:,fold]).to(device)
self.samplesperepoch = samplesperepoch
def _len(self): return self.samplesper_epoch
def _getitem_(self, idx): return self.X,self.y,self.mask, [[]]
class TadpoleDataset(torch.utils.data.Dataset):
"""Face Landmarks dataset."""
def _init(self, fold=0, split='train', samplesper_epoch=100, device='cpu'):
with open('data/train_data.pickle', 'rb') as f:
X,y,trainmask,testmask, weight_ = pickle.load(f) # Load the data
X = X[...,:30,:] # For DGM we use modality 1 (M1) for both node representation and graph learning.
self.X = torch.fromnumpy(X[:,:,fold]).float().to(device)
self.y = torch.fromnumpy(y[:,:,fold]).float().to(device)
self.weight = torch.fromnumpy(np.squeeze(weight[:1,fold])).float().to(device)
# split train set in train/val
trainmask = trainmask_[:,fold]
nval = int(train_mask.sum()*0.2)
validxs = np.random.RandomState(1).choice(np.nonzero(trainmask.flatten())[0],(nval,),replace=False)
trainmask[validxs] = 0;
valmask = trainmask*0
valmask[validxs] = 1
print('DATA STATS: train: %d val: %d' % (trainmask.sum(),valmask.sum()))
if split=='train':
self.mask = torch.fromnumpy(trainmask).to(device)
if split=='val':
self.mask = torch.fromnumpy(valmask).to(device)
if split=='test':
self.mask = torch.fromnumpy(testmask_[:,fold]).to(device)
self.samplesperepoch = samplesperepoch
def _len_(self):
return self.samplesperepoch
def _getitem_(self, idx):
return self.X,self.y,self.mask
def getplanetoiddataset(name, normalizefeatures=True, transform=None, split="complete"): path = osp.join('.', 'data', name) if split == 'complete': dataset = Planetoid(path, name) dataset[0].trainmask.fill(False) dataset[0].trainmask[:dataset[0].numnodes - 1000] = 1 dataset[0].valmask.fill(False) dataset[0].valmask[dataset[0].numnodes - 1000:dataset[0].numnodes - 500] = 1 dataset[0].testmask.fill(False) dataset[0].testmask[dataset[0].numnodes - 500:] = 1 else: dataset = Planetoid(path, name, split=split) if transform is not None and normalizefeatures: dataset.transform = T.Compose([T.NormalizeFeatures(), transform]) elif normalizefeatures: dataset.transform = T.NormalizeFeatures() elif transform is not None: dataset.transform = transform return dataset
def onehotembedding(labels, numclasses): y = torch.eye(numclasses) return y[labels]
class PlanetoidDataset(torch.utils.data.Dataset): def _init(self, split='train', samplesperepoch=100, name='Cora', device='cpu'): dataset = getplanetoiddataset(name) self.X = dataset[0].x.float().to(device) self.y = onehotembedding(dataset[0].y,dataset.numclasses).float().to(device) self.edgeindex = dataset[0].edgeindex.to(device) self.nfeatures = dataset[0].numnodefeatures self.numclasses = dataset.num_classes
if split=='train': self.mask = dataset[0].trainmask.to(device) if split=='val': self.mask = dataset[0].valmask.to(device) if split=='test': self.mask = dataset[0].test_mask.to(device)
self.samplesperepoch = samplesperepoch
def _len(self): return self.samplesper_epoch
def _getitem(self, idx): return self.X,self.y,self.mask,self.edgeindex
os.environ["CUDAVISIBLEDEVICES"]="0";
import pickle import numpy as np
import torch from torch.utils.data import DataLoader import pytorch_lightning as pl
from argparse import ArgumentParser from pytorchlightning.callbacks import ModelCheckpoint, EarlyStopping from pytorchlightning.loggers import TensorBoardLogger
def runtrainingprocess(run_params):
traindata = None testdata = None
if runparams.dataset in ['Cora', 'CiteSeer', 'PubMed']: traindata = PlanetoidDataset(split='train', name=runparams.dataset, device='cuda') valdata = PlanetoidDataset(split='val', name=runparams.dataset, samplesperepoch=1) testdata = PlanetoidDataset(split='test', name=runparams.dataset, samplesper_epoch=1)
if runparams.dataset == 'tadpole': traindata = TadpoleDataset(fold=runparams.fold,train=True, device='cuda') valdata = testdata = TadpoleDataset(fold=runparams.fold, train=False,samplesperepoch=1)
if traindata is None: raise Exception("Dataset %s not supported" % runparams.dataset)
trainloader = DataLoader(traindata, batchsize=1,numworkers=0) valloader = DataLoader(valdata, batchsize=1) testloader = DataLoader(testdata, batchsize=1)
class MyDataModule(pl.LightningDataModule): def setup(self,stage=None): pass def traindataloader(self): return trainloader def valdataloader(self): return valloader def testdataloader(self): return testloader
#configure input feature size if runparams.prefc is None or len(runparams.prefc)==0: if len(runparams.dgmlayers[0])>0: runparams.dgmlayers[0][0]=traindata.nfeatures runparams.convlayers[0][0]=traindata.nfeatures else: runparams.prefc[0]=traindata.nfeatures runparams.fclayers[-1] = traindata.numclasses
model = DGMModel(runparams)
checkpointcallback = ModelCheckpoint( savelast=True, savetopk=1, verbose=True, monitor='valloss', mode='min' ) earlystopcallback = EarlyStopping( monitor='valloss', mindelta=0.00, patience=20, verbose=False, mode='min') callbacks = [checkpointcallback,earlystopcallback]
if valdata==testdata: callbacks = None
logger = TensorBoardLogger("logs/") trainer = pl.Trainer.fromargparseargs(run_params,logger=logger, callbacks=callbacks)
trainer.fit(model, datamodule=MyDataModule()) trainer.test()
if _name == "main_":
parser = ArgumentParser() parser = pl.Trainer.addargparseargs(parser) params = parser.parseargs(['--gpus','1', '--logeverynsteps','100', '--maxepochs','100', '--progressbarrefreshrate','10', '--checkvaleverynepoch','1']) parser.addargument("--numgpus", default=10, type=int)
parser.addargument("--dataset", default='Cora') parser.addargument("--fold", default='0', type=int) #Used for k-fold cross validation in tadpole/ukbb
parser.addargument("--convlayers", default=[[32,32],[32,16],[16,8]], type=lambda x :eval(x)) parser.addargument("--dgmlayers", default= [[32,16,4],[],[]], type=lambda x :eval(x)) parser.addargument("--fclayers", default=[8,8,3], type=lambda x :eval(x)) parser.addargument("--prefc", default=[-1,32], type=lambda x :eval(x))
parser.addargument("--gfun", default='gcn') parser.addargument("--ffun", default='gcn') parser.addargument("--k", default=5, type=int) parser.addargument("--pooling", default='add') parser.add_argument("--distance", default='euclidean')
parser.addargument("--dropout", default=0.0, type=float) parser.addargument("--lr", default=1e-2, type=float) parser.addargument("--testeval", default=10, type=int)
parser.setdefaults(defaultrootpath='./log') params = parser.parseargs(namespace=params)
runtrainingprocess(params)
