multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
62ec035 verified
Raw History Blame Contribute Delete
33.6 kB
"""E-SMILES caption parsing & refactoring used by the inference postprocess."""
from __future__ import annotations
import logging
import re
from dataclasses import dataclass, replace
from enum import Enum, unique
from typing import Dict, List, Literal, Mapping, Optional, Sequence, Tuple, Union
from rdkit import Chem, RDLogger
try:
from . import chem_utils
except ImportError: # Support running from package directory as working directory.
import chem_utils
# Ambiguous element symbols are omitted.
PERIODIC_TABLE = {
"H", "He", "Li", "Be", "C", "N", "O", "F", "Ne",
"Na", "Mg", "Al", "Si", "P", "S", "Cl", "K", "Ca",
"Sc", "Ti", "Cr", "Mn", "Fe", "Co", "Ni", "Cu", "Zn",
"Ga", "Ge", "As", "Se", "Br", "Kr", "Rb", "Sr", "Zr",
"Nb", "Mo", "Tc", "Ru", "Rh", "Pd", "Ag", "Cd", "In", "Sn",
"Sb", "Te", "I", "Xe", "Cs", "Ba", "La", "Ce", "Nd",
"Pm", "Sm", "Eu", "Gd", "Tb", "Dy", "Ho", "Er", "Tm", "Yb",
"Lu", "Hf", "Ta", "Re", "Os", "Ir", "Pt", "Au", "Hg",
"Tl", "Pb", "Bi", "Po", "At", "Rn", "Fr", "Ra", "Th",
"Pa",
}
LONG_CHAR_ELEMENTS = (
"Cl", "Br", "Si", "Na", "Al", "Ca", "Sn", "As", "Hg",
"Fe", "Zn", "Cr", "Se", "Gd", "Au", "Li",
)
logger = logging.getLogger(__name__)
_MAX_AUTOMATIC_REPEAT_COUNT = 1024
def _is_safe_carbon_chain_repeat_target(atom: Chem.rdchem.Atom) -> bool:
"""Return whether ``?N`` chain insertion preserves the target chemistry."""
if atom.GetSymbol() not in {"*", "C"}:
return False
if atom.GetDegree() not in {1, 2} or atom.GetIsAromatic() or atom.IsInRing():
return False
if (
atom.GetIsotope() != 0
or atom.GetFormalCharge() != 0
or atom.GetNumRadicalElectrons() != 0
or atom.GetChiralTag() != Chem.ChiralType.CHI_UNSPECIFIED
):
return False
return all(
bond.GetBondType() == Chem.BondType.SINGLE and not bond.GetIsAromatic()
for bond in atom.GetBonds()
)
@unique
class TextType(Enum):
SYMBOL = "symbol"
SCRIPT = "script"
MULTIPLE = "multiple"
PRIME = "prime"
class Index(int):
def __new__(cls, value: int, **kwds):
assert isinstance(value, int) and value >= 0
return super().__new__(cls, value)
class AtomIndex(Index):
key: int = 1
def __hash__(self):
return hash(("atom", int(self)))
def __eq__(self, other):
if isinstance(other, AtomIndex):
return hash(self) == hash(other)
return False
class RingIndex(Index):
def __init__(self, value: int, virtual: bool = False) -> None:
self._virtual = virtual
self._key = 2_000_000 if self._virtual else 1_000_000
def __hash__(self):
return hash(("ring", int(self), self._virtual))
def __eq__(self, other):
if isinstance(other, RingIndex):
return hash(self) == hash(other)
return False
@property
def virtual(self) -> bool:
return self._virtual
@property
def key(self) -> bool | int:
return self._key
class Tokens:
"""Special token names used in the captioning format."""
atom_start = "<a>"
atom_end = "</a>"
circ_start = "<c>"
circ_end = "</c>"
dummy_start = "<d>"
dummy_end = "</d>"
substruct_start = "<s>"
substruct_end = "</s>"
sgroup_start = "<g>"
sgroup_end = "</g>"
virtual_start = "<v>"
virtual_end = "</v>"
special_id = "<id>"
ring_start = "<r>"
ring_end = "</r>"
axial_start = "<x>"
axial_end = "</x>"
dummy = "<dum>"
separator = "<sep>"
class Patterns:
"""Compiled regexes for caption parsing."""
long_char_elements_pattern = re.compile(
rf'{"|".join(LONG_CHAR_ELEMENTS) + "|."}'
)
grp_content = re.compile(
rf"(?P<{TextType.SYMBOL.value}>(?:{Tokens.special_id}|[A-Za-z0-9-\(\)]*))"
+ rf"(?P<{TextType.SCRIPT.value}>(\[\S+\])?)"
+ rf"(?P<{TextType.PRIME.value}>[\'\"]?)"
+ rf"(?P<{TextType.MULTIPLE.value}>(\?([a-z]|\d+|\d+-\d+)$)?)"
)
grp_pattern = re.compile(
rf"({Tokens.atom_start}|{Tokens.circ_start}|{Tokens.dummy_start}|{Tokens.ring_start}|{Tokens.ring_start}{Tokens.circ_start}|{Tokens.ring_start}{Tokens.virtual_start})"
+ r"(\d+:\S+?)"
+ rf"({Tokens.atom_end}|{Tokens.circ_end}|{Tokens.dummy_end}|{Tokens.ring_end})"
)
trail_pattern = re.compile(
r"(?P<groups>.*?)(?P<extension>\|Sg:[^|]+\|)?$",
re.DOTALL,
)
@dataclass
class GroupDesc:
id: Index
symbol: Optional[str] = None
script: Optional[str] = None
prime: Optional[str] = None
multiple: Optional[str] = None
is_circle: bool = False
is_dummy: bool = False
def __str__(self) -> str:
if self.is_dummy:
return "<dum>"
expr = ""
if self.is_circle:
expr = "c"
if self.symbol is not None:
expr += self.symbol
if self.script is not None:
expr = expr + "[" + self.script + "]"
if self.prime is not None:
expr += self.prime
if self.multiple is not None:
expr = expr + "?" + self.multiple
return expr
@dataclass(frozen=True)
class TranslatedMolecule:
smi: str
groups: str
caption: str
esmi: str
markush: bool
sru: bool
class Translator:
"""Caption parsing and refactor (E-SMILES) for inference postprocess."""
@classmethod
def canonicalize_smiles(cls, smi: str) -> str:
try:
mol = chem_utils.parse_smiles(smi)
if mol is None:
return smi
return Chem.MolToSmiles(mol, canonical=True, isomericSmiles=True)
except Exception:
return smi
@classmethod
def build_esmi(cls, smi: str, groups: str = "", ext: str = "") -> str:
return f"{smi}{Tokens.separator}{groups}{ext}"
@classmethod
def remove_atom_groups(cls, groups: str, atom_indices: set[int]) -> str:
if not atom_indices:
return groups
atom_group_pattern = re.compile(
rf"(?P<start>{Tokens.atom_start}|{Tokens.dummy_start})"
r"(?P<idx>\d+):.+?"
rf"(?P<end>{Tokens.atom_end}|{Tokens.dummy_end})"
)
def replace(match: re.Match) -> str:
return "" if int(match.group("idx")) in atom_indices else match.group(0)
return atom_group_pattern.sub(replace, groups)
@classmethod
def parse_caption(
cls,
caption: str,
return_mol: bool = False,
error_msg: bool = False,
) -> Optional[Tuple[Union[Chem.rdchem.Mol, str], str, str]]:
if Tokens.separator not in caption:
if error_msg:
logger.warning(f"No `{Tokens.separator}` found in caption: {caption}")
return
smi, trailing = caption.split(Tokens.separator, 1)
if error_msg:
RDLogger.EnableLog("rdApp.*")
mol = chem_utils.parse_smiles(smi)
RDLogger.DisableLog("rdApp.*")
if mol is None:
if error_msg:
logger.warning(f"Invalid SMILES: {smi}")
return
groups, ext = cls.parse_trailing(trailing)
if return_mol:
return mol, groups, ext
return smi, groups, ext
@classmethod
def parse_trailing(cls, trailing: str):
matched = re.match(Patterns.trail_pattern, trailing)
if matched is None:
return "", ""
content = matched.groupdict()
return content.get("groups") or "", content.get("extension") or ""
@classmethod
def parse_extension(cls, ext: str) -> str:
return ext.strip("|")
@classmethod
def has_symbolic_sru(cls, caption: str) -> bool:
"""Classify whole-molecule SRUs with a symbolic repeat count.
Local/nested repeats and numeric counts (including ranges) do not set
the sru flag; their structural annotations remain independently valid.
"""
if Tokens.separator not in caption:
return False
trailing = caption.split(Tokens.separator, 1)[1]
_, ext = cls.parse_trailing(trailing)
matched = re.fullmatch(
r"Sg:([A-Za-z0-9]+(?:-[A-Za-z0-9]+)?)", cls.parse_extension(ext)
)
return bool(
matched and re.fullmatch(r"\d+(?:-\d+)?", matched.group(1)) is None
)
@classmethod
def parse_groups(cls, seq: str) -> List[GroupDesc]:
if seq == "":
return []
seq = re.sub(
rf"{Tokens.substruct_start}.*?{Tokens.substruct_end}|"
rf"{Tokens.sgroup_start}.*?{Tokens.sgroup_end}|"
rf"{Tokens.virtual_start}.*?{Tokens.virtual_end}|"
rf"{Tokens.axial_start}.*?{Tokens.axial_end}",
"",
seq,
flags=re.DOTALL,
)
descriptions: List[GroupDesc] = []
for grp_start, grp_content, _ in re.findall(Patterns.grp_pattern, seq):
parsed = cls.parse_group(grp_content)
if parsed is None:
continue
idx, grp_text = parsed
if grp_start in (Tokens.atom_start, Tokens.dummy_start):
grp_desc = GroupDesc(id=AtomIndex(idx))
if len(grp_text) == 0 or grp_start == Tokens.dummy_start:
grp_desc.is_dummy = True
elif grp_start == Tokens.circ_start:
grp_desc = GroupDesc(id=AtomIndex(idx), is_circle=True)
elif grp_start in (
f"{Tokens.ring_start}{Tokens.circ_start}",
f"{Tokens.ring_start}{Tokens.virtual_start}",
):
grp_desc = GroupDesc(id=RingIndex(idx, virtual=True))
elif grp_start == Tokens.ring_start:
grp_desc = GroupDesc(id=RingIndex(idx))
else:
continue
grp_desc.symbol = grp_text.get(TextType.SYMBOL)
grp_desc.script = grp_text.get(TextType.SCRIPT)
grp_desc.prime = grp_text.get(TextType.PRIME)
grp_desc.multiple = grp_text.get(TextType.MULTIPLE)
descriptions.append(grp_desc)
return descriptions
@classmethod
def parse_group(cls, group: str) -> Optional[Tuple[int, Dict[TextType, str]]]:
items = group.split(":")
if len(items) != 2:
return
idx, content = items
if not idx.isdigit():
return
idx = int(idx)
if content == Tokens.dummy:
return idx, {}
grp_text = cls.get_group_texts(content)
if len(grp_text) == 0:
return
return idx, grp_text
@classmethod
def get_group_texts(cls, content: str) -> Dict[TextType, str]:
texts: Dict[TextType, str] = {}
matched = re.match(Patterns.grp_content, content).groupdict()
for tt_value, text in matched.items():
if len(text) == 0:
continue
if tt_value == TextType.SYMBOL.value:
texts[TextType.SYMBOL] = text
elif tt_value == TextType.SCRIPT.value:
assert text.startswith("[") and text.endswith("]")
texts[TextType.SCRIPT] = text[1:-1]
elif tt_value == TextType.PRIME.value:
texts[TextType.PRIME] = text
elif tt_value == TextType.MULTIPLE.value:
assert text.startswith("?")
texts[TextType.MULTIPLE] = text[1:]
return texts
@classmethod
def repair_atom_group_indices(
cls,
mol: Chem.rdchem.Mol,
trailing: str,
error_msg: bool = False,
) -> str:
"""Repair atom-group tags that miss their dummy atom."""
star_indices = [atom.GetIdx() for atom in mol.GetAtoms() if atom.GetSymbol() == "*"]
if not star_indices or (
Tokens.atom_start not in trailing and Tokens.dummy_start not in trailing
):
return trailing
preserved_records: List[str] = []
preserved_pattern = re.compile(
rf"{Tokens.substruct_start}.*?{Tokens.substruct_end}|"
rf"{Tokens.sgroup_start}.*?{Tokens.sgroup_end}|"
rf"{Tokens.virtual_start}.*?{Tokens.virtual_end}|"
rf"{Tokens.axial_start}.*?{Tokens.axial_end}|"
rf"{Tokens.ring_start}{Tokens.virtual_start}\d+:.+?{Tokens.ring_end}",
re.DOTALL,
)
def hide_preserved(match: re.Match) -> str:
preserved_records.append(match.group(0))
return f"@@MOLPARSER_PRESERVED_{len(preserved_records) - 1}@@"
protected_trailing = preserved_pattern.sub(hide_preserved, trailing)
atom_group_pattern = re.compile(
rf"(?P<start>{Tokens.atom_start}|{Tokens.dummy_start})"
r"(?P<idx>\d+):(?P<content>.+?)"
rf"(?P<end>{Tokens.atom_end}|{Tokens.dummy_end})"
)
matches = list(atom_group_pattern.finditer(protected_trailing))
if not matches:
return trailing
exact_targets = {
int(match.group("idx"))
for match in matches
if int(match.group("idx")) in star_indices
}
used_targets: set[int] = set()
def replace(match: re.Match) -> str:
raw_idx, content = match.group("idx"), match.group("content")
idx = int(raw_idx)
if idx in star_indices and idx not in used_targets:
used_targets.add(idx)
return match.group(0)
candidates = [
star_idx
for star_idx in star_indices
if star_idx not in used_targets and star_idx not in exact_targets
]
if not candidates:
return match.group(0)
nearest = min(candidates, key=lambda star_idx: (abs(star_idx - idx), star_idx))
used_targets.add(nearest)
if error_msg:
logger.warning(
"Repair atom group index: <a>%s:%s</a> -> <a>%s:%s</a>",
raw_idx,
content,
nearest,
content,
)
return f"{match.group('start')}{nearest}:{content}{match.group('end')}"
repaired = atom_group_pattern.sub(replace, protected_trailing)
for idx, record in enumerate(preserved_records):
repaired = repaired.replace(f"@@MOLPARSER_PRESERVED_{idx}@@", record)
return repaired
@classmethod
def _protect_stereo_for_dummy_substitution(
cls,
smi: str,
groups: str,
abbrev_map: Dict[str, str],
probe: Chem.rdchem.Mol,
error_msg: bool = False,
) -> Tuple[str, set[int]]:
"""Temporarily distinguish substitutable dummy atoms for RDKit stereo parsing."""
if "@" not in smi or "*" not in smi:
return smi, set()
star_indices = [atom.GetIdx() for atom in probe.GetAtoms() if atom.GetSymbol() == "*"]
if len(star_indices) < 2:
return smi, set()
substitutable_indices: set[int] = set()
for desc in cls.parse_groups(groups):
if not isinstance(desc.id, AtomIndex) or desc.is_dummy or not desc.symbol:
continue
lookup_symbol = desc.symbol + (desc.script or "")
if (
desc.symbol == "CH2"
or (desc.symbol == "CH" and desc.script == "2")
or lookup_symbol == "CN"
or lookup_symbol in abbrev_map
):
substitutable_indices.add(int(desc.id))
elif lookup_symbol in PERIODIC_TABLE and not desc.multiple:
substitutable_indices.add(int(desc.id))
if not substitutable_indices.intersection(star_indices):
return smi, set()
dummy_tokens: List[Tuple[int, int]] = []
cursor = 0
while cursor < len(smi):
if smi[cursor] == "[":
closing = smi.find("]", cursor + 1)
if closing < 0:
return smi, set()
if smi[cursor : closing + 1] == "[*]":
dummy_tokens.append((cursor, closing + 1))
cursor = closing + 1
continue
if smi[cursor] == "*":
dummy_tokens.append((cursor, cursor + 1))
cursor += 1
if len(dummy_tokens) != len(star_indices):
if error_msg:
logger.warning("Unable to align dummy atoms for stereo preservation")
return smi, set()
used_isotopes = {atom.GetIsotope() for atom in probe.GetAtoms()}
temporary_isotopes: set[int] = set()
next_isotope = 1
protected_parts: List[str] = []
last_end = 0
for (start, end), atom_idx in zip(dummy_tokens, star_indices):
protected_parts.append(smi[last_end:start])
if atom_idx in substitutable_indices:
while next_isotope in used_isotopes:
next_isotope += 1
protected_parts.append(f"[{next_isotope}*]")
temporary_isotopes.add(next_isotope)
used_isotopes.add(next_isotope)
next_isotope += 1
else:
protected_parts.append(smi[start:end])
last_end = end
protected_parts.append(smi[last_end:])
return "".join(protected_parts), temporary_isotopes
@classmethod
def refactor(
cls,
caption: str,
error_msg: bool = False,
) -> Optional[TranslatedMolecule]:
"""Refactor E-SMILES and detect Markush/SRU output."""
if "<sep>" not in caption:
canonical_smi = cls.canonicalize_smiles(caption)
return TranslatedMolecule(
smi=canonical_smi,
groups="",
caption=caption,
esmi=cls.build_esmi(canonical_smi),
markush=False,
sru=False,
)
smi, trailing = caption.split(Tokens.separator, 1)
if trailing == "":
canonical_smi = cls.canonicalize_smiles(smi)
return TranslatedMolecule(
smi=canonical_smi,
groups="",
caption=caption,
esmi=cls.build_esmi(canonical_smi),
markush=False,
sru=False,
)
groups, ext = cls.parse_trailing(trailing)
# Give composite abbreviations one best-effort pass before the legacy
# atom-group loop. substitute_markush remaps retained atom records,
# so recursive refactor sees current atom indices.
composite_symbols = (
"SO2",
"CO2",
"CO",
"CH2",
"CF2",
"NH",
)
composite_descs = cls.parse_groups(groups)
atom_group_ids = [
int(desc.id)
for desc in composite_descs
if isinstance(desc.id, AtomIndex)
and not desc.is_circle
and not desc.is_dummy
]
if len(atom_group_ids) != len(set(atom_group_ids)):
# Consuming one label by atom ID would also remove its siblings.
# Preserve ambiguous atom annotations before either expansion pass.
return TranslatedMolecule(
smi=smi,
groups=groups,
caption=caption,
esmi=cls.build_esmi(smi, groups, ext),
markush=True,
sru=cls.has_symbolic_sru(caption),
)
composite_abbrevs = chem_utils.get_abbrev_smi()
def keeps_an_open_site(desc: GroupDesc) -> bool:
# The legacy loop attaches a table fragment at atom 0 and keeps
# its wildcard. The Markush pass consumes that wildcard instead.
src = composite_abbrevs.get((desc.symbol or "") + (desc.script or ""))
return src is not None and "*" in src
if any(
desc.symbol
and any(token in desc.symbol for token in composite_symbols)
for desc in composite_descs
) and not any(keeps_an_open_site(desc) for desc in composite_descs):
try:
expanded = cls.substitute_markush(
caption,
{},
error_msg=error_msg,
repeat_policy="best_effort",
)
except (ValueError, RuntimeError):
expanded = caption
if isinstance(expanded, str) and expanded != caption:
translated = cls.refactor(expanded, error_msg=error_msg)
if translated is not None:
return replace(translated, caption=caption)
is_sru = cls.has_symbolic_sru(caption)
try:
abbrev_map = chem_utils.get_abbrev_smi()
raw_mol = chem_utils.parse_smiles(smi)
if raw_mol is None:
if error_msg:
logger.warning(f"Invalid SMILES: {smi}")
return TranslatedMolecule(
smi=smi,
groups=groups,
caption=caption,
esmi=cls.build_esmi(smi, groups, ext),
markush=len(groups) > 0,
sru=is_sru,
)
repaired_groups = cls.repair_atom_group_indices(raw_mol, groups, error_msg=error_msg)
if repaired_groups != groups:
groups = repaired_groups
caption = cls.build_esmi(smi, groups, ext)
stereo_safe_smi, temporary_isotopes = cls._protect_stereo_for_dummy_substitution(
smi, groups, abbrev_map, raw_mol, error_msg=error_msg
)
mol = raw_mol if stereo_safe_smi == smi else chem_utils.parse_smiles(stereo_safe_smi)
if mol is None:
if error_msg:
logger.warning(f"Invalid SMILES: {smi}")
return TranslatedMolecule(
smi=smi,
groups=groups,
caption=caption,
esmi=cls.build_esmi(smi, groups, ext),
markush=len(groups) > 0,
sru=is_sru,
)
for atom in mol.GetAtoms():
atom.SetAtomMapNum(atom.GetIdx() + 1)
ring_info = mol.GetRingInfo().AtomRings()
mapped_smiles = Chem.MolToSmiles(mol, canonical=False, isomericSmiles=True)
mapped_caption = cls.build_esmi(mapped_smiles, groups, ext)
parsed = cls.parse_caption(
mapped_caption, return_mol=True, error_msg=error_msg
)
mol, groups, ext = parsed
mol = Chem.RWMol(mol)
mapnum2idx: Dict[int, int] = {}
for a in mol.GetAtoms():
mn = a.GetAtomMapNum()
if mn > 0:
mapnum2idx[mn] = a.GetIdx()
to_remove: List[int] = []
consumed_atom_groups: set[int] = set()
preserve_dummy_groups = False
has_markush_extensions = any(
token in groups
for token in (
Tokens.substruct_start,
Tokens.sgroup_start,
Tokens.virtual_start,
f"{Tokens.ring_start}{Tokens.virtual_start}",
)
)
# Axial configuration is definite stereochemistry. Preserve and
# remap it without classifying an otherwise resolved graph as Markush.
preserve_precompat_groups = (
has_markush_extensions or Tokens.axial_start in groups
)
is_markush = has_markush_extensions
for desc in cls.parse_groups(groups):
if isinstance(desc.id, AtomIndex):
atom_idx = mapnum2idx.get(int(desc.id) + 1)
# Index repair already ran above. If the source atom still
# cannot be found, drop only this stale record and keep the
# remainder of the caption usable.
if atom_idx is None:
continue
if desc.is_circle:
is_markush = True
continue
elif isinstance(desc.id, RingIndex):
if not desc.id.virtual and int(desc.id) >= len(ring_info):
continue
is_markush = True
continue
else:
is_markush = True
continue
atom = mol.GetAtomWithIdx(atom_idx)
if atom.GetSymbol() != "*" and not desc.is_dummy:
is_markush = True
continue
if desc.is_dummy:
preserve_dummy_groups = True
continue
if not desc.symbol:
continue
# Carbon chain repetition, e.g. (CH2)n.
if (desc.symbol == "CH2") or (desc.symbol == "CH" and desc.script == "2"):
if desc.multiple and not desc.multiple.isdigit():
is_markush = True
continue
if (
desc.multiple
and int(desc.multiple) > _MAX_AUTOMATIC_REPEAT_COUNT
):
is_markush = True
continue
if desc.multiple and not _is_safe_carbon_chain_repeat_target(atom):
is_markush = True
continue
is_markush = chem_utils.carbon_chain_repetition_process(
mol, atom_idx, desc, is_markush, error_msg=False
)
consumed_atom_groups.add(int(desc.id))
continue
# A multiplicity suffix belongs to the complete group. Do not
# silently consume it as one ordinary abbreviation when the
# topology cannot be expanded to one unique structure.
if desc.multiple:
is_markush = True
continue
# Build lookup key, e.g. NO + 2 -> NO2.
lookup_symbol = desc.symbol
if desc.script:
lookup_symbol += desc.script
if lookup_symbol == "CN":
src_smi = "C(#N)"
else:
src_smi = abbrev_map.get(lookup_symbol)
if src_smi is not None:
src_mol_check = chem_utils.parse_smiles(src_smi)
if src_mol_check and src_mol_check.GetNumAtoms() == 1:
chem_utils.alter_atom(atom, smiles=src_smi)
consumed_atom_groups.add(int(desc.id))
continue
if atom.GetDegree() != 1:
if error_msg:
logger.warning(
f"Group `{lookup_symbol}` cannot be attached to atom "
f"{atom_idx} with degree {atom.GetDegree()}"
)
is_markush = True
continue
if any(not chem_utils.is_single_bond(b) for b in atom.GetBonds()):
if error_msg:
logger.warning(
f"Group `{lookup_symbol}` must link to single bond"
)
is_markush = True
continue
try:
chem_utils.merge_group(tgt_mol=mol, src=src_smi, attach_idx=atom_idx)
to_remove.append(atom_idx)
consumed_atom_groups.add(int(desc.id))
except Exception as e:
if error_msg:
logger.warning(f"Merge group failed: {e}")
is_markush = True
continue
if lookup_symbol in PERIODIC_TABLE and not desc.multiple:
try:
chem_utils.alter_atom(atom, smiles=None, element=lookup_symbol)
consumed_atom_groups.add(int(desc.id))
except Exception as e:
if error_msg:
logger.warning(
f"Failed to mutate atom {atom_idx} to {lookup_symbol}: {e}"
)
is_markush = True
continue
is_markush = True
for i in sorted(to_remove, reverse=True):
mol.RemoveAtom(i)
cleared_temporary_isotopes = False
for atom in mol.GetAtoms():
if atom.GetSymbol() == "*" and atom.GetIsotope() in temporary_isotopes:
atom.SetIsotope(0)
cleared_temporary_isotopes = True
try:
Chem.SanitizeMol(mol)
if cleared_temporary_isotopes:
Chem.AssignStereochemistry(mol, cleanIt=True, force=True)
except Exception as e:
if error_msg:
logger.error(f"Sanitize failed: {e}")
return TranslatedMolecule(
smi=smi,
groups=groups,
caption=caption,
esmi=cls.build_esmi(smi, groups, ext),
markush=True,
sru=is_sru,
)
mol = chem_utils.parse_smiles(
Chem.MolToSmiles(mol, canonical=True, isomericSmiles=True)
)
if is_markush or preserve_dummy_groups or preserve_precompat_groups:
remaining_groups = cls.remove_atom_groups(groups, consumed_atom_groups)
new_groups = chem_utils.remap_groups(mol, remaining_groups, ring_info)
else:
new_groups = ""
for atom in mol.GetAtoms():
atom.SetAtomMapNum(0)
new_smi = Chem.MolToSmiles(mol, canonical=True, isomericSmiles=True)
new_esmi = cls.build_esmi(new_smi, new_groups, ext)
return TranslatedMolecule(
smi=new_smi,
groups=new_groups,
caption=caption,
esmi=new_esmi,
markush=is_markush,
sru=is_sru,
)
except Exception as e:
if error_msg:
logger.error(f"Error while refactoring SMILES: {smi}")
logger.error(repr(e))
return TranslatedMolecule(
smi=smi,
groups=groups,
caption=caption,
esmi=cls.build_esmi(smi, groups, ext),
markush=len(groups) > 0,
sru=is_sru,
)
@classmethod
def esmiles_to_cxsmiles(cls, caption: str, error_msg: bool = False) -> str:
"""Convert E-SMILES to CXSMILES using the refactored E-SMILES output."""
raw_caption = str(caption).strip()
raw_groups = ""
if Tokens.separator in raw_caption:
trailing = raw_caption.split(Tokens.separator, 1)[1]
raw_groups, _ = cls.parse_trailing(trailing)
translated = cls.refactor(raw_caption, error_msg=error_msg)
esmi = translated.esmi if translated is not None else raw_caption
is_sru = translated.sru if translated is not None else False
try:
from .cxsmiles import _convert_refactored_esmi_to_cxsmiles
except ImportError: # Support running from package directory as working directory.
from cxsmiles import _convert_refactored_esmi_to_cxsmiles
return _convert_refactored_esmi_to_cxsmiles(
esmi,
source_groups=(translated.groups if translated is not None else raw_groups),
sru=is_sru,
)
@classmethod
def substitute_markush(
cls,
caption: str,
definitions: Mapping[str, Union[int, str, Sequence[str]]],
*,
max_outputs: int = 1024,
error_msg: bool = False,
repeat_policy: Literal["preserve", "best_effort", "strict"] = "best_effort",
terminal_policy: Literal["preserve", "hydrogen"] = "preserve",
) -> Union[str, List[str]]:
"""Substitute Markush labels and optionally physicalize repeat counts.
The import is intentionally local because ``markush`` uses the parser
types defined in this module. See :func:`molparser.utils.substitute_markush`
for repeat and terminal policy semantics.
"""
try:
from .markush import substitute_markush as _substitute_markush
except ImportError: # Support running from package directory.
from markush import substitute_markush as _substitute_markush
return _substitute_markush(
caption,
definitions,
max_outputs=max_outputs,
error_msg=error_msg,
repeat_policy=repeat_policy,
terminal_policy=terminal_policy,
)
__all__ = [
"AtomIndex",
"GroupDesc",
"Index",
"LONG_CHAR_ELEMENTS",
"Patterns",
"PERIODIC_TABLE",
"RingIndex",
"TextType",
"Tokens",
"TranslatedMolecule",
"Translator",
]