Commit e502ffb9 authored by miaecle's avatar miaecle
Browse files

Merge remote-tracking branch 'remotes/mine/MPNN' into BP

parents f36de7c1 bf83dc66
Loading
Loading
Loading
Loading
+2 −1
Original line number Diff line number Diff line
@@ -96,6 +96,7 @@ def pad_batch(batch_size, X_b, y_b, w_b, ids_b):

    # Fill in batch arrays
    start = 0
    w_out[start:start + num_samples] = w_b[:]
    while start < batch_size:
      num_left = batch_size - start
      if num_left < num_samples:
@@ -104,9 +105,9 @@ def pad_batch(batch_size, X_b, y_b, w_b, ids_b):
        increment = num_samples
      X_out[start:start + increment] = X_b[:increment]
      y_out[start:start + increment] = y_b[:increment]
      w_out[start:start + increment] = w_b[:increment]
      ids_out[start:start + increment] = ids_b[:increment]
      start += increment

    return (X_out, y_out, w_out, ids_out)


+47 −13
Original line number Diff line number Diff line
@@ -111,12 +111,11 @@ def atom_to_id(atom):
  return features_to_id(features, intervals)


def atom_features(atom, bool_id_feat=False):
def atom_features(atom, bool_id_feat=False, explicit_H=False):
  if bool_id_feat:
    return np.array([atom_to_id(atom)])
  else:
    return np.array(
        one_of_k_encoding_unk(
    results = one_of_k_encoding_unk(
            atom.GetSymbol(),
            [
                'C',
@@ -165,14 +164,27 @@ def atom_features(atom, bool_id_feat=False):
                'Unknown'
            ]) + one_of_k_encoding(atom.GetDegree(), [
                0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10
            ]) + one_of_k_encoding_unk(atom.GetTotalNumHs(), [0, 1, 2, 3, 4]) +
        one_of_k_encoding_unk(atom.GetImplicitValence(), [0, 1, 2, 3, 4, 5, 6])
        + [atom.GetFormalCharge(), atom.GetNumRadicalElectrons()] +
            ])
    if explicit_H:
      results = results + \
        one_of_k_encoding_unk(atom.GetImplicitValence(), [0, 1, 2, 3, 4, 5, 6]) + \
        [atom.GetFormalCharge(), atom.GetNumRadicalElectrons()] + \
        one_of_k_encoding_unk(atom.GetHybridization(), [
            Chem.rdchem.HybridizationType.SP, Chem.rdchem.HybridizationType.SP2,
            Chem.rdchem.HybridizationType.SP3, Chem.rdchem.HybridizationType.
            SP3D, Chem.rdchem.HybridizationType.SP3D2
        ]) + [atom.GetIsAromatic()]
    else:
      results = results + one_of_k_encoding_unk(atom.GetTotalNumHs(), [0, 1, 2, 3, 4]) + \
          one_of_k_encoding_unk(atom.GetImplicitValence(), [0, 1, 2, 3, 4, 5, 6]) + \
          [atom.GetFormalCharge(), atom.GetNumRadicalElectrons()] + \
          one_of_k_encoding_unk(atom.GetHybridization(), [
              Chem.rdchem.HybridizationType.SP, Chem.rdchem.HybridizationType.SP2,
              Chem.rdchem.HybridizationType.SP3, Chem.rdchem.HybridizationType.
              SP3D, Chem.rdchem.HybridizationType.SP3D2
        ]) + [atom.GetIsAromatic()])
          ]) + [atom.GetIsAromatic()]

    return np.array(results)


def bond_features(bond):
@@ -184,10 +196,13 @@ def bond_features(bond):
  ])


