import re import torch from transformers import AutoTokenizer, EsmForMaskedLM # The 33 standard ESM2 tokens in vocabulary-ID order. VOCAB_TOKENS = [ "", "", "", "", "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", ".", "-", "", "", ] 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 as a single token. all_sequence_tokens = [] for seq in sequences: tokens = re.split(r"()", seq) seq_tokens = [] for part in tokens: if part == "": seq_tokens.append("") 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}