Commit b024b90e authored by miaecle's avatar miaecle
Browse files

variable sizes of training set

parent 5ba7e79d
Loading
Loading
Loading
Loading
+142 −0
Original line number Diff line number Diff line
#!/usr/bin/env python2
# -*- coding: utf-8 -*-
"""
Created on Mon Aug 14 16:59:49 2017

@author: zqwu
"""

from __future__ import print_function
from __future__ import division
from __future__ import unicode_literals

import os
import numpy as np
import deepchem as dc
import argparse
import pickle
import csv
from deepchem.molnet.run_benchmark_models import benchmark_classification, benchmark_regression
from deepchem.molnet.check_availability import CheckFeaturizer, CheckSplit
from deepchem.molnet.preset_hyper_parameters import hps

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: index, random, scaffold, stratified')
parser.add_argument(
    '-m',
    action='append',
    dest='model_args',
    default=[],
    help='Choice of model: tf, tf_robust, logreg, rf, irv, graphconv, xgb,' + \
         ' dag, weave, tf_regression, tf_regression_ft, rf_regression, ' + \
         'graphconvreg, xgb_regression, dtnn, dag_regression, weave_regression')
parser.add_argument(
    '-d',
    action='append',
    dest='dataset_args',
    default=[],
    help='Choice of dataset: bace_c, bace_r, bbbp, chembl, clearance, ' +
    'clintox, delaney, hiv, hopv, kaggle, lipo, muv, nci, pcba, ' +
    'pdbbind, ppb, qm7, qm7b, qm8, qm9, sampl, sider, tox21, toxcast')
parser.add_argument(
    '-t',
    action='store_true',
    dest='test',
    default=False,
    help='Evalute performance on test set')
parser.add_argument(
    '--seed',
    action='append',
    dest='seed_args',
    default=[],
    help='Choice of random seed')
args = parser.parse_args()
#Datasets and models used in the benchmark test
splitters = args.splitter_args
models = args.model_args
datasets = args.dataset_args
test = args.test
if len(args.seed_args) > 0:
  seed = int(args.seed_args[0])
else:
  seed = 123

metrics = {
    'qm7': [dc.metrics.Metric(dc.metrics.mean_absolute_error, np.mean, mode='regression')],
    'sampl': [dc.metrics.Metric(dc.metrics.rms_score, np.mean, mode='regression')],
    'pdbbind': [dc.metrics.Metric(dc.metrics.rms_score, np.mean, mode='regression')],
    'tox21': [dc.metrics.Metric(dc.metrics.roc_auc_score, np.mean, mode='classification')],
    'bace_c': [dc.metrics.Metric(dc.metrics.roc_auc_score, np.mean, mode='classification')],
    }
out_path = '.'
for dataset in datasets:
  for split in splitters:
    for model in models:
      with open(os.path.join(out_path, dataset + model + '.pkl'), 'r') as f:
        hyper_parameters = pickle.load(f)

      metric = metrics[dataset]
      if dataset in ['bace_c', 'tox21']:
        mode = 'classification'
      elif dataset in ['pdbbind', 'qm7', 'sampl']:
        mode = 'regression'


      pair = (dataset, model)
      if pair in CheckFeaturizer:
        featurizer = CheckFeaturizer[pair][0]
        n_features = CheckFeaturizer[pair][1]

      loading_functions = {
        'bace_c': dc.molnet.load_bace_classification,
        'pdbbind': dc.molnet.load_pdbbind_grid,
        'qm7': dc.molnet.load_qm7_from_mat,
        'sampl': dc.molnet.load_sampl,
        'tox21': dc.molnet.load_tox21
      }

      tasks, all_dataset, transformers = loading_functions[dataset](
          featurizer=featurizer, reload=reload)
      all_dataset = dc.data.DiskDataset.merge(all_dataset)

      for frac_train in [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9]:
        splitters = {
            'index': dc.splits.IndexSplitter(),
            'random': dc.splits.RandomSplitter(),
            'scaffold': dc.splits.ScaffoldSplitter()
        }
        splitter = splitters[split]
        train, valid, test = splitter.train_valid_test_split(all_dataset,
                                                             frac_train=frac_train,
                                                             frac_valid=1-frac_train,
                                                             frac_test=0.)
        if mode == 'classification':
          train_score, valid_score, test_score = benchmark_classification(
              train, valid, test, tasks, transformers, n_features, metric,
              model, test=False, hyper_parameters=hyper_parameters, seed=seed)
        elif mode == 'regression':
          train_score, valid_score, test_score = benchmark_regression(
              train, valid, test, tasks, transformers, n_features, metric,
              model, test=False, hyper_parameters=hyper_parameters, seed=seed)
        with open(os.path.join(out_path, 'results_variable.csv'), 'a') as f:
          writer = csv.writer(f)
          model_name = list(train_score.keys())[0]
          for i in train_score[model_name]:
            output_line = [
                dataset, str(split), mode, model_name, i, 'train',
                train_score[model_name][i], 'valid', valid_score[model_name][i]
            ]
            output_line.extend([
                'frac_train', frac_train])
            writer.writerow(output_line)