dmlc--dgl
3192beb42d
* jtnn model zoo * poke ci * fix line sep * fix * Fix import order * fix render * fix render * revert * fix * Resolve conflict * dix * remove create_var * refactor * fix * refactor * readme * format * fix lint * fix lint * pylint * lint * fix lint * fix lint * add hint * fix * Remove vocab * Add explanation for warning * add directory * Load model to cpu by default * Update
32 行
763 B
Python
32 行
763 B
Python
# pylint: disable=C0111, C0103, E1101, W0611, W0612
|
|
import copy
|
|
|
|
import rdkit
|
|
import rdkit.Chem as Chem
|
|
|
|
|
|
def get_slots(smiles):
|
|
mol = Chem.MolFromSmiles(smiles)
|
|
return [(atom.GetSymbol(), atom.GetFormalCharge(), atom.GetTotalNumHs())
|
|
for atom in mol.GetAtoms()]
|
|
|
|
|
|
class Vocab(object):
|
|
|
|
def __init__(self, smiles_list):
|
|
self.vocab = smiles_list
|
|
self.vmap = {x: i for i, x in enumerate(self.vocab)}
|
|
self.slots = [get_slots(smiles) for smiles in self.vocab]
|
|
|
|
def get_index(self, smiles):
|
|
return self.vmap[smiles]
|
|
|
|
def get_smiles(self, idx):
|
|
return self.vocab[idx]
|
|
|
|
def get_slots(self, idx):
|
|
return copy.deepcopy(self.slots[idx])
|
|
|
|
def size(self):
|
|
return len(self.vocab)
|