From ada270aa3c9a28d1931e7cd42c1599e63c7dca8b Mon Sep 17 00:00:00 2001 From: jebus Date: Wed, 12 Aug 2026 21:01:41 +1200 Subject: [PATCH] [Fix] Canonicalize SMILES when checking overlaps in smiles_covered_by (#443) Co-authored-by: Tim Lorsbach Reviewed-on: https://git.envipath.com/enviPath/enviPy/pulls/443 --- utilities/chem.py | 26 ++++++++++++++++++-------- utilities/ml.py | 30 ++++++++++++++++++------------ 2 files changed, 36 insertions(+), 20 deletions(-) diff --git a/utilities/chem.py b/utilities/chem.py index cbfba09e..476b450b 100644 --- a/utilities/chem.py +++ b/utilities/chem.py @@ -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,8 +486,10 @@ class FormatConverter(object): if standardize: for smi in l_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: @@ -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: diff --git a/utilities/ml.py b/utilities/ml.py index 06e80dcf..a633ee2e 100644 --- a/utilities/ml.py +++ b/utilities/ml.py @@ -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