Views
No views yet
facebook/esm2_t33_650M_UR50D using the masked language modeling objective.
The model was finetuned for two epochs on concatenated pairs of interacting proteins, clustered using persistent homology
landscapes as explained in this post. The dataset consists of 10,000
protein pairs, which can be found here.
This is a very new method for clustering protein-protein complexes.1import numpy as np
2from transformers import AutoTokenizer, EsmForMaskedLM
3import torch
4
5# Load the base model and tokenizer
6tokenizer = AutoTokenizer.from_pretrained("facebook/esm2_t33_650M_UR50D")
7model = EsmForMaskedLM.from_pretrained("AmelieSchreiber/esm_mlmppi_ph50")
8
9# Ensure the model is in evaluation mode
10model.eval()
11
12# Define the protein of interest and its potential binders
13protein_of_interest = "MLTEVMEVWHGLVIAVVSLFLQACFLTAINYLLSRHMAHKSEQILKAASLQVPRPSPGHHHPPAVKEMKETQTERDIPMSDSLYRHDSDTPSDSLDSSCSSPPACQATEDVDYTQVVFSDPGELKNDSPLDYENIKEITDYVNVNPERHKPSFWYFVNPALSEPAEYDQVAM"
14potential_binders = [
15 # Known to interact
16 "MASPGSGFWSFGSEDGSGDSENPGTARAWCQVAQKFTGGIGNKLCALLYGDAEKPAESGGSQPPRAAARKAACACDQKPCSCSKVDVNYAFLHATDLLPACDGERPTLAFLQDVMNILLQYVVKSFDRSTKVIDFHYPNELLQEYNWELADQPQNLEEILMHCQTTLKYAIKTGHPRYFNQLSTGLDMVGLAADWLTSTANTNMFTYEIAPVFVLLEYVTLKKMREIIGWPGGSGDGIFSPGGAISNMYAMMIARFKMFPEVKEKGMAALPRLIAFTSEHSHFSLKKGAAALGIGTDSVILIKCDERGKMIPSDLERRILEAKQKGFVPFLVSATAGTTVYGAFDPLLAVADICKKYKIWMHVDAAWGGGLLMSRKHKWKLSGVERANSVTWNPHKMMGVPLQCSALLVREEGLMQNCNQMHASYLFQQDKHYDLSYDTGDKALQCGRHVDVFKLWLMWRAKGTTGFEAHVDKCLELAEYLYNIIKNREGYEMVFDGKPQHTNVCFWYIPPSLRTLEDNEERMSRLSKVAPVIKARMMEYGTTMVSYQPLGDKVNFFRMVISNPAATHQDIDFLIEEIERLGQDL",
17 "MAAGVAGWGVEAEEFEDAPDVEPLEPTLSNIIEQRSLKWIFVGGKGGVGKTTCSCSLAVQLSKGRESVLIISTDPAHNISDAFDQKFSKVPTKVKGYDNLFAMEIDPSLGVAELPDEFFEEDNMLSMGKKMMQEAMSAFPGIDEAMSYAEVMRLVKGMNFSVVVFDTAPTGHTLRLLNFPTIVERGLGRLMQIKNQISPFISQMCNMLGLGDMNADQLASKLEETLPVIRSVSEQFKDPEQTTFICVCIAEFLSLYETERLIQELAKCKIDTHNIIVNQLVFPDPEKPCKMCEARHKIQAKYLDQMEDLYEDFHIVKLPLLPHEVRGADKVNTFSALLLEPYKPPSAQ",
18 "EKTGLSIRGAQEEDPPDPQLMRLDNMLLAEGVSGPEKGGGSAAAAAAAAASGGSSDNSIEHSDYRAKLTQIRQIYHTELEKYEQACNEFTTHVMNLLREQSRTRPISPKEIERMVGIIHRKFSSIQMQLKQSTCEAVMILRSRFLDARRKRRNFSKQATEILNEYFYSHLSNPYPSEEAKEELAKKCSITVSQSLVKDPKERGSKGSDIQPTSVVSNWFGNKRIRYKKNIGKFQEEANLYAAKTAVTAAHAVAAAVQNNQTNSPTTPNSGSSGSFNLPNSGDMFMNMQSLNGDSYQGSQVGANVQSQVDTLRHVINQTGGYSDGLGGNSLYSPHNLNANGGWQDATTPSSVTSPTEGPGSVHSDTSN",
19 # Not known to interact
20 "MRQRLLPSVTSLLLVALLFPGSSQARHVNHSATEALGELRERAPGQGTNGFQLLRHAVKRDLLPPRTPPYQVHISHREARGPSFRICVDFLGPRWARGCSTGN",
21 "MSGIALSRLAQERKAWRKDHPFGFVAVPTKNPDGTMNLMNWECAIPGKKGTPWEGGLFKLRMLFKDDYPSSPPKCKFEPPLFHPNVYPSGTVCLSILEEDKDWRPAITIKQILLGIQELLNEPNIQDPAQAEAYTIYCQNRVEYEKRVRAQAKKFAPS"
22] # Add potential binding sequences here
23
24def compute_mlm_loss(protein, binder, iterations=5):
25 total_loss = 0.0
26
27 for _ in range(iterations):
28 # Concatenate protein sequences with a separator
29 concatenated_sequence = protein + binder
30
31 # Mask a subset of amino acids in the concatenated sequence (excluding the separator)
32 tokens = list(concatenated_sequence)
33 mask_rate = 0.35 # For instance, masking 35% of the sequence
34 num_mask = int(len(tokens) * mask_rate)
35
36 # Exclude the separator from potential mask indices
37 available_indices = [i for i, token in enumerate(tokens) if token != ":"]
38 probs = torch.ones(len(available_indices))
39 mask_indices = torch.multinomial(probs, num_mask, replacement=False)
40
41 for idx in mask_indices:
42 tokens[available_indices[idx]] = tokenizer.mask_token
43
44 masked_sequence = "".join(tokens)
45 inputs = tokenizer(masked_sequence, return_tensors="pt", truncation=True, max_length=1024, padding='max_length', add_special_tokens=False)
46
47 # Compute the MLM loss
48 with torch.no_grad():
49 outputs = model(**inputs, labels=inputs["input_ids"])
50 loss = outputs.loss
51
52 total_loss += loss.item()
53
54 # Return the average loss
55 return total_loss / iterations
56
57# Compute MLM loss for each potential binder
58mlm_losses = {}
59for binder in potential_binders:
60 loss = compute_mlm_loss(protein_of_interest, binder)
61 mlm_losses[binder] = loss
62
63# Rank binders based on MLM loss
64ranked_binders = sorted(mlm_losses, key=mlm_losses.get)
65
66print("Ranking of Potential Binders:")
67for idx, binder in enumerate(ranked_binders, 1):
68 print(f"{idx}. {binder} - MLM Loss: {mlm_losses[binder]}")1import networkx as nx
2import numpy as np
3import torch
4from transformers import AutoTokenizer, AutoModelForMaskedLM, EsmForMaskedLM
5import plotly.graph_objects as go
6from ipywidgets import interact
7from ipywidgets import widgets
8
9# Check if CUDA is available and set the default device accordingly
10device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
11
12# Load the pretrained (or fine-tuned) ESM-2 model and tokenizer
13tokenizer = AutoTokenizer.from_pretrained("facebook/esm2_t33_650M_UR50D")
14model = EsmForMaskedLM.from_pretrained("AmelieSchreiber/esm_mlmppi_ph50")
15
16# Send the model to the device (GPU or CPU)
17model.to(device)
18
19# Ensure the model is in evaluation mode
20model.eval()
21
22# Define Protein Sequences (Replace with your list)
23all_proteins = [
24 "MFLSILVALCLWLHLALGVRGAPCEAVRIPMCRHMPWNITRMPNHLHHSTQENAILAIEQYEELVDVNCSAVLRFFLCAMYAPICTLEFLHDPIKPCKSVCQRARDDCEPLMKMYNHSWPESLACDELPVYDRGVCISPEAIVTDLPEDVKWIDITPDMMVQERPLDVDCKRLSPDRCKCKKVKPTLATYLSKNYSYVIHAKIKAVQRSGCNEVTTVVDVKEIFKSSSPIPRTQVPLITNSSCQCPHILPHQDVLIMCYEWRSRMMLLENCLVEKWRDQLSKRSIQWEERLQEQRRTVQDKKKTAGRTSRSNPPKPKGKPPAPKPASPKKNIKTRSAQKRTNPKRV",
25 "MDAVEPGGRGWASMLACRLWKAISRALFAEFLATGLYVFFGVGSVMRWPTALPSVLQIAITFNLVTAMAVQVTWKASGAHANPAVTLAFLVGSHISLPRAVAYVAAQLVGATVGAALLYGVMPGDIRETLGINVVRNSVSTGQAVAVELLLTLQLVLCVFASTDSRQTSGSPATMIGISVALGHLIGIHFTGCSMNPARSFGPAIIIGKFTVHWVFWVGPLMGALLASLIYNFVLFPDTKTLAQRLAILTGTVEVGTGAGAGAEPLKKESQPGSGAVEMESV",
26 "MKFLLDILLLLPLLIVCSLESFVKLFIPKRRKSVTGEIVLITGAGHGIGRLTAYEFAKLKSKLVLWDINKHGLEETAAKCKGLGAKVHTFVVDCSNREDIYSSAKKVKAEIGDVSILVNNAGVVYTSDLFATQDPQIEKTFEVNVLAHFWTTKAFLPAMTKNNHGHIVTVASAAGHVSVPFLLAYCSSKFAAVGFHKTLTDELAALQITGVKTTCLCPNFVNTGFIKNPSTSLGPTLEPEEVVNRLMHGILTEQKMIFIPSSIAFLTTLERILPERFLAVLKRKISVKFDAVIGYKMKAQ",
27
28 "MAAAVPRRPTQQGTVTFEDVAVNFSQEEWCLLSEAQRCLYRDVMLENLALISSLGCWCGSKDEEAPCKQRISVQRESQSRTPRAGVSPKKAHPCEMCGLILEDVFHFADHQETHHKQKLNRSGACGKNLDDTAYLHQHQKQHIGEKFYRKSVREASFVKKRKLRVSQEPFVFREFGKDVLPSSGLCQEEAAVEKTDSETMHGPPFQEGKTNYSCGKRTKAFSTKHSVIPHQKLFTRDGCYVCSDCGKSFSRYVSFSNHQRDHTAKGPYDCGECGKSYSRKSSLIQHQRVHTGQTAYPCEECGKSFSQKGSLISHQLVHTGEGPYECRECGKSFGQKGNLIQHQQGHTGERAYHCGECGKSFRQKFCFINHQRVHTGERPYKCGECGKSFGQKGNLVHHQRGHTGERPYECKECGKSFRYRSHLTEHQRLHTGERPYNCRECGKLFNRKYHLLVHERVHTGERPYACEVCGKLFGNKHSVTIHQRIHTGERPYECSECGKSFLSSSALHVHKRVHSGQKPYKCSECGKSFSECSSLIKHRRIHTGERPYECTKCGKTFQRSSTLLHHQSSHRRKAL",
29 "MGQPWAAGSTDGAPAQLPLVLTALWAAAVGLELAYVLVLGPGPPPLGPLARALQLALAAFQLLNLLGNVGLFLRSDPSIRGVMLAGRGLGQGWAYCYQCQSQVPPRSGHCSACRVCILRRDHHCRLLGRCVGFGNYRPFLCLLLHAAGVLLHVSVLLGPALSALLRAHTPLHMAALLLLPWLMLLTGRVSLAQFALAFVTDTCVAGALLCGAGLLFHGMLLLRGQTTWEWARGQHSYDLGPCHNLQAALGPRWALVWLWPFLASPLPGDGITFQTTADVGHTAS",
30 "MGLRIHFVVDPHGWCCMGLIVFVWLYNIVLIPKIVLFPHYEEGHIPGILIIIFYGISIFCLVALVRASITDPGRLPENPKIPHGEREFWELCNKCNLMRPKRSHHCSRCGHCVRRMDHHCPWINNCVGEDNHWLFLQLCFYTELLTCYALMFSFCHYYYFLPLKKRNLDLFVFRHELAIMRLAAFMGITMLVGITGLFYTQLIGIITDTTSIEKMSNCCEDISRPRKPWQQTFSEVFGTRWKILWFIPFRQRQPLRVPYHFANHV",
31
32 "MLLLGAVLLLLALPGHDQETTTQGPGVLLPLPKGACTGWMAGIPGHPGHNGAPGRDGRDGTPGEKGEKGDPGLIGPKGDIGETGVPGAEGPRGFPGIQGRKGEPGEGAYVYRSAFSVGLETYVTIPNMPIRFTKIFYNQQNHYDGSTGKFHCNIPGLYYFAYHITVYMKDVKVSLFKKDKAMLFTYDQYQENNVDQASGSVLLHLEVGDQVWLQVYGEGERNGLYADNDNDSTFTGFLLYHDTN",
33 "MGLLAFLKTQFVLHLLVGFVFVVSGLVINFVQLCTLALWPVSKQLYRRLNCRLAYSLWSQLVMLLEWWSCTECTLFTDQATVERFGKEHAVIILNHNFEIDFLCGWTMCERFGVLGSSKVLAKKELLYVPLIGWTWYFLEIVFCKRKWEEDRDTVVEGLRRLSDYPEYMWFLLYCEGTRFTETKHRVSMEVAAAKGLPVLKYHLLPRTKGFTTAVKCLRGTVAAVYDVTLNFRGNKNPSLLGILYGKKYEADMCVRRFPLEDIPLDEKEAAQWLHKLYQEKDALQEIYNQKGMFPGEQFKPARRPWTLLNFLSWATILLSPLFSFVLGVFASGSPLLILTFLGFVGAASFGVRRLIGVTEIEKGSSYGNQEFKKKE",
34 "MDLAGLLKSQFLCHLVFCYVFIASGLIINTIQLFTLLLWPINKQLFRKINCRLSYCISSQLVMLLEWWSGTECTIFTDPRAYLKYGKENAIVVLNHKFEIDFLCGWSLSERFGLLGGSKVLAKKELAYVPIIGWMWYFTEMVFCSRKWEQDRKTVATSLQHLRDYPEKYFFLIHCEGTRFTEKKHEISMQVARAKGLPRLKHHLLPRTKGFAITVRSLRNVVSAVYDCTLNFRNNENPTLLGVLNGKKYHADLYVRRIPLEDIPEDDDECSAWLHKLYQEKDAFQEEYYRTGTFPETPMVPPRRPWTLVNWLFWASLVLYPFFQFLVSMIRSGSSLTLASFILVFFVASVGVRWMIGVTEIDKGSAYGNSDSKQKLND",
35
36 "MALLLCFVLLCGVVDFARSLSITTPEEMIEKAKGETAYLPCKFTLSPEDQGPLDIEWLISPADNQKVDQVIILYSGDKIYDDYYPDLKGRVHFTSNDLKSGDASINVTNLQLSDIGTYQCKVKKAPGVANKKIHLVVLVKPSGARCYVDGSEEIGSDFKIKCEPKEGSLPLQYEWQKLSDSQKMPTSWLAEMTSSVISVKNASSEYSGTYSCTVRNRVGSDQCLLRLNVVPPSNKAGLIAGAIIGTLLALALIGLIIFCCRKKRREEKYEKEVHHDIREDVPPPKSRTSTARSYIGSNHSSLGSMSPSNMEGYSKTQYNQVPSEDFERTPQSPTLPPAKVAAPNLSRMGAIPVMIPAQSKDGSIV",
37 "MSYVFVNDSSQTNVPLLQACIDGDFNYSKRLLESGFDPNIRDSRGRTGLHLAAARGNVDICQLLHKFGADLLATDYQGNTALHLCGHVDTIQFLVSNGLKIDICNHQGATPLVLAKRRGVNKDVIRLLESLEEQEVKGFNRGTHSKLETMQTAESESAMESHSLLNPNLQQGEGVLSSFRTTWQEFVEDLGFWRVLLLIFVIALLSLGIAYYVSGVLPFVENQPELVH",
38 "MRVAGAAKLVVAVAVFLLTFYVISQVFEIKMDASLGNLFARSALDTAARSTKPPRYKCGISKACPEKHFAFKMASGAANVVGPKICLEDNVLMSGVKNNVGRGINVALANGKTGEVLDTKYFDMWGGDVAPFIEFLKAIQDGTIVLMGTYDDGATKLNDEARRLIADLGSTSITNLGFRDNWVFCGGKGIKTKSPFEQHIKNNKDTNKYEGWPEVVEMEGCIPQKQD",
39
40 "MAPAAATGGSTLPSGFSVFTTLPDLLFIFEFIFGGLVWILVASSLVPWPLVQGWVMFVSVFCFVATTTLIILYIIGAHGGETSWVTLDAAYHCTAALFYLSASVLEALATITMQDGFTYRHYHENIAAVVFSYIATLLYVVHAVFSLIRWKSS",
41 "MRLQGAIFVLLPHLGPILVWLFTRDHMSGWCEGPRMLSWCPFYKVLLLVQTAIYSVVGYASYLVWKDLGGGLGWPLALPLGLYAVQLTISWTVLVLFFTVHNPGLALLHLLLLYGLVVSTALIWHPINKLAALLLLPYLAWLTVTSALTYHLWRDSLCPVHQPQPTEKSD",
42 "MEESVVRPSVFVVDGQTDIPFTRLGRSHRRQSCSVARVGLGLLLLLMGAGLAVQGWFLLQLHWRLGEMVTRLPDGPAGSWEQLIQERRSHEVNPAAHLTGANSSLTGSGGPLLWETQLGLAFLRGLSYHDGALVVTKAGYYYIYSKVQLGGVGCPLGLASTITHGLYKRTPRYPEELELLVSQQSPCGRATSSSRVWWDSSFLGGVVHLEAGEKVVVRVLDERLVRLRDGTRSYFGAFMV"
43]
44
45def compute_average_mlm_loss(protein1, protein2, iterations=10):
46 total_loss = 0.0
47 connector = "G" * 25 # Connector sequence of G's
48 for _ in range(iterations):
49 concatenated_sequence = protein1 + connector + protein2
50 inputs = tokenizer(concatenated_sequence, return_tensors="pt", padding=True, truncation=True, max_length=1024)
51
52 mask_prob = 0.35
53 mask_indices = torch.rand(inputs["input_ids"].shape, device=device) < mask_prob
54
55 # Locate the positions of the connector 'G's and set their mask indices to False
56 connector_indices = tokenizer.encode(connector, add_special_tokens=False)
57 connector_length = len(connector_indices)
58 start_connector = len(tokenizer.encode(protein1, add_special_tokens=False))
59 end_connector = start_connector + connector_length
60
61 # Avoid masking the connector 'G's
62 mask_indices[0, start_connector:end_connector] = False
63
64 # Apply the mask to the input IDs
65 inputs["input_ids"][mask_indices] = tokenizer.mask_token_id
66 inputs = {k: v.to(device) for k, v in inputs.items()} # Send inputs to the device
67
68 with torch.no_grad():
69 outputs = model(**inputs, labels=inputs["input_ids"])
70
71 loss = outputs.loss
72 total_loss += loss.item()
73
74 return total_loss / iterations
75
76# Compute all average losses to determine the maximum threshold for the slider
77all_losses = []
78for i, protein1 in enumerate(all_proteins):
79 for j, protein2 in enumerate(all_proteins[i+1:], start=i+1):
80 avg_loss = compute_average_mlm_loss(protein1, protein2)
81 all_losses.append(avg_loss)
82
83# Set the maximum threshold to the maximum loss computed
84max_threshold = max(all_losses)
85print(f"Maximum loss (maximum threshold for slider): {max_threshold}")
86
87def plot_graph(threshold):
88 G = nx.Graph()
89
90 # Add all protein nodes to the graph
91 for i, protein in enumerate(all_proteins):
92 G.add_node(f"protein {i+1}")
93
94 # Loop through all pairs of proteins and calculate average MLM loss
95 loss_idx = 0 # Index to keep track of the position in the all_losses list
96 for i, protein1 in enumerate(all_proteins):
97 for j, protein2 in enumerate(all_proteins[i+1:], start=i+1):
98 avg_loss = all_losses[loss_idx]
99 loss_idx += 1
100
101 # Add an edge if the loss is below the threshold
102 if avg_loss < threshold:
103 G.add_edge(f"protein {i+1}", f"protein {j+1}", weight=round(avg_loss, 3))
104
105 # 3D Network Plot
106 # Adjust the k parameter to bring nodes closer. This might require some experimentation to find the right value.
107 k_value = 2 # Lower value will bring nodes closer together
108 pos = nx.spring_layout(G, dim=3, seed=42, k=k_value)
109
110 edge_x = []
111 edge_y = []
112 edge_z = []
113 for edge in G.edges():
114 x0, y0, z0 = pos[edge[0]]
115 x1, y1, z1 = pos[edge[1]]
116 edge_x.extend([x0, x1, None])
117 edge_y.extend([y0, y1, None])
118 edge_z.extend([z0, z1, None])
119
120 edge_trace = go.Scatter3d(x=edge_x, y=edge_y, z=edge_z, mode='lines', line=dict(width=0.5, color='grey'))
121
122 node_x = []
123 node_y = []
124 node_z = []
125 node_text = []
126 for node in G.nodes():
127 x, y, z = pos[node]
128 node_x.append(x)
129 node_y.append(y)
130 node_z.append(z)
131 node_text.append(node)
132
133 node_trace = go.Scatter3d(x=node_x, y=node_y, z=node_z, mode='markers', marker=dict(size=5), hoverinfo='text', hovertext=node_text)
134
135 layout = go.Layout(title='Protein Interaction Graph', title_x=0.5, scene=dict(xaxis=dict(showbackground=False), yaxis=dict(showbackground=False), zaxis=dict(showbackground=False)))
136
137 fig = go.Figure(data=[edge_trace, node_trace], layout=layout)
138 fig.show()
139
140# Create an interactive slider for the threshold value with a default of 8.50
141interact(plot_graph, threshold=widgets.FloatSlider(min=0.0, max=max_threshold, step=0.05, value=8.25))