forked from enviPath/enviPy
[Fix] Canonicalize SMILES when checking overlaps in smiles_covered_by (#443)
Co-authored-by: Tim Lorsbach <tim@lorsba.ch> Reviewed-on: enviPath/enviPy#443
This commit is contained in:
@ -68,6 +68,8 @@ class PredictionResult(object):
|
|||||||
|
|
||||||
|
|
||||||
class FormatConverter(object):
|
class FormatConverter(object):
|
||||||
|
tautomer_enumerator = rdMolStandardize.TautomerEnumerator()
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def mass(smiles):
|
def mass(smiles):
|
||||||
return Descriptors.MolWt(FormatConverter.from_smiles(smiles))
|
return Descriptors.MolWt(FormatConverter.from_smiles(smiles))
|
||||||
@ -240,8 +242,9 @@ class FormatConverter(object):
|
|||||||
Chem.RemoveStereochemistry(res_mol)
|
Chem.RemoveStereochemistry(res_mol)
|
||||||
|
|
||||||
if canonicalize_tautomers:
|
if canonicalize_tautomers:
|
||||||
te = rdMolStandardize.TautomerEnumerator() # idem
|
tautomers = FormatConverter.tautomer_enumerator.Enumerate(res_mol)
|
||||||
res_mol = te.Canonicalize(res_mol)
|
if len(tautomers) >= 1:
|
||||||
|
res_mol = FormatConverter.tautomer_enumerator.PickCanonical(tautomers)
|
||||||
|
|
||||||
return Chem.MolToSmiles(res_mol, kekuleSmiles=True)
|
return Chem.MolToSmiles(res_mol, kekuleSmiles=True)
|
||||||
|
|
||||||
@ -389,7 +392,7 @@ class FormatConverter(object):
|
|||||||
prods.append(p)
|
prods.append(p)
|
||||||
|
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
logger.error(f"Sanitizing and converting failed:\n{e}")
|
logger.debug(f"Sanitizing and converting failed:\n{e}")
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if len(prods):
|
if len(prods):
|
||||||
@ -397,7 +400,8 @@ class FormatConverter(object):
|
|||||||
pss.add(ps)
|
pss.add(ps)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Applying {smirks} on {smiles} failed:\n{e}")
|
logger.debug(f"Applying {smirks} on {smiles} failed:\n{e}")
|
||||||
|
pass
|
||||||
|
|
||||||
return list(pss)
|
return list(pss)
|
||||||
|
|
||||||
@ -482,8 +486,10 @@ class FormatConverter(object):
|
|||||||
if standardize:
|
if standardize:
|
||||||
for smi in l_smiles:
|
for smi in l_smiles:
|
||||||
try:
|
try:
|
||||||
smi = FormatConverter.standardize(
|
smi = FormatConverter.canonicalize(
|
||||||
smi, remove_stereo=True, canonicalize_tautomers=canonicalize_tautomers
|
FormatConverter.standardize(
|
||||||
|
smi, remove_stereo=True, canonicalize_tautomers=canonicalize_tautomers
|
||||||
|
)
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
# :shrug:
|
# :shrug:
|
||||||
@ -497,8 +503,12 @@ class FormatConverter(object):
|
|||||||
if standardize:
|
if standardize:
|
||||||
for smi in r_smiles:
|
for smi in r_smiles:
|
||||||
try:
|
try:
|
||||||
smi = FormatConverter.standardize(
|
smi = FormatConverter.canonicalize(
|
||||||
smi, remove_stereo=True, canonicalize_tautomers=canonicalize_tautomers
|
FormatConverter.standardize(
|
||||||
|
smi,
|
||||||
|
remove_stereo=True,
|
||||||
|
canonicalize_tautomers=canonicalize_tautomers,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
# :shrug:
|
# :shrug:
|
||||||
|
|||||||
@ -238,9 +238,11 @@ class RuleBasedDataset(Dataset):
|
|||||||
):
|
):
|
||||||
if feat_funcs is None:
|
if feat_funcs is None:
|
||||||
feat_funcs = [FormatConverter.maccs]
|
feat_funcs = [FormatConverter.maccs]
|
||||||
_structures = set() # Get all the structures
|
|
||||||
|
_structures = set()
|
||||||
for r in reactions:
|
for r in reactions:
|
||||||
_structures.update(r.educts.all())
|
_structures.update(r.educts.all())
|
||||||
|
|
||||||
if not educts_only:
|
if not educts_only:
|
||||||
_structures.update(r.products.all())
|
_structures.update(r.products.all())
|
||||||
|
|
||||||
@ -282,17 +284,14 @@ class RuleBasedDataset(Dataset):
|
|||||||
if key not in triggered:
|
if key not in triggered:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# standardize products from reactions for comparison
|
if FormatConverter.smiles_covered_by(
|
||||||
standardized_products = []
|
[prod.smiles for prod in r.products.all()],
|
||||||
for cs in r.products.all():
|
list(triggered[key]),
|
||||||
smi = cs.smiles
|
standardize=True,
|
||||||
try:
|
canonicalize_tautomers=True,
|
||||||
smi = FormatConverter.standardize(smi, remove_stereo=True)
|
):
|
||||||
except Exception:
|
|
||||||
logger.debug(f"Standardizing SMILES failed for {smi}")
|
|
||||||
standardized_products.append(smi)
|
|
||||||
if len(set(standardized_products).difference(triggered[key])) == 0:
|
|
||||||
observed.add(key)
|
observed.add(key)
|
||||||
|
|
||||||
feat_columns = []
|
feat_columns = []
|
||||||
for feat_func in feat_funcs:
|
for feat_func in feat_funcs:
|
||||||
if isinstance(feat_func, Descriptor):
|
if isinstance(feat_func, Descriptor):
|
||||||
@ -301,6 +300,7 @@ class RuleBasedDataset(Dataset):
|
|||||||
feats = feat_func(compounds[0].smiles)
|
feats = feat_func(compounds[0].smiles)
|
||||||
start_i = len(feat_columns)
|
start_i = len(feat_columns)
|
||||||
feat_columns.extend([f"feature_{start_i + i}" for i, _ in enumerate(feats)])
|
feat_columns.extend([f"feature_{start_i + i}" for i, _ in enumerate(feats)])
|
||||||
|
|
||||||
ds_columns = (
|
ds_columns = (
|
||||||
["structure_id"]
|
["structure_id"]
|
||||||
+ feat_columns
|
+ feat_columns
|
||||||
@ -334,7 +334,9 @@ class RuleBasedDataset(Dataset):
|
|||||||
obs.append(None)
|
obs.append(None)
|
||||||
else:
|
else:
|
||||||
obs.append(0)
|
obs.append(0)
|
||||||
|
|
||||||
rows.append([str(comp.uuid)] + feats + trig + obs)
|
rows.append([str(comp.uuid)] + feats + trig + obs)
|
||||||
|
|
||||||
ds = RuleBasedDataset(len(applicable_rules), ds_columns, data=rows)
|
ds = RuleBasedDataset(len(applicable_rules), ds_columns, data=rows)
|
||||||
return ds
|
return ds
|
||||||
|
|
||||||
@ -680,9 +682,13 @@ class RelativeReasoning:
|
|||||||
|
|
||||||
def predict(self, X):
|
def predict(self, X):
|
||||||
res = np.zeros((len(X), (self.end_index + 1 - self.start_index)))
|
res = np.zeros((len(X), (self.end_index + 1 - self.start_index)))
|
||||||
|
immutable_res = np.zeros((len(X), (self.end_index + 1 - self.start_index)))
|
||||||
|
|
||||||
# Loop through all instances
|
# Loop through all instances
|
||||||
for inst_idx, inst in enumerate(X):
|
for inst_idx, inst in enumerate(X):
|
||||||
|
for i, t in enumerate(inst[self.start_index : self.end_index + 1]):
|
||||||
|
immutable_res[inst_idx][i] = t
|
||||||
|
|
||||||
# Loop through all "triggered" features
|
# Loop through all "triggered" features
|
||||||
for i, t in enumerate(inst[self.start_index : self.end_index + 1]):
|
for i, t in enumerate(inst[self.start_index : self.end_index + 1]):
|
||||||
# Set label
|
# Set label
|
||||||
@ -696,7 +702,7 @@ class RelativeReasoning:
|
|||||||
if i2 in self.winmap.get(i, []):
|
if i2 in self.winmap.get(i, []):
|
||||||
# if thatat rule also triggered, it dominated the current
|
# if thatat rule also triggered, it dominated the current
|
||||||
# set label to 0
|
# set label to 0
|
||||||
if X[inst_idx][i2]:
|
if immutable_res[inst_idx][i2]:
|
||||||
res[inst_idx][i] = 0
|
res[inst_idx][i] = 0
|
||||||
|
|
||||||
return res
|
return res
|
||||||
|
|||||||
Reference in New Issue
Block a user