Commit 03a118b7 authored by evanfeinberg's avatar evanfeinberg
Browse files

fixes to multitask

parent cf6b3cd6
Loading
Loading
Loading
Loading
+1 −1
Original line number Diff line number Diff line
@@ -162,7 +162,7 @@ class Model(object):
      for j in range(len(interval_points)-1):
        indices = range(interval_points[j], interval_points[j+1])
        y_pred_on_batch = self.predict_on_batch(X[indices, :])
        y_pred_on_batch = np.reshape(y_pred_on_batch, (len(indices),))
        #y_pred_on_batch = np.reshape(y_pred_on_batch, (len(indices),))
        y_preds.append(y_pred_on_batch)

      y_pred = np.concatenate(y_preds)
+6 −3
Original line number Diff line number Diff line
@@ -280,6 +280,10 @@ def parse_args(input_args=None):
  add_model_command(subparsers)
  return parser.parse_args(input_args)

def shard_inputs(input_file):
  input_file_no_ext = os.path.splitext(input_file)


def featurize_inputs(feature_dir, data_dir, input_files,
                     user_specified_features, tasks, smiles_field,
                     split_field, id_field, threshold, protein_pdb_field, 
@@ -321,9 +325,8 @@ def featurize_input(input_file, feature_dir, user_specified_features, tasks,
                              ligand_mol2_field=ligand_mol2_field,
                              user_specified_features=user_specified_features,
                              verbose=True)
  out = os.path.join(
      feature_dir, "%s.joblib" %(os.path.splitext(os.path.basename(input_file))[0]))
  featurizer.featurize(input_file, FeaturizedSamples.feature_types, out)

  featurizer.featurize(input_file, FeaturizedSamples.feature_types, feature_dir)

def train_test_split(input_transforms, output_transforms,
                     feature_types, splittype, data_dir):
+3 −1
Original line number Diff line number Diff line
@@ -11,6 +11,7 @@ from sklearn.externals import joblib
import gzip
import cPickle as pickle
import pandas as pd
import numpy as np

def save_to_disk(dataset, filename):
  """Save a dataset to file."""
@@ -40,4 +41,5 @@ def load_pandas_from_disk(filename):
  else:
    # First line of user-specified CSV *must* be header.
    df = pd.read_csv(filename, header=0)
    df = df.replace(np.nan,str(""), regex=True)
    return df