diff --git a/src/xlm/configs/lightning_train/post_hoc_evaluator/denovo.yaml b/src/xlm/configs/lightning_train/post_hoc_evaluator/denovo.yaml index c7bb63a0..56222cec 100644 --- a/src/xlm/configs/lightning_train/post_hoc_evaluator/denovo.yaml +++ b/src/xlm/configs/lightning_train/post_hoc_evaluator/denovo.yaml @@ -5,3 +5,4 @@ compute_validity: True compute_uniqueness: True compute_qed: True compute_sa: True +convert_to_smile: True diff --git a/src/xlm/tasks/molgen.py b/src/xlm/tasks/molgen.py index f4437218..8a64ad97 100644 --- a/src/xlm/tasks/molgen.py +++ b/src/xlm/tasks/molgen.py @@ -502,6 +502,7 @@ def __init__( compute_uniqueness: bool = True, compute_qed: bool = True, compute_sa: bool = True, + convert_to_smile: bool = True ): self.use_bracket_safe = use_bracket_safe self.compute_diversity = compute_diversity @@ -509,6 +510,7 @@ def __init__( self.compute_uniqueness = compute_uniqueness self.compute_qed = compute_qed self.compute_sa = compute_sa + self.convert_to_smile = convert_to_smile # Lazy-loaded TDC oracles (to avoid import overhead) self._oracle_qed = None @@ -575,7 +577,7 @@ def eval( all_smiles = [] for pred in predictions: safe_str = pred.get("text", "") - smiles = self._safe_to_smiles_with_bracket_handling(safe_str) + smiles = self._safe_to_smiles_with_bracket_handling(safe_str) if self.convert_to_smile else safe_str pred["smiles"] = smiles # Add SMILES to prediction dict if ( smiles is not None