CoolFace
Modelpublic

cvdpkgfl/DGM_Z_Version

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
Model Card

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)