Views
No views yet

1import torch
2import torch.nn as nn
3from transformers import AutoModel,AutoModelForMaskedLM,AutoTokenizer
4import os
5import torch.nn.functional as F
6
7class PLMinteract(nn.Module):
8 def __init__(self,model_name,num_labels,embedding_size):
9 super(PLMinteract,self).__init__()
10 self.esm_mask = AutoModelForMaskedLM.from_pretrained(model_name)
11 self.embedding_size=embedding_size
12 self.classifier = nn.Linear(embedding_size,1) # embedding_size
13 self.num_labels=num_labels
14
15 def forward_test(self,features):
16 embedding_output = self.esm_mask.base_model(**features, return_dict=True)
17 embedding=embedding_output.last_hidden_state[:,0,:] #cls token
18 embedding = F.relu(embedding)
19 logits = self.classifier(embedding)
20 logits=logits.view(-1)
21 probability = torch.sigmoid(logits)
22 return probability
23
24# folder_huggingface_download : the download model from huggingface, such as "danliu1226/PLM-interact-650M-humanV11"
25# model_name: the ESM2 model that PLM-interact trained
26# embedding_size: the embedding size of ESM2 model
27
28folder_huggingface_download='download_huggingface_folder/'
29model_name= 'facebook/esm2_t33_650M_UR50D'
30embedding_size =1280
31
32protein1 ="EGCVSNLMVCNLAYSGKLEELKESILADKSLATRTDQDSRTALHWACSAGHTEIVEFLLQLGVPVNDKDDAGWSPLHIAASAGRDEIVKALLGKGAQVNAVNQNGCTPLHYAASKNRHEIAVMLLEGGANPDAKDHYEATAMHRAAAKGNLKMIHILLYYKASTNIQDTEGNTPLHLACDEERVEEAKLLVSQGASIYIENKEEKTPLQVAKGGLGLILKRMVEG"
33
34protein2= "MGQSQSGGHGPGGGKKDDKDKKKKYEPPVPTRVGKKKKKTKGPDAASKLPLVTPHTQCRLKLLKLERIKDYLLMEEEFIRNQEQMKPLEEKQEEERSKVDDLRGTPMSVGTLEEIIDDNHAIVSTSVGSEHYVSILSFVDKDLLEPGCSVLLNHKVHAVIGVLMDDTDPLVTVMKVEKAPQETYADIGGLDNQIQEIKESVELPLTHPEYYEEMGIKPPKGVILYGPPGTGKTLLAKAVANQTSATFLRVVGSELIQKYLGDGPKLVRELFRVAEEHAPSIVFIDEIDAIGTKRYDSNSGGEREIQRTMLELLNQLDGFDSRGDVKVIMATNRIETLDPALIRPGRIDRKIEFPLPDEKTKKRIFQIHTSRMTLADDVTLDDLIMAKDDLSGADIKAICTEAGLMALRERRMKVTNEDFKKSKENVLYKKQEGTPEGLYL"
35
36DEVICE = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')
37tokenizer = AutoTokenizer.from_pretrained(model_name)
38PLMinter= PLMinteract(model_name, 1, embedding_size)
39load_model = torch.load(f"{folder_huggingface_download}pytorch_model.bin")
40PLMinter.load_state_dict(load_model)
41
42texts=[protein1, protein2]
43tokenized = tokenizer(*texts, padding=True, truncation='longest_first', return_tensors="pt", max_length=1603)
44tokenized = tokenized.to(DEVICE)
45
46PLMinter.eval()
47PLMinter.to(DEVICE)
48with torch.no_grad():
49 probability = PLMinter.forward_test(tokenized)
50 print(probability.item())