| import re |
| import torch |
| from transformers import AutoTokenizer, EsmForMaskedLM |
|
|
| |
| 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() |
|
|
| |
| 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] |
|
|
| |
| |
| 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 |
|
|
| results = [] |
| for i, seq_tokens in enumerate(all_sequence_tokens): |
| n_tokens = len(seq_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} |
|
|