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}