[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):
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:

View File

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