Commit c76fb4d9 authored by miaecle's avatar miaecle
Browse files

import low data benchmark

parent cbaa4c3a
Loading
Loading
Loading
Loading
+3 −1
Original line number Diff line number Diff line
@@ -501,7 +501,7 @@ if __name__ == '__main__':
    models = ['tf', 'tf_robust', 'logreg', 'graphconv', 'tf_regression']
  if len(datasets) == 0:
    datasets = ['tox21', 'sider', 'muv', 'toxcast', 'pcba', 
                'kaggle', 'delaney']
                'delaney', 'kaggle']

  #input hyperparameters
  #tf: dropouts, learning rate, layer_sizes, weight initial stddev,penalty,
@@ -548,6 +548,8 @@ if __name__ == '__main__':
                                       model=model, split=split, 
                                       verbosity='high', out_path='.')
      else:
        if dataset in ['kaggle']:
          datasets.remove('kaggle') #kaggle only needs to be run once
        for model in models:
          if model in ['tf_regression']:
             benchmark_loading_datasets(base_dir_o, hps, dataset=dataset, 
+280 −0
Original line number Diff line number Diff line
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Created on Thu Dec  8 16:48:05 2016

@author: Michael Wu

Low data benchmark test
Giving performances of: Siamese, attention-based embedding, residual embedding
                    
on datasets: muv, sider, tox21

time estimation listed in README file
"""
from __future__ import print_function
from __future__ import division
from __future__ import unicode_literals

import sys
import os
import numpy as np
import shutil
import time
import deepchem as dc
import tensorflow as tf
import argparse
from keras import backend as K

from low_data.datasets import load_tox21_convmol
from low_data.datasets import load_muv_convmol
from low_data.datasets import load_sider_convmol


def low_data_benchmark_loading_datasets(hyper_parameters, cross_valid=False,
                               dataset='tox21', model='siamese', split='task',
                               verbosity='high', out_path='.'):
  """
  Loading dataset for low data benchmark test
  
  Parameters
  ----------
  hyper_parameters : dict of list
      hyper parameters including batch size, learning rate, etc.

  cross_valid : boolean, optional (default=False)
      whether implement cross validation on datasets
  
  dataset : string, optional (default='tox21')
      choice of which dataset to use, should be: tox21, muv, sider
      
  model : string,  optional (default='siamese')
      choice of which model to use, should be: siamese, attn, res
  
  split : string,  optional (default='task')
      choice of splitter function, only task splitter is supported

  out_path : string, optional(default='.')
      path of result file
  """
  # Check input
  if dataset in ['muv','tox21','sider']:
    mode = 'classification'
  else:
    raise ValueError('Dataset not supported')
  
  if not model in ['siamese','attn','res']:
    raise ValueError('Model not supported')

  if not split in ['task']:
    raise ValueError('Only task splitter is supported')
  
  loading_functions = {'tox21': load_tox21_convmol, 'muv': load_muv_convmol,
                       'sider': load_sider_convmol}
  
  print('-------------------------------------')
  print('Low data benchmark %s on dataset: %s' % (model, dataset))
  print('-------------------------------------')
  time_start = time.time()
  #loading datasets
  tasks, all_dataset, transformers = loading_functions[dataset]()
  time_finish_loading = time.time()
  
  #defining splitter function
  splitters = {'task': dc.splits.TaskSplitter()}
  splitter = splitters[split]

  #running model
  for count, hp in enumerate(hyper_parameters[model]):
    # Loading general settings
    # Number of folds for split 
    K = hp['K']
    n_feat = hp['n_feat']

    fold_datasets = splitter.k_fold_split(all_dataset, K)
    if cross_valid:
        num_iter = K # K iterations for cross validation
    else:
        num_iter = 1
    for count_iter in range(num_iter):
      # Assembling train and valid datasets
      train_folds = fold_datasets[:num_iter-count_iter-1] + fold_datasets[num_iter-count_iter:]
      train_dataset = dc.splits.merge_fold_datasets(train_folds)
      valid_dataset = fold_datasets[num_iter-count_iter-1]

      time_start_fitting = time.time()
      valid_scores = low_data_benchmark_classification(
                         train_dataset, valid_dataset, hp, n_feat,
                         model=model, verbosity=verbosity)
      time_finish_fitting = time.time() 
      with open(os.path.join(out_path, 'results.csv'),'a') as f:
        f.write('\n'+str(count)+','+str(count_iter)+',')
        f.write(dataset+','+model+',')
        f.write('valid,')
        for i in valid_scores:
          f.write(str(valid_scores[i])+',')
        f.write('time_for_running,'+
              str(time_finish_fitting-time_start_fitting)+',')

  return None

def low_data_benchmark_classification(train_dataset, valid_dataset, 
                                      hyper_parameters, n_features, 
                                      model='siamese', seed=123, 
                                      verbosity='high'):
  """
  Calculate low data benchmark performance
  
  Parameters
  ----------
  train_dataset : dataset struct
      loaded dataset, ConvMol struct, used for training
      
  valid_dataset : dataset struct
      loaded dataset, ConvMol struct, used for validation
  
  hyper_parameters : dict
      hyper parameters including batch size, learning rate, etc.
 
  n_features : integer
      number of features, or length of binary fingerprints
  
  model : string,  optional (default='siamese')
      choice of which model to use, should be: siamese, attn, res

  Returns
  -------
  scores : dict
	predicting results(AUC) on valid set

  """
  scores = {}
  
  # Initialize metrics
  classification_metric = dc.metrics.Metric(dc.metrics.roc_auc_score, np.mean,
                                            verbosity=verbosity,
                                            mode="classification")

  assert model in ['siamese','attn','res']

  # Loading hyperparameters
  # num positive/negative ligands
  n_pos = hyper_parameters['n_pos']
  n_neg = hyper_parameters['n_neg']
  # Set batch sizes for network
  test_batch_size = hyper_parameters['test_batch_size']
  support_batch_size = n_pos + n_neg
  # Model structure
  n_filters = hyper_parameters['n_filters']
  n_fully_connected_nodes = hyper_parameters['n_fully_connected_nodes']

  # Traning settings
  nb_epochs = hyper_parameters['nb_epochs']
  n_train_trials = hyper_parameters['n_train_trials']
  n_eval_trials = hyper_parameters['n_eval_trials'] 

  learning_rate = hyper_parameters['learning_rate']

  g = tf.Graph()
  sess = tf.Session(graph=g)
  K.set_session(sess)
  # Building graph convolution model
  with g.as_default():
    tf.set_random_seed(seed)
    support_graph = dc.nn.SequentialSupportGraph(n_features)
    
    for count, n_filter in enumerate(n_filters):
      support_graph.add(dc.nn.GraphConv(int(n_filter), activation='relu'))
      support_graph.add(dc.nn.GraphPool())
    
    for count, n_fcnode in enumerate(n_fully_connected_nodes):
      support_graph.add(dc.nn.Dense(int(n_fcnode), activation='tanh'))
      
    support_graph.add_test(dc.nn.GraphGather(test_batch_size, 
                                             activation='tanh'))
    support_graph.add_support(dc.nn.GraphGather(support_batch_size, 
                                                activation='tanh'))
    if model in ['siamese']:
      pass
    elif model in ['attn']:
      max_depth = hyper_parameters['max_depth']
      support_graph.join(dc.nn.AttnLSTMEmbedding(
          test_batch_size, support_batch_size, max_depth))
    elif model in ['res']:
      max_depth = hyper_parameters['max_depth']
      support_graph.join(dc.nn.ResiLSTMEmbedding(
          test_batch_size, support_batch_size, max_depth))
      
    with tf.Session() as sess:
      model_low_data = dc.models.SupportGraphClassifier(
          sess, support_graph, test_batch_size=test_batch_size,
          support_batch_size=support_batch_size, learning_rate=learning_rate,
          verbosity="high")
        
      print('-------------------------------------')
      print('Start fitting by graph convolution')
      # Fit trained model
      model_low_data.fit(train_dataset, nb_epochs=nb_epochs,
            n_episodes_per_epoch=n_train_trials,
            n_pos=n_pos, n_neg=n_neg,
            log_every_n_samples=50)
      # Evaluating graph convolution model
      scores[model] = model_low_data.evaluate(valid_dataset, 
                          classification_metric, n_pos, n_neg, 
                          n_trials=n_eval_trials)
      
  return scores

    
if __name__ == '__main__':
  # Global variables
  np.random.seed(123)
  verbosity = 'high'
  
  #Working folder initialization
  base_dir_o="/tmp/benchmark_test_"+time.strftime("%Y_%m_%d", time.localtime())
  if os.path.exists(base_dir_o):
    shutil.rmtree(base_dir_o)
  os.makedirs(base_dir_o)
  
  parser = argparse.ArgumentParser(description='Deepchem benchmark: '+
      'giving performances of different learning models on datasets')
  parser.add_argument('-s', action='append', dest='splitter_args', default=[],
      help='Choice of splitting function: task')
  parser.add_argument('-m', action='append', dest='model_args', default=[], 
      help='Choice of model: siamese, attn, res')
  parser.add_argument('-d', action='append', dest='dataset_args', default=[], 
      help='Choice of dataset: tox21, sider, muv')
  args = parser.parse_args()
  #Datasets and models used in the benchmark test
  splitters = args.splitter_args
  models = args.model_args
  datasets = args.dataset_args

  if len(splitters) == 0:
    splitters = ['task']
  if len(models) == 0:
    models = ['siamese', 'attn', 'res']
  if len(datasets) == 0:
    datasets = ['tox21', 'sider', 'muv']

  #input hyperparameters
  #tf: dropouts, learning rate, layer_sizes, weight initial stddev,penalty,
  #    batch_size
  hps = {}
  hps = {}
  hps['siamese'] = [{'K': 4, 'n_feat': 71, 'n_pos': 1, 'n_neg': 1,
                     'test_batch_size': 128, 'n_filters': [64, 128, 64],
                     'n_fully_connected_nodes': [128], 'max_depth': 3,
                     'nb_epochs': 1, 'n_train_trials': 2000, 
                     'n_eval_trials': 20, 'learning_rate': 1e-4}]
  hps['res'] = hps['siamese']
  hps['attn'] = hps['siamese']

  for split in splitters:
    for dataset in datasets:
      for model in models:
        low_data_benchmark_loading_datasets(hps, dataset=dataset, 
                                            model=model, split=split, 
                                            verbosity='high', out_path='.')

examples/best_results.csv

deleted100644 → 0
+0 −18
Original line number Diff line number Diff line
tox21,logreg,train,0.910,valid,0.759,learning_rate,0.004~0.008,penalty(l1 or l2),0.3~0.6
sider,logreg,train,0.900,valid,0.620,learning_rate,0.004~0.008,penalty(l2),0.4~1
muv,logreg,train,0.910,valid,0.744,learning_rate,0.002~0.004,penalty(l2),0.4~0.7
toxcast,logreg,train,0.762,valid,0.622,learning_rate,0.004~0.1,penalty(l2),0.2~0.6
pcba,logreg,train,0.794,valid,0.762,learning_rate,0.004,penalty(l2),0.3

sider,tf,train,0.931,valid,0.647,learning_rate,0.0003~0.003,dropouts,0.2~0.3,layer_sizes,1000~1200
tox21,tf,train,0.996,valid,0.763,learning_rate,0.001~0.003,dropouts,0.4~0.5,layer_sizes,1000~1200
toxcast,tf,train,0.926,valid,0.705,learning_rate,0.0008~0.0012,dropouts,0.3~0.5,layer_sizes,1200
muv,tf,train,0.980,valid,0.710,learning_rate,0.0008~0.0013,dropouts,0.25~0.4,layer_sizes,1200
pcba,tf,train,0.970,valid,0.779,learning_rate,0.0008,dropouts,0.35,layer_sizes,1200

pcba,graphconv,train,0.866,valid,0.836,learning_rate,0.0001,learning_rate_decay,1500,layer_structure,2layers&64filters
muv,graphconv,train,0.881,valid,0.832,learning_rate,0.0004~0.0008,learning_rate_decay,1500,layer_structure,2layers&64~128filters
tox21,graphconv,train,0.930,valid,0.819,learning_rate,0.0003~0.002,learning_rate_decay,1000~2000,layer_structure,2layers&128filters
sider,graphconv,train,0.845,valid,0.646,learning_rate,0.0004~0.002,,,layer_structure,2layers&64filters
toxcast,graphconv,train,0.906,valid,0.725,learning_rate,0.0004,,,layer_structure,2layers&96filters
+65 −0
Original line number Diff line number Diff line
,,,,,,,,,,,
0,tox21,index,classification,train,tf,0.8560179742,valid,tf,0.7629072715,time_for_running,54.7687239647
0,tox21,index,classification,train,tf_robust,0.8574218666,valid,tf_robust,0.7666630477,time_for_running,94.5911459923
0,tox21,index,classification,train,logreg,0.9030599129,valid,logreg,0.705362221,time_for_running,59.253813982
0,tox21,index,classification,train,graphconv,0.8715826349,valid,graphconv,0.7980162563,time_for_running,165.248687983
0,sider,index,classification,train,tf,0.7752189616,valid,tf,0.6337334294,time_for_running,74.6054239273
0,sider,index,classification,train,tf_robust,0.8034814576,valid,tf_robust,0.631875281,time_for_running,134.711124897
0,sider,index,classification,train,logreg,0.9330257816,valid,logreg,0.620001812,time_for_running,83.4092998505
0,sider,index,classification,train,graphconv,0.7075047065,valid,graphconv,0.5939296433,time_for_running,51.1861209869
0,muv,index,classification,train,tf,0.9038733856,valid,tf,0.7644975895,time_for_running,361.176653862
0,muv,index,classification,train,tf_robust,0.9335517868,valid,tf_robust,0.7806120283,time_for_running,537.99600482
0,muv,index,classification,train,logreg,0.9631884897,valid,logreg,0.7660728782,time_for_running,430.140232086
0,muv,index,classification,train,graphconv,0.8399166531,valid,graphconv,0.8227625489,time_for_running,1931.37933302
0,toxcast,index,classification,train,tf,0.8301354685,valid,tf,0.6782261344,time_for_running,2157.31364393
0,toxcast,index,classification,train,tf_robust,0.8251617547,valid,tf_robust,0.6797981747,time_for_running,4165.26211691
0,toxcast,index,classification,train,logreg,0.7212350411,valid,logreg,0.5746445348,time_for_running,2639.25799894
0,toxcast,index,classification,train,graphconv,0.8211459131,valid,graphconv,0.7197369862,time_for_running,911.413060904
0,pcba,index,classification,train,tf,0.8258423231,valid,tf,0.8017449997,time_for_running,8316.01786399
0,pcba,index,classification,train,tf_robust,0.8094070952,valid,tf_robust,0.7826523072,time_for_running,18856.9995341
0,pcba,index,classification,train,logreg,0.8087514041,valid,logreg,0.7757913202,time_for_running,11075.2340579
0,pcba,index,classification,train,graphconv,0.87647472,valid,graphconv,0.8523348204,time_for_running,14497.7339029
0,delaney,index,regression,train,tf_regression,0.7830983671,valid,tf_regression,0.5789729655,time_for_running,41.1367759705
0,kaggle,None,regression,train,tf_regression,0.7480423542,valid,tf_regression,0.4516795145,time_for_running,3238.91535401
0,tox21,random,classification,train,tf,0.8565178786,valid,tf,0.7834036936,time_for_running,53.8197240829
0,tox21,random,classification,train,tf_robust,0.8549658589,valid,tf_robust,0.7735497329,time_for_running,88.9351768494
0,tox21,random,classification,train,logreg,0.9028113641,valid,logreg,0.7350604574,time_for_running,60.2267189026
0,tox21,random,classification,train,graphconv,0.8649231702,valid,graphconv,0.8268737631,time_for_running,159.461936951
0,sider,random,classification,train,tf,0.7786895104,valid,tf,0.665646893,time_for_running,75.2554209232
0,sider,random,classification,train,tf_robust,0.7607831717,valid,tf_robust,0.620646631,time_for_running,145.425393105
0,sider,random,classification,train,logreg,0.9315624982,valid,logreg,0.628537773,time_for_running,83.4710030556
0,sider,random,classification,train,graphconv,0.7059736283,valid,graphconv,0.6376782717,time_for_running,52.9893791676
0,muv,random,classification,train,tf,0.8953200915,valid,tf,0.7396286547,time_for_running,376.079932213
0,muv,random,classification,train,tf_robust,0.9142442363,valid,tf_robust,0.6672445505,time_for_running,564.25266695
0,muv,random,classification,train,logreg,0.9608985781,valid,logreg,0.6956532843,time_for_running,446.927948952
0,muv,random,classification,train,graphconv,0.8460695576,valid,graphconv,0.7755425716,time_for_running,1871.64626002
0,toxcast,random,classification,train,tf,0.8306986464,valid,tf,0.6840850646,time_for_running,2287.80405498
0,toxcast,random,classification,train,tf_robust,0.8140336561,valid,tf_robust,0.6921436476,time_for_running,4112.52007484
0,toxcast,random,classification,train,logreg,0.7373983296,valid,logreg,0.5433428006,time_for_running,2637.63535094
0,toxcast,random,classification,train,graphconv,0.82008075,valid,graphconv,0.6925051431,time_for_running,894.896769047
0,pcba,random,classification,train,tf,0.810876083,valid,tf,0.787010593,time_for_running,8910.69580603
0,pcba,random,classification,train,tf_robust,0.8092694566,valid,tf_robust,0.7776478785,time_for_running,14205.0440209
0,pcba,random,classification,train,logreg,0.8065139555,valid,logreg,0.7724261671,time_for_running,10074.3197708
0,pcba,random,classification,train,graphconv,0.8750284011,valid,graphconv,0.8443492695,time_for_running,14665.9515259
0,delaney,random,regression,train,tf_regression,0.7791066217,valid,tf_regression,0.6164873014,time_for_running,35.6433098316
0,tox21,scaffold,classification,train,tf,0.8626085326,valid,tf,0.7030201614,time_for_running,63.5685660839
0,tox21,scaffold,classification,train,tf_robust,0.8608722489,valid,tf_robust,0.7100530015,time_for_running,101.614424944
0,tox21,scaffold,classification,train,logreg,0.9004137009,valid,logreg,0.650190286,time_for_running,60.018599987
0,tox21,scaffold,classification,train,graphconv,0.8848841104,valid,graphconv,0.7317094773,time_for_running,158.603221893
0,sider,scaffold,classification,train,tf,0.7758662546,valid,tf,0.5565867138,time_for_running,74.5987181664
0,sider,scaffold,classification,train,tf_robust,0.7970038738,valid,tf_robust,0.5598527985,time_for_running,138.48903203
0,sider,scaffold,classification,train,logreg,0.925883573,valid,logreg,0.5917837968,time_for_running,85.1897189617
0,sider,scaffold,classification,train,graphconv,0.7216965584,valid,graphconv,0.5834509597,time_for_running,51.2979419231
0,muv,scaffold,classification,train,tf,0.8985808597,valid,tf,0.7622587057,time_for_running,540.442322016
0,muv,scaffold,classification,train,tf_robust,0.9442240423,valid,tf_robust,0.7257131907,time_for_running,731.773438931
0,muv,scaffold,classification,train,logreg,0.9472815201,valid,logreg,0.7672245889,time_for_running,444.053046942
0,muv,scaffold,classification,train,graphconv,0.8720846286,valid,graphconv,0.7948635787,time_for_running,1857.96855307
0,toxcast,scaffold,classification,train,tf,0.8282567273,valid,tf,0.6168555208,time_for_running,2748.30164099
0,toxcast,scaffold,classification,train,tf_robust,0.8298692148,valid,tf_robust,0.6143518759,time_for_running,2758.56494784
0,toxcast,scaffold,classification,train,logreg,0.7160825956,valid,logreg,0.4917417082,time_for_running,2630.97518611
0,toxcast,scaffold,classification,train,graphconv,0.8319269009,valid,graphconv,0.6384323545,time_for_running,888.159852028
0,pcba,scaffold,classification,train,tf,0.8142493334,valid,tf,0.7599029916,time_for_running,15337.6803091
0,pcba,scaffold,classification,train,tf_robust,0.8177093675,valid,tf_robust,0.7559237683,time_for_running,22051.4083931
0,pcba,scaffold,classification,train,logreg,0.8099796593,valid,logreg,0.7423270057,time_for_running,9959.83747697
0,pcba,scaffold,classification,train,graphconv,0.8743221913,valid,graphconv,0.8166550236,time_for_running,14184.1512611
0,delaney,scaffold,regression,train,tf_regression,0.7893516465,valid,tf_regression,0.4218847009,time_for_running,35.2720739841