Commit 52631518 authored by Bharath Ramsundar's avatar Bharath Ramsundar
Browse files

required changes to load CANVAS features

parent b7da7209
Loading
Loading
Loading
Loading
+30 −16
Original line number Diff line number Diff line
@@ -3,6 +3,12 @@ Top level script to featurize input, train models, and evaluate them.
"""
import argparse
import numpy as np
from deep_chem.utils.featurize import generate_directories
from deep_chem.utils.featurize import extract_data
from deep_chem.utils.featurize import generate_targets
from deep_chem.utils.featurize import generate_features
from deep_chem.utils.featurize import generate_fingerprints
from deep_chem.utils.featurize import generate_descriptors
from deep_chem.models.deep import fit_singletask_mlp
from deep_chem.models.deep import fit_multitask_mlp
from deep_chem.models.deep3d import fit_3D_convolution
@@ -16,8 +22,7 @@ def parse_args(input_args=None):
  parser = argparse.ArgumentParser()
  subparsers = parser.add_subparsers(title='Modes')
 
  # TODO(rbharath): This function should invoke process_dataset under the hood
  # to avoid having to invoke two scripts.
  # FEATURIZE FLAGS
  featurize_cmd = subparsers.add_parser("featurize",
                      help="Featurize raw input data.")
  featurize_cmd.add_argument("--input-file", required=1,
@@ -27,29 +32,30 @@ def parse_args(input_args=None):
                      help="Type of input file. If pandas, input must be a pkl.gz\n"
                           "containing a pandas dataframe. If sdf, should be in\n"
                           "(perhaps gzipped) sdf file.")
  featurize_cmd.add_argument("--delimiter", default="\t",
  featurize_cmd.add_argument("--delimiter", default=",",
                      help="If csv input, delimiter to use for read csv file")
  featurize_cmd.add_argument("--fields", required=1, nargs="+",
                      help = "Names of fields.")
  featurize_cmd.add_argument("--field-types", required=1, nargs="+",
                      choices=["string", "float", "list-string", "list-float", "ndarray"],
                      help="Type of data in fields.")
  featurize_cmd.add_argument("--feature-endpoint", type=str,
  featurize_cmd.add_argument("--feature-endpoints", type=str, nargs="+",
                      help="Optional endpoint that holds pre-computed feature vector")
  featurize_cmd.add_argument("--prediction-endpoint", type=str, required=1,
                      help="Name of measured endpoint to predict.")
  featurize_cmd.add_argument("--split-endpoint", type=str, default=None,
                      help="Name of endpoint specifying train/test split.")
  featurize_cmd.add_argument("--smiles-endpoint", type=str, default="smiles",
                      help="Name of endpoint specifying SMILES for molecule.")
  featurize_cmd.add_argument("--threshold", type=float, default=None,
                      help="If specified, will be used to binarize real-valued prediction-endpoint.")
  featurize_cmd.add_argument("--has-colnames", type=bool, default=False,
                      help="Input has column labels which should be skipped (only for csv/xlsx).")
  featurize_cmd.add_argument("--name", required=1,
                      help="Name of the dataset.")
  featurize_cmd.add_argument("--out", required=1,
                      help="Folder to generate processed dataset in.")
  featurize_cmd.set_defaults(func=featurize_input)

  # TRAIN FLAGS
  train_cmd = subparsers.add_parser("train",
                  help="Train a model on specified data.")
  group = train_cmd.add_argument_group("load-and-transform")
@@ -104,6 +110,8 @@ def parse_args(input_args=None):

  eval_cmd = subparsers.add_parser("eval",
                help="Evaluate trained model on specified data.")
  eval_cmd.add_argument("--paths", nargs="+", required=1,
                      help="Paths to input datasets.")
  eval_cmd.add_argument("--splittype", type=str, default="scaffold",
                       choices=["scaffold", "random", "specified"],
                       help="Type of train/test data-splitting.\n"
@@ -125,21 +133,19 @@ def featurize_input(args):
  if len(args.fields) != len(args.field_types):
    raise ValueError("number of fields does not equal number of field types")
  out_x_pkl, out_y_pkl, out_sdf = generate_directories(args.name, args.out, 
      args.feature_endpoint)
      args.feature_endpoints)
  df, mols = extract_data(args.input_file, args.input_type, args.fields,
      args.field_types, args.prediction_endpoint,
      args.threshold, args.delimiter, args.has_colnames)
  generate_targets(df, mols, args.prediction_endpoint, args.split_endpoint, out_y_pkl, out_sdf)
  generate_features(df, args.feature_endpoint, out_x_pkl)
      args.field_types, args.prediction_endpoint, args.smiles_endpoint,
      args.threshold, args.delimiter)
  generate_targets(df, mols, args.prediction_endpoint, args.split_endpoint,
      args.smiles_endpoint, out_y_pkl, out_sdf)
  generate_features(df, args.feature_endpoints, args.smiles_endpoint, out_x_pkl)
  generate_fingerprints(args.name, args.out)
  generate_descriptors(args.name, args.out)

def main():
  args = parse_args()
  paths = {}

def train_model(args):
  """Builds model from featurized data."""
  paths = args.paths

  targets = get_target_names(paths)
  task_types = {target: args.task_type for target in targets}
  input_transforms = args.input_transforms 
@@ -173,11 +179,19 @@ def main():
        batch_size=args.batch_size)
  else:
    models = fit_singletask_models(per_task_data, args.model, task_types)
  # TODO(rbharath): Save trained model.

def eval_trained_model(args):
  results, aucs, r2s, rms = compute_model_performance(per_task_data, models,
    args.compute_aucs, args.compute_r2s, args.compute_rms) 
  if args.csv_out is not None:
    results_to_csv(results, args.csv_out, task_type=args.task_type)

def main():
  args = parse_args()
  args.func(args)



if __name__ == "__main__":
  main()
+1 −1
Original line number Diff line number Diff line
# Usage ./process_bace.sh INPUT_SDF_FILE OUT_DIR DATASET_NAME
python -m deep_chem.scripts.process_dataset --input-file $1 --input-type sdf --fields Name smiles pIC50 Model --field-types string string float string --name $3 --out $2 --prediction-endpoint pIC50
python -m deep_chem.scripts.modeler featurize --input-file $1 --input-type sdf --fields Name smiles pIC50 Model --field-types string string float string --name $3 --out $2 --prediction-endpoint pIC50
+70 −49
Original line number Diff line number Diff line
@@ -4,6 +4,8 @@ Process an input dataset into a format suitable for machine learning.
import os
import cPickle as pickle
import gzip
import functools
import itertools
import pandas as pd
import openpyxl as px
import numpy as np
@@ -13,13 +15,8 @@ from rdkit import Chem
import subprocess
from vs_utils.utils import SmilesGenerator, ScaffoldGenerator

def parse_args(input_args=None):
  """Parse command-line arguments."""
  parser = argparse.ArgumentParser()
  return parser.parse_args(input_args)

def generate_directories(name, out, feature_endpoint):
  """Generate processed dataset."""
def generate_directories(name, out, feature_endpoints):
  """Generate directory structure for featurized dataset."""
  dataset_dir = os.path.join(out, name)
  if not os.path.exists(dataset_dir):
    os.makedirs(dataset_dir)
@@ -35,15 +32,16 @@ def generate_directories(name, out, feature_endpoint):
  shards_dir = os.path.join(dataset_dir, "shards")
  if not os.path.exists(shards_dir):
    os.makedirs(shards_dir)
  if feature_endpoint is not None:
    feature_endpoint_dir = os.path.join(dataset_dir, feature_endpoint)
  if feature_endpoints is not None:
    feature_endpoint_dir = os.path.join(dataset_dir, "features")
    if not os.path.exists(feature_endpoint_dir):
      os.makedirs(feature_endpoint_dir)

  # Return names of files to be generated
  out_y_pkl = os.path.join(target_dir, "%s.pkl.gz" % name)
  out_sdf = os.path.join(shards_dir, "%s-0.sdf.gz" % name)
  out_x_pkl = os.path.join(feature_endpoint_dir, "%s.pkl.gz" %name) if feature_endpoint is not None else None
  out_x_pkl = (os.path.join(feature_endpoint_dir, "%s.pkl.gz" %name)
      if feature_endpoints is not None else None)
  return out_x_pkl, out_y_pkl, out_sdf

def parse_float_input(val):
@@ -115,32 +113,38 @@ def get_rows(input_file, input_type, delimiter):
        mols = [mol for mol in supp if mol is not None]
      return mols

def get_row_data(row, input_type, fields, field_types):
  """Extract information from row data."""
def get_colnames(row, input_type):
  """Get names of all columns."""
  if input_type == "xlsx":
    return [cell.internal_value for cell in row]
  elif input_type == "csv":
    return row

def get_row_data(row, input_type, fields, smiles_endpoint, colnames=None):
  """Extract information from row data."""
  row_data = {}
  if input_type == "xlsx":
    for ind, colname in enumerate(colnames):
      if colname in fields:
        row_data[colname] = row[ind].internal_value
  elif input_type == "csv":
    for ind, colname in enumerate(colnames):
      if colname in fields:
        row_data[colname] = row[ind]
  elif input_type == "pandas":
    # pandas rows are tuples (row_num, row_info)
    row, row_data = row[1], {}
    # pandas rows are keyed by field-name. Change to key by index to match
    # csv/xlsx handling
    for ind, field in enumerate(fields):
      row_data[ind] = row[field]
    return row_data
    # pandas rows are tuples (row_num, row_data)
    row = row[1]
    for field in fields:
      row_data[field] = row[field]
  elif input_type == "sdf":
    row_data, mol = {}, row
    for ind, (field, field_type) in enumerate(zip(fields, field_types)):
      # TODO(rbharath): SDF files typically don't have smiles, so we manually
      # generate smiles in this case. This is a kludgey solution...
      if field == "smiles":
        row_data[ind] = Chem.MolToSmiles(mol)
        continue
      if not mol.HasProp(field):
        row_data[ind] = None
    mol = {}
    for field in fields:
      if field == smiles_endpoint:
        row_data[field] = Chem.MolToSmiles(mol)
      elif not mol.HasProp(field):
        row_data[field] = None
      else:
        row_data[ind] = mol.GetProp(field)
        row_data[field] = mol.GetProp(field)
  return row_data

def process_field(data, field_type):
@@ -159,14 +163,14 @@ def process_field(data, field_type):
  elif field_type == "ndarray":
    return data 

def generate_targets(df, mols, prediction_endpoint, split_endpoint, out_pkl, out_sdf):
def generate_targets(df, mols, prediction_endpoint, split_endpoint, smiles_endpoint, out_pkl, out_sdf):
  """Process input data file, generate labels, i.e. y"""
  #TODO(enf, rbharath): Modify package unique identifier to take user-specified 
    #unique identifier instead of assuming smiles string
  if split_endpoint is not None:
    labels_df = df[["smiles", prediction_endpoint, split_endpoint]]
    labels_df = df[[smiles_endpoint, prediction_endpoint, split_endpoint]]
  else:
    labels_df = df[["smiles", prediction_endpoint]]
    labels_df = df[[smiles_endpoint, prediction_endpoint]]

  # Write pkl.gz file
  with gzip.open(out_pkl, "wb") as f:
@@ -178,45 +182,62 @@ def generate_targets(df, mols, prediction_endpoint, split_endpoint, out_pkl, out
      w.write(mol)
    w.close()

def generate_scaffold(smiles_elt, include_chirality=False):
  smiles_string = smiles_elt["smiles"]
def generate_scaffold(smiles_elt, include_chirality=False, smiles_endpoint="smiles"):
  smiles_string = smiles_elt[smiles_endpoint]
  mol = Chem.MolFromSmiles(smiles_string)
  engine = ScaffoldGenerator(include_chirality=include_chirality)
  scaffold = engine.get_scaffold(mol)
  return(scaffold)

def generate_features(df, feature_endpoint, out_pkl):
  if feature_endpoint is None:
def generate_features(df, feature_endpoints, smiles_endpoint, out_pkl):
  if feature_endpoints is None:
    print("No feature endpoint specified by user.")
    return

  features_df = df[["smiles"]]
  features_df["features"] = df[[feature_endpoint]]
  features_df["scaffolds"] = df[["smiles"]].apply(generate_scaffold, axis=1)
  features_df["mol_id"] = df[["smiles"]].apply(lambda s : "", axis=1)
  features_df = df[[smiles_endpoint]]
  #features_df["features"] = df[[feature_endpoint]]
  features_data = []
  for row in df.iterrows():
    # pandas rows are tuples (row_num, row_data)
    row, feature_list = row[1], []
    for feature in feature_endpoints:
      feature_list.append(row[feature])
    features_data.append({"row": np.array(feature_list)})
  features_df["features"] = pd.DataFrame(features_data)
  
    
  features_df["scaffolds"] = df[[smiles_endpoint]].apply(
    functools.partial(generate_scaffold, smiles_endpoint=smiles_endpoint),
    axis=1)
  features_df["mol_id"] = df[[smiles_endpoint]].apply(lambda s : "", axis=1)

  with gzip.open(out_pkl, "wb") as f:
    pickle.dump(features_df, f, pickle.HIGHEST_PROTOCOL)

def extract_data(input_file, input_type, fields, field_types, 
      prediction_endpoint, threshold, delimiter, has_colnames):
      prediction_endpoint, smiles_endpoint, threshold, delimiter):
  """Extracts data from input as Pandas data frame"""

  rows, mols, smiles = [], [], SmilesGenerator()
  colnames = [] 
  for row_index, raw_row in enumerate(get_rows(input_file, input_type, delimiter)):
    print row_index
    # Skip row labels if necessary.
    if has_colnames and (row_index == 0 or raw_row is None):  
    # Skip empty rows
    if raw_row is None:
      continue
    # TODO(rbharath): The script expects that all columns in xlsx/csv files
    # have column names attached. Check that this holds true somewhere
    # Get column names if xlsx/csv and continue
    if (input_type == "xlsx" or input_type == "csv") and row_index == 0:  
      colnames = get_colnames(raw_row, input_type)
      continue
    row, row_data = {}, get_row_data(raw_row, input_type, fields, field_types)
    row, row_data = {}, get_row_data(raw_row, input_type, fields, smiles_endpoint, colnames)
    for ind, (field, field_type) in enumerate(zip(fields, field_types)):
      if field == prediction_endpoint and threshold is not None:
        raw_val = process_field(row_data[ind], field_type)
        raw_val = process_field(row_data[field], field_type)
        row[field] = 1 if raw_val > threshold else 0 
      else:
        row[field] = process_field(row_data[ind], field_type)
    
    mol = Chem.MolFromSmiles(row["smiles"])
        row[field] = process_field(row_data[field], field_type)
    mol = Chem.MolFromSmiles(row[smiles_endpoint])
    row["smiles"] = smiles.get_smiles(mol)
    mols.append(mol)
    rows.append(row)