Commit 4e012017 authored by miaecle's avatar miaecle
Browse files

MPNN structures

parent 529c8134
Loading
Loading
Loading
Loading
+112 −0
Original line number Diff line number Diff line
@@ -799,3 +799,115 @@ class DAGGather(Layer):
      outputs = tf.nn.xw_plus_b(outputs, W, b_list[idw])
      outputs = self.activation(outputs)
    return outputs

class MessagePassing(Layer):
  """ General class for MPNN """

  def __init__(self,
               T,
               message_fn='enn',
               update_fn='gru',
               n_hidden=100,
               **kwargs):
    """
        Parameters
        ----------
        T: int
          Number of message passing steps
        message_fn: str, optional
          message function in the model
        update_fn: str, optional
          update function in the model
        n_hidden: int, optional
          number of hidden units in the passing phase
        """
    
    self.T = T
    self.message_fn = message_fn
    self.update_fn = update_fn
    self.n_hidden = n_hidden
    super(MessagePassing, self).__init__(**kwargs)

  def build(self, pair_features, n_pair_features):
    if self.message_fn == 'enn':
      self.message_function = EdgeNetwork(pair_features, 
                                          n_pair_features,
                                          self.n_hidden)
    if self.update_fn == 'gru':
      self.update_function = GatedRecurrentUnit(self.n_hidden)
  
  def create_tensor(self, in_layers=None, set_tensors=True, **kwargs):
    """ Perform T steps of message passing """
    if in_layers is None:
      in_layers = self.in_layers
    in_layers = convert_to_layers(in_layers)

    # Extract atom_features
    atom_features = in_layers[0].out_tensor
    pair_features = in_layers[1].out_tensor
    atom_to_pair = in_layers[2].out_tensor
    n_atom_features = atom_features.get_shape().as_list()[-1]
    n_pair_features = pair_features.get_shape().as_list()[-1]
    # Add trainable weights
    self.build(pair_features, n_pair_features)
    
    if n_atom_features < self.n_hidden:
      pad_length = self.n_hidden - n_atom_features
      out = tf.pad(atom_features, ((0,0), (0, pad_length)), mode='CONSTANT')
    elif n_atom_features > self.n_hidden:
      raise ValueError("Too large initial feature vector")
    
    for i in range(self.T):
      message = self.message_function.forward(out, atom_to_pair)
      out = self.update_function.forward(out, message)
    
    out_tensor = out
    
    if set_tensors:
      self.variables = self.trainable_weights
      self.out_tensor = out_tensor
    return out_tensor

class EdgeNetwork(object):
  """ Submodule for Message Passing """
  def __init__(self, 
               pair_features,
               n_pair_features=8, 
               n_hidden=100, 
               init='glorot_uniform'):
    self.n_pair_features = n_pair_features
    self.n_hidden = n_hidden
    self.init = initializations.get(init)
    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)
    
  
  def forward(self, atom_features, atom_to_pair):
    return tf.gather(atom_features, atom_to_pair[:,1]) * self.A

class GatedRecurrentUnit(object):
  """ Submodule for Message Passing """
  def __init__(self, n_hidden=100, init='glorot_uniform'):
    self.n_hidden = n_hidden
    self.init = initializations.get(init)
    self.Wz = self.init([n_hidden, n_hidden])
    self.Wr = self.init([n_hidden, n_hidden])
    self.Wh = self.init([n_hidden, n_hidden])
    self.Uz = self.init([n_hidden, n_hidden])
    self.Ur = self.init([n_hidden, n_hidden])
    self.Uh = self.init([n_hidden, n_hidden])
    self.bz = model_ops.zeros(shape=(n_hidden,))
    self.br = model_ops.zeros(shape=(n_hidden,))
    self.bh = model_ops.zeros(shape=(n_hidden,))
    
  def forward(self, inputs, messages):
    z = tf.nn.sigmoid(tf.matmul(messages, self.Wz) + \
                      tf.matmul(inputs, self.Uz) + self.bz)
    r = tf.nn.sigmoid(tf.matmul(messages, self.Wr) + \
                      tf.matmul(inputs, self.Ur) + self.br)
    h = (1-z) * tf.nn.tanh(tf.matmul(messages, self.Wh) + \
                           tf.matmul(inputs * r, self.Uh) + self.bh) + \
         z * inputs
    return h
    
 No newline at end of file
+142 −0
Original line number Diff line number Diff line
@@ -632,3 +632,145 @@ class GraphConvTensorGraph(TensorGraph):
    y_ = self.predict_on_generator(generator, transformers)

    return y_.reshape(-1, n_tasks)[:n_smiles]

