Commit f5576f7b authored by miaecle's avatar miaecle
Browse files

MPNN first example

parent 22b74a3b
Loading
Loading
Loading
Loading
+16 −5
Original line number Diff line number Diff line
@@ -835,6 +835,8 @@ class MessagePassing(Layer):
                                          self.n_hidden)
    if self.update_fn == 'gru':
      self.update_function = GatedRecurrentUnit(self.n_hidden)
    self.trainable_weights = self.message_function.trainable_weights + \
        self.update_function.trainable_weights

  def create_tensor(self, in_layers=None, set_tensors=True, **kwargs):
    """ Perform T steps of message passing """
@@ -881,10 +883,14 @@ class EdgeNetwork(object):
    W = self.init([n_pair_features, n_hidden*n_hidden])
    b = model_ops.zeros(shape=(n_hidden*n_hidden,))
    self.A = tf.nn.xw_plus_b(pair_features, W, b)

    self.A = tf.reshape(self.A, (-1, n_hidden, n_hidden))
    self.trainable_weights = [W, b]

  def forward(self, atom_features, atom_to_pair):
    return tf.gather(atom_features, atom_to_pair[:,1]) * self.A
    out = tf.expand_dims(tf.gather(atom_features, atom_to_pair[:,1]), 2)
    out = tf.reduce_sum(out * self.A, axis=1)
    out = tf.segment_sum(out, atom_to_pair[:,0])
    return  out

class GatedRecurrentUnit(object):
  """ Submodule for Message Passing """
@@ -900,6 +906,9 @@ class GatedRecurrentUnit(object):
    self.bz = model_ops.zeros(shape=(n_hidden,))
    self.br = model_ops.zeros(shape=(n_hidden,))
    self.bh = model_ops.zeros(shape=(n_hidden,))
    self.trainable_weights = [self.Wz, self.Wr, self.Wh,
                              self.Uz, self.Ur, self.Uh,
                              self.bz, self.br, self.bh]

  def forward(self, inputs, messages):
    z = tf.nn.sigmoid(tf.matmul(messages, self.Wz) + \
@@ -939,13 +948,13 @@ class SetGather(Layer):
    self.init = initializations.get(init)
    super(SetGather, self).__init__(**kwargs)

  def build(self, pair_features, n_pair_features):
  def build(self):
    self.U = self.init((2*self.n_hidden, 4*self.n_hidden))
    self.b = tf.Variable(
        np.concatenate((np.zeros(self.n_hidden), np.ones(self.n_hidden),
                        np.zeros(self.n_hidden), np.zeros(self.n_hidden))),
        dtype=tf.float32)
    
    self.trainable_weights = [self.U, self.b]

  def create_tensor(self, in_layers=None, set_tensors=True, **kwargs):
    """ Perform T steps of message passing """
@@ -965,7 +974,9 @@ class SetGather(Layer):
      q_expanded = tf.gather(h, atom_split)
      e = tf.reduce_sum(atom_features * q_expanded, 1)
      e_mols = tf.dynamic_partition(e, atom_split, self.batch_size)
      a = tf.concat([tf.nn.softmax(e_mol) for e_mol in e_mols], 0)
      # Add another value(~-Inf) to prevent error in softmax
      e_mols = [tf.concat([e_mol, tf.constant([-1000.])], 0) for e_mol in e_mols]
      a = tf.concat([tf.nn.softmax(e_mol)[:-1] for e_mol in e_mols], 0)
      r = tf.segment_sum(tf.reshape(a, [-1,1]) * atom_features, atom_split)
      q_star = tf.concat([h, r], axis=1)
      h, c = self.LSTMStep(q_star, c)
+12 −2
Original line number Diff line number Diff line
@@ -684,7 +684,6 @@ class MPNNTensorGraph(TensorGraph):

  def __init__(self,
               n_tasks,
               batch_size,
               n_atom_feat=70,
               n_pair_feat=8,
               n_hidden=100,
@@ -707,7 +706,6 @@ class MPNNTensorGraph(TensorGraph):

        """
    self.n_tasks = n_tasks
    self.batch_size = batch_size
    self.n_atom_feat = n_atom_feat
    self.n_pair_feat = n_pair_feat
    self.n_hidden = n_hidden
@@ -826,3 +824,15 @@ class MPNNTensorGraph(TensorGraph):
        feed_dict[self.atom_split] = np.array(atom_split)
        feed_dict[self.atom_to_pair] = np.concatenate(atom_to_pair, axis=0)
        yield feed_dict

  def predict(self, dataset, transformers=[], batch_size=None):
    length_dataset = dataset.y.shape[0]
    generator = self.default_generator(dataset, predict=True, pad_batches=True)
    y_pred = self.predict_on_generator(generator, transformers)
    return y_pred[:length_dataset]

  def predict_proba(self, dataset, transformers=[], batch_size=None):
    length_dataset = dataset.y.shape[0]
    generator = self.default_generator(dataset, predict=True, pad_batches=True)
    y_pred = self.predict_proba_on_generator(generator, transformers)
    return y_pred[:length_dataset]
 No newline at end of file
+51 −0
Original line number Diff line number Diff line
"""
Script that trains MPNN models on qm8 dataset.
"""
from __future__ import print_function
from __future__ import division
from __future__ import unicode_literals

import numpy as np
np.random.seed(123)
import tensorflow as tf
tf.set_random_seed(123)
import deepchem as dc

# Load QM8 dataset
tasks, datasets, transformers = dc.molnet.load_qm8(featurizer='MP')
train_dataset, valid_dataset, test_dataset = datasets

# Fit models
metric = [
    dc.metrics.Metric(dc.metrics.mean_absolute_error, np.mean, mode="regression"),
    dc.metrics.Metric(dc.metrics.pearson_r2_score, np.mean, mode="regression")
]

# Batch size of models
batch_size = 32
n_atom_feat = 70
n_pair_feat = 8

model = dc.models.MPNNTensorGraph(
    len(tasks),
    n_atom_feat=n_atom_feat,
    n_pair_feat=n_pair_feat,
    T=5,
    M=10,
    batch_size=batch_size,
    learning_rate=0.0001,
    use_queue=False,
    mode="regression")

# Fit trained model
model.fit(train_dataset, nb_epoch=100)

print("Evaluating model")
train_scores = model.evaluate(train_dataset, metric, transformers)
valid_scores = model.evaluate(valid_dataset, metric, transformers)

print("Train scores")
print(train_scores)

print("Validation scores")
print(valid_scores)