def pair_features(mol, edge_list, canon_adj_list, bt_len=6):
def pair_features(mol, edge_list, canon_adj_list, bt_len=6, graph_distance=True):
  if graph_distance:
    max_distance = 7
  features = np.zeros(
      (mol.GetNumAtoms(), mol.GetNumAtoms(), bt_len + max_distance + 1))
  else:
    max_distance = 1
  N = mol.GetNumAtoms()
  features = np.zeros((N, N, bt_len + max_distance + 1))
  num_atoms = mol.GetNumAtoms()
  rings = mol.GetRingInfo().AtomRings()
  for a1 in range(num_atoms):
@@ -201,9 +216,18 @@ def pair_features(mol, edge_list, canon_adj_list, bt_len=6):
        features[a1, ring, bt_len] = 1
        features[a1, a1, bt_len] = 0.
    # graph distance between two atoms
    if graph_distance:
      distance = find_distance(
          a1, num_atoms, canon_adj_list, max_distance=max_distance)
      features[a1, :, bt_len + 1:] = distance
  if not graph_distance:
    coords = np.zeros((N, 3))
    for atom in range(N):
      pos = mol.GetConformer(0).GetAtomPosition(atom)
      coords[atom, :] = [pos.x, pos.y, pos.z]
    features[:,:,-1] = np.sqrt(np.sum(np.square(
        np.stack([coords] * N, axis=1) - \
        np.stack([coords] * N, axis=0)), axis=2))

  return features

@@ -263,14 +287,24 @@ class WeaveFeaturizer(Featurizer):

  name = ['weave_mol']

  def __init__(self):
  def __init__(self, graph_distance=True, explicit_H=None):
    # Set dtype
    self.graph_distance = graph_distance
    self.dtype = object
    self.check_H = False
    if explicit_H is None:
      self.explicit_H = False
      self.check_H = True

  def _featurize(self, mol):
    """Encodes mol as a WeaveMol object."""
    # Atom features
    idx_nodes = [(a.GetIdx(), atom_features(a)) for a in mol.GetAtoms()]
    if self.check_H and not self.explicit_H:
      for a in mol.GetAtoms():
        if a.GetSymbol() == 'H':
          self.explicit_H = True
          break
    idx_nodes = [(a.GetIdx(), atom_features(a, explicit_H=self.explicit_H)) for a in mol.GetAtoms()]
    idx_nodes.sort()  # Sort by ind to ensure same order as rd_kit
    idx, nodes = list(zip(*idx_nodes))

@@ -290,6 +324,6 @@ class WeaveFeaturizer(Featurizer):
      canon_adj_list[edge[1]].append(edge[0])

    # Calculate pair features
    pairs = pair_features(mol, edge_list, canon_adj_list, bt_len=6)
    pairs = pair_features(mol, edge_list, canon_adj_list, bt_len=6, graph_distance=self.graph_distance)

    return WeaveMol(nodes, pairs)
+1 −1
Original line number Diff line number Diff line
@@ -28,5 +28,5 @@ from deepchem.models.tensorflow_models.progressive_multitask import ProgressiveM
from deepchem.models.tensorflow_models.progressive_joint import ProgressiveJointRegressor
from deepchem.models.tensorflow_models.IRV import TensorflowMultiTaskIRVClassifier
from deepchem.models.tensorgraph.tensor_graph import TensorGraph
from deepchem.models.tensorgraph.models.graph_models import WeaveTensorGraph, DTNNTensorGraph, DAGTensorGraph, GraphConvTensorGraph
from deepchem.models.tensorgraph.models.graph_models import WeaveTensorGraph, DTNNTensorGraph, DAGTensorGraph, GraphConvTensorGraph, MPNNTensorGraph
from deepchem.models.tensorgraph.models.symmetry_function_regression import BPSymmetryFunctionRegression, ANIRegression
+189 −0
Original line number Diff line number Diff line
@@ -799,3 +799,192 @@ 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

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

  def __init__(self,
               M,
               batch_size,
               n_hidden=100,
               init='orthogonal',
               **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.M = M
    self.batch_size = batch_size
    self.n_hidden = n_hidden
    self.init = initializations.get(init)
    super(SetGather, self).__init__(**kwargs)

  def build(self, pair_features, n_pair_features):
    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)
    

  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)
    
    self.build()
    # Extract atom_features
    atom_features = in_layers[0].out_tensor
    atom_split = in_layers[1].out_tensor

    c = tf.zeros((self.batch_size, self.n_hidden))
    h = tf.zeros((self.batch_size, self.n_hidden))
    
    for i in range(self.M):
      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)
      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)

    out_tensor = q_star
    if set_tensors:
      self.variables = self.trainable_weights
      self.out_tensor = out_tensor
    return out_tensor

  def LSTMStep(self, h, c, x=None):

    # Taken from Keras code [citation needed]
    z = tf.nn.xw_plus_b(h, self.U, self.b)
    i = tf.nn.sigmoid(z[:, :self.n_hidden])
    f = tf.nn.sigmoid(z[:, self.n_hidden:2 * self.n_hidden])
    o = tf.nn.sigmoid(z[:, 2 * self.n_hidden:3 * self.n_hidden])
    z3 = z[:, 3 * self.n_hidden:]
    c_out = f * c + i * tf.nn.tanh(z3)
    h_out = o * tf.nn.tanh(c_out)

    return h_out, c_out
 No newline at end of file
