import html import json import logging import os import re from collections import defaultdict from typing import Dict, List, Tuple import pandas as pd from gt4sd.algorithms import ( RegressionTransformerMolecules, RegressionTransformerProteins, ) from gt4sd.algorithms.core import AlgorithmConfiguration from rdkit import Chem from rdkit.Chem import Draw from terminator.selfies import decoder logger = logging.getLogger(__name__) logger.addHandler(logging.NullHandler()) _SELFIES_TOKEN_PATTERN = re.compile(r"\[[^\]]+\]") _PROPERTY_TOKEN_PATTERN = re.compile(r"<[^>]+>") def _extract_molecule_sequence(sequence: str) -> str: """Extract the molecule SELFIES/SMILES part from an RT input sequence.""" sequence = sequence.split("|")[-1].strip() without_properties = _PROPERTY_TOKEN_PATTERN.sub("", sequence) selfies_tokens = [ token for token in _SELFIES_TOKEN_PATTERN.findall(without_properties) if token != "[MASK]" ] return "".join(selfies_tokens) or without_properties.replace("[MASK]", "").strip() def _sequence_to_smiles(sequence: str, domain: str) -> str: if domain == "Proteins": mol = Chem.MolFromFASTA(sequence) if mol is None: raise ValueError(f"Could not parse protein sequence {sequence}") return Chem.MolToSmiles(mol) molecule_sequence = _extract_molecule_sequence(sequence) candidates = [molecule_sequence] if molecule_sequence != sequence: candidates.append(sequence) for candidate in candidates: mol = Chem.MolFromSmiles(candidate) if mol is not None: return Chem.MolToSmiles(mol) if "[" in candidate: try: smiles = decoder(candidate) except Exception: continue mol = Chem.MolFromSmiles(smiles) if mol is not None: return Chem.MolToSmiles(mol) raise ValueError(f"Could not parse molecule sequence {sequence}") def _draw_molecule_svg(smiles: str, size: Tuple[int, int]) -> str: mol = Chem.MolFromSmiles(smiles) if mol is None: return f"
{html.escape(smiles)}
" svg = Draw.MolsToGridImage([mol], molsPerRow=1, subImgSize=size, useSVG=True) return str(svg).replace("", "") def _draw_static_grid( result_df: pd.DataFrame, n_cols: int, size: Tuple[int, int] ) -> str: columns = [column for column in result_df.columns if column != "SMILES"] if "Name" in columns: columns = ["Name"] + [column for column in columns if column != "Name"] cards = [] for _, row in result_df.iterrows(): smiles = str(row["SMILES"]) details = "".join( f"
{html.escape(str(column))}
" f"
{html.escape(str(row[column]))}
" for column in columns ) details += ( "
SMILES
" f"
{html.escape(smiles)}
" ) cards.append( "
" f"
{_draw_molecule_svg(smiles, size)}
" f"
{details}
" "
" ) return f"""
{''.join(cards)}
""" def get_application(application: str) -> AlgorithmConfiguration: """ Convert application name to AlgorithmConfiguration. Args: application: Molecules or Proteins Returns: The corresponding AlgorithmConfiguration """ if application == "Molecules": application = RegressionTransformerMolecules elif application == "Proteins": application = RegressionTransformerProteins else: raise ValueError( "Currently only models for molecules and proteins are supported" ) return application def get_inference_dict( application: AlgorithmConfiguration, algorithm_version: str ) -> Dict: """ Get inference dictionary for a given application and algorithm version. Args: application: algorithm application (Molecules or Proteins) algorithm_version: algorithm version (e.g. qed) Returns: A dictionary with the inference parameters. """ config = application(algorithm_version=algorithm_version) with open(os.path.join(config.ensure_artifacts(), "inference.json"), "r") as f: data = json.load(f) return data def get_rt_name(x: Dict) -> str: """ Get the UI display name of the regression transformer. Args: x: dictionary with the inference parameters Returns: The display name """ return ( x["algorithm_application"].split("Transformer")[-1] + ": " + x["algorithm_version"].capitalize() ) def draw_grid_predict(prediction: str, target: str, domain: str) -> str: """ Uses RDKit SVGs to draw a HTML grid for the prediction Args: prediction: Predicted sequence. target: Target molecule domain: Domain of the prediction (molecules or proteins) Returns: HTML to display """ if domain not in ["Molecules", "Proteins"]: raise ValueError(f"Unsupported domain {domain}") seq = target.split("|")[-1] try: seq = _sequence_to_smiles(seq, domain=domain) except Exception: logger.warning(f"Could not draw sequence {seq}") result = {"SMILES": [seq], "Name": ["Target"]} # Add properties for prop in prediction.split("<")[1:]: result[ prop.split(">")[0] ] = f"{prop.split('>')[0].capitalize()} = {prop.split('>')[1]}" result_df = pd.DataFrame(result) return _draw_static_grid(result_df, n_cols=1, size=(600, 700)) def draw_grid_generate( samples: List[Tuple[str]], domain: str, n_cols: int = 5, size=(140, 200) ) -> str: """ Uses RDKit SVGs to draw a HTML grid for the generated molecules Args: samples: The generated samples (with properties) domain: Domain of the prediction (molecules or proteins) n_cols: Number of columns in grid. Defaults to 5. size: Size of molecule in grid. Defaults to (140, 200). Returns: HTML to display """ if domain not in ["Molecules", "Proteins"]: raise ValueError(f"Unsupported domain {domain}") smis = [] for sample in samples: sequence = sample[0] try: smis.append(_sequence_to_smiles(sequence, domain=domain)) except Exception: logger.warning(f"Could not convert sequence {sequence}") smis.append(sequence) result = defaultdict(list) result.update({"SMILES": smis, "Name": [f"sample_{i}" for i in range(len(smis))]}) # Create properties properties = [s.split("<")[1] for s in samples[0][1].split(">")[:-1]] # Fill properties for sample in samples: for prop in properties: value = float(sample[1].split(prop)[-1][1:].split("<")[0]) result[prop].append(f"{prop} = {value}") result_df = pd.DataFrame(result) return _draw_static_grid(result_df, n_cols=n_cols, size=size)