class MPNNTensorGraph(TensorGraph):

  def __init__(self,
               n_tasks,
               n_atom_feat=70,
               n_pair_feat=8,
               n_hidden=100,
               T=5,
               **kwargs):
    """
        Parameters
        ----------
        n_tasks: int
          Number of tasks
        n_atom_feat: int, optional
          Number of features per atom.
        n_pair_feat: int, optional
          Number of features per pair of atoms.
        n_hidden: int, optional
          Number of units(convolution depths) in corresponding hidden layer
        n_graph_feat: int, optional
          Number of output features for each molecule(graph)

        """
    self.n_tasks = n_tasks
    self.n_atom_feat = n_atom_feat
    self.n_pair_feat = n_pair_feat
    self.n_hidden = n_hidden
    self.T = T
    super(MPNNTensorGraph, self).__init__(**kwargs)
    self.build_graph()

  def build_graph(self):
    """Building graph structures:
        Features => WeaveLayer => WeaveLayer => Dense => WeaveGather => Classification or Regression
        """
    self.atom_features = Feature(shape=(None, self.n_atom_feat))
    self.pair_features = Feature(shape=(None, self.n_pair_feat))
    self.pair_split = Feature(shape=(None,), dtype=tf.int32)
    self.atom_split = Feature(shape=(None,), dtype=tf.int32)
    self.atom_to_pair = Feature(shape=(None, 2), dtype=tf.int32)
    
    message_passing = MessagePassing(self.T, 
                                     message_fn='enn', 
                                     update_fn='gru',
                                     self.n_hidden,
                                     in_layers=[self.atom_features,
                                                self.pair_features,
                                                self.atom_to_pair])
    
    

    costs = []
    self.labels_fd = []
    for task in range(self.n_tasks):
      if self.mode == "classification":
        classification = Dense(
            out_channels=2, activation_fn=None, in_layers=[weave_gather])
        softmax = SoftMax(in_layers=[classification])
        self.add_output(softmax)

        label = Label(shape=(None, 2))
        self.labels_fd.append(label)
        cost = SoftMaxCrossEntropy(in_layers=[label, classification])
        costs.append(cost)
      if self.mode == "regression":
        regression = Dense(
            out_channels=1, activation_fn=None, in_layers=[weave_gather])
        self.add_output(regression)

        label = Label(shape=(None, 1))
        self.labels_fd.append(label)
        cost = L2Loss(in_layers=[label, regression])
        costs.append(cost)
    if self.mode == "classification":
      all_cost = Concat(in_layers=costs, axis=1)
    elif self.mode == "regression":
      all_cost = Stack(in_layers=costs, axis=1)
    self.weights = Weights(shape=(None, self.n_tasks))
    loss = WeightedError(in_layers=[all_cost, self.weights])
    self.set_loss(loss)

  def default_generator(self,
                        dataset,
                        epochs=1,
                        predict=False,
                        pad_batches=True):
    """ TensorGraph style implementation
        similar to deepchem.models.tf_new_models.graph_topology.AlternateWeaveTopology.batch_to_feed_dict
        """
    for epoch in range(epochs):
      if not predict:
        print('Starting epoch %i' % epoch)
      for (X_b, y_b, w_b, ids_b) in dataset.iterbatches(
          batch_size=self.batch_size,
          deterministic=True,
          pad_batches=pad_batches):

        feed_dict = dict()
        if y_b is not None and not predict:
          for index, label in enumerate(self.labels_fd):
            if self.mode == "classification":
              feed_dict[label] = to_one_hot(y_b[:, index])
            if self.mode == "regression":
              feed_dict[label] = y_b[:, index:index + 1]
        if w_b is not None and not predict:
          feed_dict[self.weights] = w_b

        atom_feat = []
        pair_feat = []
        atom_split = []
        atom_to_pair = []
        pair_split = []
        start = 0
        for im, mol in enumerate(X_b):
          n_atoms = mol.get_num_atoms()
          # number of atoms in each molecule
          atom_split.extend([im] * n_atoms)
          # index of pair features
          C0, C1 = np.meshgrid(np.arange(n_atoms), np.arange(n_atoms))
          atom_to_pair.append(
              np.transpose(
                  np.array([C1.flatten() + start,
                            C0.flatten() + start])))
          # number of pairs for each atom
          pair_split.extend(C1.flatten() + start)
          start = start + n_atoms

          # atom features
          atom_feat.append(mol.get_atom_features())
          # pair features
          pair_feat.append(
              np.reshape(mol.get_pair_features(), (n_atoms * n_atoms,
                                                   self.n_pair_feat)))

        feed_dict[self.atom_features] = np.concatenate(atom_feat, axis=0)
        feed_dict[self.pair_features] = np.concatenate(pair_feat, axis=0)
        feed_dict[self.pair_split] = np.array(pair_split)
        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
 No newline at end of file