+151 −2
Original line number Diff line number Diff line
@@ -5,8 +5,10 @@ import tensorflow as tf
from deepchem.feat.mol_graphs import ConvMol
from deepchem.metrics import to_one_hot, from_one_hot
from deepchem.models.tensorgraph.graph_layers import WeaveLayer, WeaveGather, \
    Combine_AP, Separate_AP, DTNNEmbedding, DTNNStep, DTNNGather, DAGLayer, DAGGather, DTNNExtract
from deepchem.models.tensorgraph.layers import Dense, Concat, SoftMax, SoftMaxCrossEntropy, GraphConv, BatchNorm, \
    Combine_AP, Separate_AP, DTNNEmbedding, DTNNStep, DTNNGather, DAGLayer, \
    DAGGather, DTNNExtract, MessagePassing, SetGather
from deepchem.models.tensorgraph.layers import Dense, Concat, SoftMax, \
    SoftMaxCrossEntropy, GraphConv, BatchNorm, \
    GraphPool, GraphGather, WeightedError, Dropout, BatchNormalization, Stack
from deepchem.models.tensorgraph.layers import L2Loss, Label, Weights, Feature
from deepchem.models.tensorgraph.tensor_graph import TensorGraph
@@ -677,3 +679,150 @@ class GraphConvTensorGraph(TensorGraph):
      y_ = undo_transforms(y_, transformers)

    return y_

class MPNNTensorGraph(TensorGraph):

  def __init__(self,
               n_tasks,
               batch_size,
               n_atom_feat=70,
               n_pair_feat=8,
               n_hidden=100,
               T=5,
               M=10,
               **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.batch_size = batch_size
    self.n_atom_feat = n_atom_feat
    self.n_pair_feat = n_pair_feat
    self.n_hidden = n_hidden
    self.T = T
    self.M = M
    super(MPNNTensorGraph, self).__init__(**kwargs)
    self.build_graph()

  def build_graph(self):
    self.atom_features = Feature(shape=(None, self.n_atom_feat))
    self.pair_features = Feature(shape=(None, self.n_pair_feat))
    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',
                                     n_hidden=self.n_hidden,
                                     in_layers=[self.atom_features,
                                                self.pair_features,
                                                self.atom_to_pair])
    atom_embeddings = Dense(self.n_hidden, in_layers=[message_passing])
    mol_embeddings = SetGather(self.M, 
                               self.batch_size, 
                               n_hidden=self.n_hidden,
                               in_layers=[atom_embeddings, self.atom_split])
    
    dense1 = Dense(out_channels=2*self.n_hidden, 
                   activation_fn=tf.nn.relu, 
                   in_layers=[mol_embeddings])
    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=[dense1])
        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=[dense1])
        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.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
Loading