Repository navigation
Expand file tree
/
Copy pathevofea_embedding.py
More file actions
134 lines (110 loc) · 6.09 KB
/
Copy pathevofea_embedding.py
File metadata and controls
134 lines (110 loc) · 6.09 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
# This file is to generate protein residual level embedding using ProtT5
from transformers import T5EncoderModel, T5Tokenizer
import torch
import h5py
import time
import argparse
import os
from Bio import SeqIO
device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')
#@title Load ProtT5 in half-precision. { display-mode: "form" }
# Load ProtT5 in half-precision (more specifically: the encoder-part of ProtT5-XL-U50)
def get_T5_model():
#model = T5EncoderModel.from_pretrained("./protT5/protT5_checkpoint/", torch_dtype=torch.float16).to(device)
model = T5EncoderModel.from_pretrained("Rostlab/prot_t5_xl_half_uniref50-enc").to(device)
model.half()
model = model.eval() # set model to evaluation model
tokenizer = T5Tokenizer.from_pretrained("Rostlab/prot_t5_xl_uniref50", do_lower_case=False )
return model, tokenizer
# @title Read in file in fasta format. { display-mode: "form" }
def read_fasta(fasta_path, split_char="!", id_field=0):
'''
Reads in fasta file containing multiple sequences.
Split_char and id_field allow to control identifier extraction from header.
E.g.: set split_char="|" and id_field=1 for SwissProt/UniProt Headers.
Returns dictionary holding multiple sequences or only single
sequence, depending on input file.
'''
seqs = dict()
with open(fasta_path, 'r') as fasta_f:
for line in fasta_f:
# get uniprot ID from header and create new entry
if line.startswith('>'):
uniprot_id = line.replace('>', '').strip().split(split_char)[id_field]
# replace tokens that are mis-interpreted when loading h5
uniprot_id = uniprot_id.replace("/", "_").replace(".", "_")
seqs[uniprot_id] = ''
else:
# repl. all whie-space chars and join seqs spanning multiple lines, drop gaps and cast to upper-case
seq = ''.join(line.split()).upper().replace("-", "")
# repl. all non-standard AAs and map them to unknown/X
seq = seq.replace('U', 'X').replace('Z', 'X').replace('O', 'X')
seqs[uniprot_id] += seq
return seqs
# @title Generate embeddings. { display-mode: "form" }
# Generate embeddings via batch-processing
# per_residue indicates that embeddings for each residue in a protein should be returned.
# per_protein indicates that embeddings for a whole protein should be returned (average-pooling)
# max_residues gives the upper limit of residues within one batch
# max_seq_len gives the upper sequences length for applying batch-processing
# max_batch gives the upper number of sequences per batch
def get_embeddings(model, tokenizer, seqs, per_residue, per_protein, max_residues=4000, max_seq_len=1000, max_batch=100):
results = {"residue_embs": dict(),"protein_embs": dict() }
# sort sequences according to length (reduces unnecessary padding --> speeds up embedding)
seq_dict = sorted(seqs.items(), key=lambda kv: len(seqs[kv[0]]), reverse=True)
start = time.time()
batch = list()
for seq_idx, (pdb_id, seq) in enumerate(seq_dict, 1):
seq = seq
if len(seq) > 1000:
seq = seq[:500] + seq[-500:]
seq_len = len(seq)
seq = ' '.join(list(seq))
batch.append((pdb_id, seq, seq_len))
n_res_batch = sum([s_len for _, _, s_len in batch]) + seq_len
if len(batch) >= max_batch or n_res_batch >= max_residues or seq_idx == len(seq_dict) or seq_len > max_seq_len:
pdb_ids, seqs, seq_lens = zip(*batch)
batch = list()
# add_special_tokens adds extra token at the end of each sequence
token_encoding = tokenizer.batch_encode_plus(seqs, add_special_tokens=True, padding="longest")
input_ids = torch.tensor(token_encoding['input_ids']).to(device)
attention_mask = torch.tensor(token_encoding['attention_mask']).to(device)
try:
with torch.no_grad():
# returns: ( batch-size x max_seq_len_in_minibatch x embedding_dim )
embedding_repr = model(input_ids, attention_mask=attention_mask)
except RuntimeError:
print("RuntimeError during embedding for {} (L={})".format(pdb_id, seq_len))
continue
for batch_idx, identifier in enumerate(pdb_ids): # for each protein in the current mini-batch
s_len = seq_lens[batch_idx]
# slice off padding --> batch-size x seq_len x embedding_dim
emb = embedding_repr.last_hidden_state[batch_idx, :s_len]
if per_residue: # store per-residue embeddings (Lx1024)
results["residue_embs"][identifier] = emb.detach().cpu().numpy().squeeze()
if per_protein: # apply average-pooling to derive per-protein embeddings (1024-d)
protein_emb = emb.mean(dim=0)
results["protein_embs"][identifier] = protein_emb.detach().cpu().numpy().squeeze()
passed_time = time.time() - start
avg_time = passed_time / len(results["residue_embs"]) if per_residue else passed_time / len(results["protein_embs"])
print('\n############# EMBEDDING STATS #############')
print('Total number of per-residue embeddings: {}'.format(len(results["residue_embs"])))
print('Total number of per-protein embeddings: {}'.format(len(results["protein_embs"])))
print("Time for generating embeddings: {:.1f}[m] ({:.3f}[s/protein])".format(passed_time / 60, avg_time))
print('\n############# END #############')
return results
def save_embeddings(emb_dict,out_path):
with h5py.File(str(out_path), "w") as hf:
for sequence_id, embedding in emb_dict.items():
hf.create_dataset(sequence_id, data=embedding)
return None
def write_prediction_fasta(predictions, out_path):
class_mapping = {0:"H",1:"E",2:"L"}
with open(out_path, 'w+') as out_f:
out_f.write( '\n'.join(
[ ">{}\n{}".format(
seq_id, ''.join( [class_mapping[j] for j in yhat] ))
for seq_id, yhat in predictions.items()
]
) )
return None