""" Inference-mode mirror of training/model.py's TransformerDMSRegressor. Kept as a separate copy (not an import) because this file is the ONLY directory git-synced to the HF Space — it must be fully self-contained. """ import torch import torch.nn as nn from transformers import EsmForMaskedLM, EsmTokenizer DEFAULT_MODEL_NAME = "facebook/esm2_t12_35M_UR50D" class TransformerDMSRegressor(nn.Module): def __init__(self, model_name: str = DEFAULT_MODEL_NAME): super().__init__() self.tokenizer = EsmTokenizer.from_pretrained(model_name) self.backbone = EsmForMaskedLM.from_pretrained(model_name) for param in self.backbone.parameters(): param.requires_grad = False for param in self.backbone.esm.encoder.layer[-2:].parameters(): param.requires_grad = True self.regression_head = nn.Linear(1, 1) @torch.no_grad() def score(self, sequence: str, position: int, ref_aa: str, alt_aa: str, device: torch.device): max_len = 1022 # ESM tokenizer budget minus special tokens start = max(0, position - 1 - max_len // 2) end = min(len(sequence), start + max_len) start = max(0, end - max_len) window = sequence[start:end] local_pos = position - start # 1-based within window, aligns with offset encoding = self.tokenizer(window, return_tensors="pt", truncation=True, max_length=max_len + 2) input_ids = encoding["input_ids"].to(device) attention_mask = encoding["attention_mask"].to(device) seq_len = input_ids.size(1) mutation_idx = torch.tensor([min(local_pos, seq_len - 1)], device=device) ref_id = torch.tensor([self.tokenizer.convert_tokens_to_ids(ref_aa)], device=device) alt_id = torch.tensor([self.tokenizer.convert_tokens_to_ids(alt_aa)], device=device) outputs = self.backbone(input_ids=input_ids, attention_mask=attention_mask) logits = outputs.logits logits_at_mut = logits[0, mutation_idx[0], :] log_probs = torch.log_softmax(logits_at_mut, dim=-1) llr = (log_probs[alt_id[0]] - log_probs[ref_id[0]]).item() fitness = self.regression_head(torch.tensor([[llr]], device=device)).item() return llr, fitness