ammonb commited on
Commit
a70e95c
·
verified ·
1 Parent(s): 6ce6bfe

Upload handler.py with huggingface_hub

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