File size: 2,532 Bytes
a70e95c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 | import re
import torch
from transformers import AutoTokenizer, EsmForMaskedLM
# The 33 standard ESM2 tokens in vocabulary-ID order.
VOCAB_TOKENS = [
"<cls>", "<pad>", "<eos>", "<unk>",
"L", "A", "G", "V", "S", "E", "R", "T", "I", "D",
"P", "K", "Q", "N", "F", "Y", "M", "H", "W", "C",
"X", "B", "U", "Z", "O", ".", "-",
"<null_1>", "<mask>",
]
BASE_MODEL = "facebook/esm2_t36_3B_UR50D"
class EndpointHandler:
def __init__(self, path: str):
self.tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL)
self.model = EsmForMaskedLM.from_pretrained(
BASE_MODEL, torch_dtype=torch.float16
)
self.model.to("cuda")
self.model.eval()
# Pre-compute and validate vocab_tokens from the actual tokenizer.
self.vocab_tokens = [
self.tokenizer.convert_ids_to_tokens(i)
for i in range(self.tokenizer.vocab_size)
]
def __call__(self, data: dict) -> dict:
items = data.get("items", data.get("inputs", []))
sequences = [item["sequence"] for item in items]
# Build sequence_tokens for each sequence: split each character as its
# own token, but keep <mask> as a single token.
all_sequence_tokens = []
for seq in sequences:
tokens = re.split(r"(<mask>)", seq)
seq_tokens = []
for part in tokens:
if part == "<mask>":
seq_tokens.append("<mask>")
else:
seq_tokens.extend(list(part))
all_sequence_tokens.append(seq_tokens)
encoded = self.tokenizer(
sequences,
return_tensors="pt",
padding=True,
truncation=True,
).to("cuda")
with torch.no_grad():
output = self.model(**encoded)
logits = output.logits # (batch, seq_len_with_special, vocab_size)
results = []
for i, seq_tokens in enumerate(all_sequence_tokens):
n_tokens = len(seq_tokens)
# Slice out CLS (position 0) and EOS/padding at the end.
# Positions 1..n_tokens correspond to the actual sequence tokens.
seq_logits = logits[i, 1 : n_tokens + 1, :].float().cpu().tolist()
results.append(
{
"logits": seq_logits,
"sequence_tokens": seq_tokens,
"vocab_tokens": self.vocab_tokens,
}
)
return {"results": results}
|