Download classifier_code/nonfouling_wt.py from ChatterjeeLab/moPPIt: direct link, hf CLI and curl.
- Browser
- Download file 2.79 kB
-
https://huggingface.co/ChatterjeeLab/moPPIt/resolve/main/classifier_code/nonfouling_wt.py
- Command line
-
hf download hf://ChatterjeeLab/moPPIt/classifier_code/nonfouling_wt.py
-
curl -L -o nonfouling_wt.py https://huggingface.co/ChatterjeeLab/moPPIt/resolve/main/classifier_code/nonfouling_wt.py
2.79 kB
| import numpy as np | |
| import xgboost as xgb | |
| import torch | |
| from transformers import EsmModel, AutoTokenizer | |
| import torch.nn as nn | |
| import pdb | |
| # ======================== MLP ========================================= | |
| # Still need mean pooling along lengths | |
| class MaskedMeanPool(nn.Module): | |
| def forward(self, X, M): # X: (B,L,H), M: (B,L) | |
| Mf = M.unsqueeze(-1).float() | |
| denom = Mf.sum(dim=1).clamp(min=1.0) | |
| return (X * Mf).sum(dim=1) / denom # (B,H) | |
| class MLPClassifier(nn.Module): | |
| def __init__(self, in_dim, hidden=512, dropout=0.1): | |
| super().__init__() | |
| self.pool = MaskedMeanPool() | |
| self.net = nn.Sequential( | |
| nn.Linear(in_dim, hidden), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| nn.Linear(hidden, 1), | |
| ) | |
| def forward(self, X, M): | |
| z = self.pool(X, M) | |
| return self.net(z).squeeze(-1) # logits | |
| # ======================== MLP ========================================= | |
| class NonfoulingModel: | |
| def __init__(self, device): | |
| ckpt = torch.load('../classifier_ckpt/wt_nonfouling.pt', weights_only=False, map_location=device) | |
| best_params = ckpt["best_params"] | |
| self.predictor = MLPClassifier(in_dim=1280, hidden=int(best_params["hidden"]), dropout=float(best_params.get("dropout", 0.1))) | |
| self.predictor.load_state_dict(ckpt["state_dict"]) | |
| self.predictor = self.predictor.to(device) | |
| self.predictor.eval() | |
| self.model = EsmModel.from_pretrained("facebook/esm2_t33_650M_UR50D").to(device) | |
| # self.model.eval() | |
| self.device = device | |
| def generate_embeddings(self, input_ids, attention_mask): | |
| """Generate ESM embeddings for protein sequences""" | |
| with torch.no_grad(): | |
| embeddings = self.model(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state | |
| return embeddings | |
| def get_scores(self, input_ids, attention_mask): | |
| features = self.generate_embeddings(input_ids, attention_mask) | |
| keep = (input_ids != 0) & (input_ids != 1) & (input_ids != 2) | |
| attention_mask[keep==False] = 0 | |
| scores = self.predictor(features, attention_mask) | |
| return scores.detach().cpu().numpy() | |
| def __call__(self, input_ids, attention_mask): | |
| scores = self.get_scores(input_ids, attention_mask) | |
| return 1.0 / (1.0 + np.exp(-scores)) | |
| def unittest(): | |
| device = 'cuda:0' | |
| nf = NonfoulingModel(device=device) | |
| seq = ["HAIYPRH", "HAEGTFTSDVSSYLEGQAAKEFIAWLVKGR"] | |
| tokenizer = AutoTokenizer.from_pretrained('facebook/esm2_t33_650M_UR50D') | |
| seq_tokens = tokenizer(seq, padding=True, return_tensors='pt').to(device) | |
| scores = nf(**seq_tokens) | |
| print(scores) | |
| if __name__ == '__main__': | |
| unittest() |