pranamanam's picture Jingjie00's picture
Upload Staplebridge files (#1)
bb6d2aa
Raw History Blame Contribute Delete
7.29 kB
"""Multi-character monomer tokenizer for the hydrocarbon branch.
The lactam path builds states with ``StapleState.from_sequence``, which is
``list(seq)`` - one character per residue. That is correct for K/D/E lactam and
is left untouched. Hydrocarbon anchors are two characters (``S5``, ``R8``), so
this module provides a *separate* tokenizing constructor. Lactam never calls it.
Model-vocabulary projection
---------------------------
``staplebridge.data.vocab.ALL_TOKENS`` has 24 entries and
:class:`~staplebridge.models.embeddings.TokenMLPEncoder` sizes its embedding as
``nn.Embedding(len(TOKEN_TO_ID), ...)``. Appending ``S5``/``R8`` to that list
would change the embedding matrix shape and break loading of every existing
checkpoint. So this module does **not** touch the vocab. Instead it projects
hydrocarbon monomers onto the ncAA tokens the vocab already carries:
S3, S5, S8 -> "X" (X is already the StaPep alias for S5)
R3, R5, R8 -> "B"
Aib -> "X"
Nle -> "B"
The projection is only applied where a tensor of token ids is needed. The
authoritative state keeps the true monomer tokens, so anchor typing, catalog
matching and the endpoint prior all see ``S5``/``R8`` exactly.
"""
from __future__ import annotations
from typing import Final
from staplebridge.chemistry.state import StapleState
from staplebridge.data.vocab import TOKEN_TO_ID
NATURAL_AA: Final[frozenset[str]] = frozenset("ACDEFGHIKLMNPQRSTVWY")
#: Hydrocarbon anchor monomers recognised by the tokenizer. Longest-first so a
#: greedy scan consumes ``S5`` before it can mistake ``S`` for serine.
HYDROCARBON_ANCHOR_TOKENS: Final[tuple[str, ...]] = (
"S3",
"S5",
"S8",
"R3",
"R5",
"R8",
)
#: Non-anchor non-natural monomers the tokenizer accepts as whole segments.
OTHER_MONOMER_TOKENS: Final[tuple[str, ...]] = ("Aib", "AIB", "Nle", "NLE")
#: Terminal modifications: no residue index, no contribution to length.
N_TERMINAL_MODS: Final[frozenset[str]] = frozenset({"AC"})
C_TERMINAL_MODS: Final[frozenset[str]] = frozenset({"NH2"})
#: Projection onto tokens the existing model vocabulary already contains. See
#: the module docstring for why the vocab itself is not extended.
_MODEL_TOKEN_PROJECTION: Final[dict[str, str]] = {
"S3": "X",
"S5": "X",
"S8": "X",
"R3": "B",
"R5": "B",
"R8": "B",
"AIB": "X",
"NLE": "B",
}
_MULTI_CHAR: Final[tuple[str, ...]] = tuple(
sorted(
{t.upper() for t in HYDROCARBON_ANCHOR_TOKENS + OTHER_MONOMER_TOKENS},
key=len,
reverse=True,
)
)
class HydrocarbonTokenizationError(ValueError):
"""Raised when a sequence cannot be fully consumed into known monomers."""
def is_anchor_token(token: str) -> bool:
"""True when ``token`` is a hydrocarbon staple anchor monomer."""
return token.upper() in {t.upper() for t in HYDROCARBON_ANCHOR_TOKENS}
def normalize_monomer(token: str) -> str:
"""Canonicalise one monomer token (``s5`` -> ``S5``, ``Aib`` -> ``AIB``)."""
upper = token.upper()
if upper in {t.upper() for t in HYDROCARBON_ANCHOR_TOKENS}:
return upper
if upper in {"AIB", "NLE"}:
return upper
return upper
def tokenize_sequence(sequence: str) -> list[str]:
"""Tokenize a hydrocarbon-style sequence into monomer tokens.
Supports dash-delimited segments, undelimited runs and mixtures, plus
``Ac-``/``-NH2`` terminal modifications (which are dropped from the residue
list because they carry no residue index).
Args:
sequence: e.g. ``"TSFR8EYWALLS5"``, ``"Ac-ISF-R8-ELLDYY-S5-ESGS"``.
Returns:
One canonical monomer token per residue.
Raises:
HydrocarbonTokenizationError: if any part cannot be consumed. Nothing is
guessed.
"""
if sequence is None or not str(sequence).strip():
raise HydrocarbonTokenizationError("empty sequence")
text = str(sequence).strip().strip("-")
segments = [s for s in text.split("-") if s]
tokens: list[str] = []
for position, segment in enumerate(segments):
upper = segment.upper()
if upper in N_TERMINAL_MODS and position == 0:
continue
if upper in C_TERMINAL_MODS and position == len(segments) - 1:
continue
if upper in _MULTI_CHAR:
tokens.append(normalize_monomer(upper))
continue
tokens.extend(_tokenize_run(segment))
if not tokens:
raise HydrocarbonTokenizationError(
f"sequence {sequence!r} contained no residues"
)
return tokens
def _tokenize_run(run: str) -> list[str]:
"""Longest-match scan of one undelimited run."""
upper = run.upper()
tokens: list[str] = []
index = 0
while index < len(upper):
for candidate in _MULTI_CHAR:
if upper.startswith(candidate, index):
tokens.append(normalize_monomer(candidate))
index += len(candidate)
break
else:
char = upper[index]
if char in NATURAL_AA:
tokens.append(char)
index += 1
else:
raise HydrocarbonTokenizationError(
f"unrecognized character {char!r} at offset {index} of {run!r}"
)
return tokens
def state_from_sequence(sequence: str, **kwargs: object) -> StapleState:
"""Build a :class:`StapleState` whose tokens are hydrocarbon monomers.
This is the hydrocarbon counterpart of ``StapleState.from_sequence``. It is
a separate function precisely so the lactam constructor keeps its exact
per-character behaviour.
"""
return StapleState(sequence_tokens=tokenize_sequence(sequence), **kwargs)
def anchor_positions(tokens: list[str]) -> list[int]:
"""Residue indices of the hydrocarbon anchors, N-to-C."""
return [i for i, t in enumerate(tokens) if is_anchor_token(t)]
def to_model_tokens(tokens: list[str]) -> list[str]:
"""Project monomer tokens onto tokens present in the existing model vocab.
Multi-character hydrocarbon monomers are mapped onto the ``X``/``B`` ncAA
tokens that ``staplebridge.data.vocab`` already defines, so the embedding
matrix keeps its original size and old checkpoints still load. Natural
residues pass through unchanged.
"""
projected: list[str] = []
for token in tokens:
upper = token.upper()
if upper in _MODEL_TOKEN_PROJECTION:
projected.append(_MODEL_TOKEN_PROJECTION[upper])
elif token in TOKEN_TO_ID:
projected.append(token)
elif upper in TOKEN_TO_ID:
projected.append(upper)
else:
projected.append("<unk>")
return projected
def to_display_sequence(tokens: list[str]) -> str:
"""Human-readable sequence string, dash-separating multi-character monomers.
``"".join`` would render ``[..., "S5", ...]`` ambiguously against a real
``S`` followed by a literal ``5``, so multi-character monomers are set off
with dashes.
"""
parts: list[str] = []
for token in tokens:
if len(token) > 1:
parts.append(f"-{token}-")
else:
parts.append(token)
return "".join(parts).replace("--", "-").strip("-")