Spaces:
Sleeping
Sleeping
File size: 5,567 Bytes
4d0da28 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 | import torch
import numpy as np
from rdkit import Chem
from rdkit.Chem.Descriptors import ExactMolWt
from torch_geometric.data import Data as GraphData
ALLOWED_ATOMIC_NUMS = {6, 7, 8, 9, 15, 16, 17, 35}
MIN_NUM_ATOMS = 3
MAX_NUM_ATOMS = 60
def get_node_feats(mol: Chem.Mol, atom: Chem.Atom):
'''
:param mol: The mol from which the node features will be computed. Currently not used
:param atom: The atom from which the node features will be computed
:return:
'''
atom_num = atom.GetAtomicNum()
valence = atom.GetTotalValence()
charge = atom.GetFormalCharge()
degree = atom.GetDegree()
is_aromatic = atom.GetIsAromatic()
return [float(elem) for elem in [atom_num, valence, charge, degree, is_aromatic]]
# return [float(elem) for elem in [atom_num, valence, charge, is_aromatic]]
# results = one_of_k_encoding_unk(
# atom.GetSymbol(),
# [
# 'B',
# 'C',
# 'N',
# 'O',
# 'F',
# 'Si',
# 'P',
# 'S',
# 'Cl',
# 'As',
# 'Se',
# 'Br',
# 'Te',
# 'I',
# 'At',
# 'other'
# ]) + one_of_k_encoding(atom.GetDegree(),
# [0, 1, 2, 3, 4, 5]) + \
# [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, 'other'
# ]) + [atom.GetIsAromatic()]
def get_bond_feats(mol: Chem.Mol, bond: Chem.Bond):
'''
:param mol: The mol from which the edge features will be computed. Currently not used
:param bond: The bond from which the edge features will be computed
:return:
'''
if bond.GetIsAromatic():
bond_num = 1.5
else:
bond_num = bond.GetBondType()
bond_num = float(bond_num)
if bond_num > 3:
bond_num = 4
# return [bond_num ]
bondInfo = [bond_num, int(bond.GetIsAromatic()), int(bond.GetIsConjugated()), int(bond.IsInRing()) ]
# bondInfo += [["STEREONONE", "STEREOANY", "STEREOZ", "STEREOE"].index(str(bond.GetStereo()))]
return bondInfo
def smiles_to_graph(smi):
mol = Chem.MolFromSmiles(smi)
if mol is None:
return None
# remove stereo information, such as inward and outward edges
Chem.RemoveStereochemistry(mol)
nodes = []
edges_index, edges_attr = [], []
cur_atom_id = 0
rdkit_atomId_to_atom_id = {}
for bond in mol.GetBonds():
atom1_idx = bond.GetBeginAtomIdx()
atom2_idx = bond.GetEndAtomIdx()
atom1 = mol.GetAtomWithIdx(atom1_idx)
atom2 = mol.GetAtomWithIdx(atom2_idx)
atomicNum1 = atom1.GetAtomicNum()
atomicNum2 = atom2.GetAtomicNum()
if atomicNum1 == 0 or atomicNum2 == 0:
continue
if atomicNum1 not in ALLOWED_ATOMIC_NUMS or atomicNum2 not in ALLOWED_ATOMIC_NUMS:
return None
if atom1_idx not in rdkit_atomId_to_atom_id:
rdkit_atomId_to_atom_id[atom1_idx] = cur_atom_id
nodes.append( get_node_feats(mol, atom1) )
cur_atom_id += 1
if atom2_idx not in rdkit_atomId_to_atom_id:
rdkit_atomId_to_atom_id[atom2_idx] = cur_atom_id
nodes.append( get_node_feats(mol, atom2) )
cur_atom_id += 1
idx1, idx2 = rdkit_atomId_to_atom_id[atom1_idx], rdkit_atomId_to_atom_id[atom2_idx]
# print(atom1_idx, idx1, atom2_idx, idx2)
edges_index.append([idx1, idx2])
edges_index.append([idx2, idx1])
bond_feats = get_bond_feats(mol, bond)
edges_attr.extend([bond_feats, bond_feats])
x = np.array(nodes, dtype=np.float32)
if x.shape[0] < MIN_NUM_ATOMS or x.shape[0] > MAX_NUM_ATOMS:
return None
edges_index = np.array( edges_index).T.copy().astype( np.int64 )
edges_attr = np.array(edges_attr, dtype=np.float32)
graph_dict = dict(x=x, edge_index=edges_index, edge_attr=edges_attr)
graph = GraphData(**{key: torch.tensor(val) for key, val in graph_dict.items()} )
return graph
def fromPerGramToPerMMolPrice(price, smi):
mol = Chem.MolFromSmiles(smi)
if mol is None:
return None
return price * ExactMolWt(mol) / 1000
def compute_nodes_degree(graphs, max_degree=7):
from torch_geometric.utils import degree
import torch
deg = torch.zeros(max_degree, dtype=torch.long)
if graphs is None:
return deg
for data in graphs:
d = degree(data.edge_index[1], num_nodes=data.num_nodes, dtype=torch.long)
deg += torch.bincount(d, minlength=deg.numel())[:deg.numel()]
return deg
def test(smi = "C1=CC=C(C=C1)CC(C(=O)O)N"):
mol = Chem.MolFromSmiles(smi)
# print(mol)
graph = smiles_to_graph(smi)
print(graph.keys)
print(graph.x)
print(graph.edge_index)
print( graph.edge_attr )
print( graph.x.shape, graph.edge_index.shape, graph.edge_attr.shape)
for idx in range(mol.GetNumAtoms()):
mol.GetAtomWithIdx(idx).SetProp('molAtomMapNumber', str(mol.GetAtomWithIdx(idx).GetIdx()))
from matplotlib import pyplot as plt; from rdkit.Chem import Draw; plt.imshow(Draw.MolsToGridImage([mol], molsPerRow=2)); plt.show()
if __name__ == "__main__":
smis_list = ["[C@H](C)1CCCO1", "O[C@@H](N)C", "C1=CC=C(C=C1)CC(C(=O)O)N"]
for smi in smis_list:
test(smi) |