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)