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

kerasify changes

parent 997c598e
Loading
Loading
Loading
Loading
+11 −1
Original line number Diff line number Diff line
@@ -324,6 +324,10 @@ class Dataset(object):
      shard_perm = np.arange(num_shards)
    for i in range(num_shards):
      X, y, w, ids = self.get_shard(shard_perm[i])
      ############################################################ DEBUG
      print("X.shape, y.shape, w.shape, ids.shape")
      print(X.shape, y.shape, w.shape, ids.shape)
      ############################################################ DEBUG
      n_samples = X.shape[0]
      if not deterministic:
        sample_perm = np.random.permutation(n_samples)
@@ -338,7 +342,13 @@ class Dataset(object):
      for j in range(len(interval_points)-1):
        indices = range(interval_points[j], interval_points[j+1])
        perm_indices = sample_perm[indices]
        X_batch = X[perm_indices, :]
        ############################################################# DEBUG
        #print("len(indices)")
        #print(len(indices))
        #print("perm_indices")
        #print(perm_indices)
        ############################################################# DEBUG
        X_batch = X[perm_indices]
        y_batch = y[perm_indices]
        w_batch = w[perm_indices]
        ids_batch = ids[perm_indices]
+3 −2
Original line number Diff line number Diff line
@@ -99,7 +99,7 @@ class Splitter(object):
    return train_dataset, valid_dataset, test_dataset

  def train_test_split(self, samples, train_dir, test_dir, seed=None,
                       frac_train=.8):
                       frac_train=.8, compute_feature_statistics=True):
    """
    Splits self into train/test sets.
    Returns Dataset objects.
@@ -107,7 +107,8 @@ class Splitter(object):
    valid_dir = tempfile.mkdtemp()
    train_samples, _, test_samples = self.train_valid_test_split(
      samples, train_dir, valid_dir, test_dir,
      frac_train=frac_train, frac_test=1-frac_train, frac_valid=0.)
      frac_train=frac_train, frac_test=1-frac_train, frac_valid=0.,
      compute_feature_statistics=compute_feature_statistics)
    return train_samples, test_samples

  def split(self, dataset, frac_train=None, frac_valid=None, frac_test=None,