[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:
2026-08-12 21:01:41 +12:00
parent 093daa5ecf
commit ada270aa3c
2 changed files with 36 additions and 20 deletions

View File

@ -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:

View File

@ -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