Commit 8738a7b5 authored by leswing's avatar leswing
Browse files

Test Featurizer

parent 5add5bac
Loading
Loading
Loading
Loading
+2 −3
Original line number Diff line number Diff line
@@ -268,11 +268,10 @@ class ConvMol(object):

    num_mols = len(mols)

    atoms_per_mol = [mol.get_num_atoms() for mol in mols]

    # Get atoms by degree
    # Results should be sorted by (atom_degree, mol_index)
    atoms_by_deg = np.concatenate([x.atom_features for x in mols])
    degree_vector = np.concatenate([x.degree_list for x in mols], axis=0)
    # Mergesort is a "stable" sort, so the array maintains it's secondary sort of mol_index
    all_atoms = atoms_by_deg[degree_vector.argsort(kind='mergesort')]

    # Sort all atoms by degree.
+41 −0
Original line number Diff line number Diff line
from unittest import TestCase

import numpy as np
from deepchem.feat import ConvMolFeaturizer
from deepchem.feat.mol_graphs import ConvMol
from deepchem.molnet import load_bace_classification


class TestConvMol(TestCase):

  def get_molecules(self):
    tasks, all_dataset, transformers = load_bace_classification(
        featurizer="Raw")
    return all_dataset[0].X

  def test_mol_ordering(self):
    mols = self.get_molecules()
    featurizer = ConvMolFeaturizer()
    featurized_mols = featurizer.featurize(mols)
    for i in range(len(featurized_mols)):
      atom_features = featurized_mols[i].atom_features
      degree_list = np.expand_dims(featurized_mols[i].degree_list, axis=1)
      atom_features = np.concatenate([degree_list, atom_features], axis=1)
      featurized_mols[i].atom_features = atom_features

    conv_mol = ConvMol.agglomerate_mols(featurized_mols)

    for start, end in conv_mol.deg_slice.tolist():
      members = conv_mol.membership[start:end]
      sorted_members = np.array(sorted(members))
      members = np.array(members)
      self.assertTrue(np.all(sorted_members == members))

    conv_mol_atom_features = conv_mol.get_atom_features()

    adj_number = 0
    for start, end in conv_mol.deg_slice.tolist():
      deg_features = conv_mol_atom_features[start:end]
      adj_number_array = deg_features[:, 0]
      self.assertTrue(np.all(adj_number_array == adj_number))
      adj_number += 1