[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

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