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