Commit 3eb3be38 authored by peastman's avatar peastman
Browse files

Created example for MAML

parent 8edb82dd
Loading
Loading
Loading
Loading
+27 −5
Original line number Diff line number Diff line
"""Model-Agnostic Meta-Learning (MAML) algorithm for low data learning."""

from deepchem.models.tensorgraph.layers import Layer
from deepchem.models.tensorgraph.optimizers import Adam
from deepchem.models.tensorgraph.optimizers import Adam, GradientDescent
import os
import shutil
import tempfile
@@ -68,6 +68,8 @@ class MAML(object):

  To use this class, create a subclass of MetaLearner that encapsulates the model
  and data for your learning problem.  Pass it to a MAML object and call fit().
  You can then use train_on_current_task() to fine tune the model for a particular
  task.
  """

  def __init__(self,
@@ -151,11 +153,14 @@ class MAML(object):
            self._loss, replacements)
      self._meta_loss = updated_loss

      # Create the optimizer for meta-optimization.
      # Create the optimizers for meta-optimization and task optimization.

      self._global_step = tf.placeholder(tf.int32, [])
      self._tf_optimizer = optimizer._create_optimizer(self._global_step)
      self._train_op = self._tf_optimizer.minimize(self._meta_loss)
      self._meta_train_op = optimizer._create_optimizer(
          self._global_step).minimize(self._meta_loss)
      task_optimizer = GradientDescent(learning_rate=self._learning_rate)
      self._task_train_op = task_optimizer._create_optimizer(
          self._global_step).minimize(self._loss)
      self._session = tf.Session()

  def __del__(self):
@@ -199,7 +204,7 @@ class MAML(object):
        feed_dict[self._global_step] = i
        for key, value in self.learner.get_batch().items():
          feed_dict[self._meta_placeholders[key]] = value
        self._session.run(self._train_op, feed_dict=feed_dict)
        self._session.run(self._meta_train_op, feed_dict=feed_dict)

        # Do checkpointing.

@@ -218,3 +223,20 @@ class MAML(object):
    with self._graph.as_default():
      saver = tf.train.Saver(self.learner.variables)
      saver.restore(self._session, last_checkpoint)

  def train_on_current_task(self, optimization_steps=1, restore=True):
    """Perform a few steps of gradient descent to fine tune the model on the current task.

    Parameters
    ----------
    optimization_steps: int
      the number of steps of gradient descent to perform
    restore: bool
      if True, restore the model from the most recent checkpoint before optimizing
    """
    if restore:
      self.restore()
    with self._graph.as_default():
      feed_dict = self.learner.get_batch()
      for i in range(optimization_steps):
        self._session.run(self._task_train_op, feed_dict=feed_dict)
+84 −0
Original line number Diff line number Diff line
from __future__ import print_function

import deepchem as dc
import numpy as np
import random

# Load the data.

tasks, datasets, transformers = dc.molnet.load_toxcast()
(train_dataset, valid_dataset, test_dataset) = datasets
x = train_dataset.X
y = train_dataset.y
w = train_dataset.w
n_features = x.shape[1]
n_molecules = y.shape[0]
n_tasks = y.shape[1]

# For each task, create a list of the molecules for which we have data.

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.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.set_task_index(0)

  @property
  def loss(self):
    return model.loss

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

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

  def get_batch(self):
    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
    weights = w[mols, self.task].reshape((-1, 1))
    feed_dict = {}
    feed_dict[model.features[0].out_tensor] = x[mols, :]
    feed_dict[model.labels[0].out_tensor] = labels
    feed_dict[model.task_weights[0].out_tensor] = weights
    self.batch_start += self.batch_size
    return feed_dict

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

n_epochs = 30
learner = ToxcastLearner()
maml = dc.metalearning.MAML(learner)
maml.fit(n_epochs*learner.n_training_tasks)

# Validate on the remaining tasks.

def compute_loss(steps):
  maml.restore()
  losses = []
  for task in range(learner.n_training_tasks, len(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)

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