Commit 315b642d authored by peastman's avatar peastman
Browse files

Improvements to MAML example

parent 4569f305
Loading
Loading
Loading
Loading
+31 −19
Original line number Diff line number Diff line
@@ -2,7 +2,6 @@ from __future__ import print_function

import deepchem as dc
import numpy as np
import random

# Load the data.

@@ -16,24 +15,28 @@ n_molecules = y.shape[0]
n_tasks = y.shape[1]

# For each task, create a list of the molecules for which we have data.
# Since we are interested in low data learning, we will only use the first
# 20 molecules with data for each task, even though most tasks have much
# more than that.

task_molecules = []
for i in range(n_tasks):
  task_molecules.append(w[:, i].nonzero()[0])
tasks = [i for i,m in enumerate(task_molecules) if len(m) >= 100]
random.shuffle(tasks)

# Create the model to train.

model = dc.models.TensorGraphMultiTaskClassifier(1, n_features, layer_sizes=[1000], dropouts=[0.0])
model = dc.models.TensorGraphMultiTaskClassifier(
    1, n_features, layer_sizes=[1000], dropouts=[0.0])
model.build()

# Define a MetaLearner describing the learning problem.


class ToxcastLearner(dc.metalearning.MetaLearner):

  def __init__(self):
    self.n_training_tasks = int(len(tasks)*0.8)
    self.batch_size = 50
    self.n_training_tasks = int(n_tasks * 0.8)
    self.batch_size = 10
    self.set_task_index(0)

  @property
@@ -41,17 +44,18 @@ class ToxcastLearner(dc.metalearning.MetaLearner):
    return model.loss

  def set_task_index(self, index):
    self.task_index = index
    self.task = tasks[index]
    self.task = index
    self.batch_start = 0

  def select_task(self):
    self.set_task_index((self.task_index+1) % self.n_training_tasks)
    self.set_task_index((self.task + 1) % self.n_training_tasks)

  def get_batch(self):
    mols = task_molecules[self.task][self.batch_start:self.batch_start+self.batch_size]
    mols = task_molecules[self.task][self.batch_start:
                                     self.batch_start + self.batch_size]
    labels = np.zeros((self.batch_size, 1, 2))
    labels[np.arange(self.batch_size), 0, y[mols, self.task].astype(np.int64)] = 1
    labels[np.arange(self.batch_size), 0, y[mols, self.task].astype(
        np.int64)] = 1
    weights = w[mols, self.task].reshape((-1, 1))
    feed_dict = {}
    feed_dict[model.features[0].out_tensor] = x[mols, :]
@@ -60,6 +64,7 @@ class ToxcastLearner(dc.metalearning.MetaLearner):
    self.batch_start += self.batch_size
    return feed_dict


# Run meta-learning on 80% of the tasks.

n_epochs = 40
@@ -70,16 +75,23 @@ maml.fit(steps)

# Validate on the remaining tasks.


def compute_loss(steps):
  maml.restore()
  losses = []
  for task in range(learner.n_training_tasks, len(tasks)):
  y_true = []
  y_pred = []
  for task in range(learner.n_training_tasks, n_tasks):
    learner.set_task_index(task)
    if steps > 0:
      maml.train_on_current_task(optimization_steps=steps)
    with model._get_tf("Graph").as_default():
      losses.append(maml._session.run(learner.loss, feed_dict=learner.get_batch()))
  return np.average(losses)
      feed_dict = learner.get_batch()
      y_true.append(feed_dict[model.labels[0].out_tensor][0])
      y_pred.append(maml._session.run(model.outputs[0], feed_dict=feed_dict)[0])
  y_true = np.concatenate(y_true)
  y_pred = np.concatenate(y_pred)
  return dc.metrics.compute_roc_auc_scores(y_true, y_pred)


print('Loss before training:', compute_loss(0))
print('Loss after training:', compute_loss(1))
print('AUC before fine tuning:', compute_loss(0))
print('AUC after fine tuning:', compute_loss(1))