OneScience-Group/TemStaPro-main
017
1"""2Process the data set before the inference process.3"""4 5import numpy6import torch7from hashlib import sha2568from os import path9 10def get_sequences_without_embeddings(sequences, emb_dir, per_res=False):11 """12 Collecting sequences that do not have generated embeddings.13 14 sequences - DICT of all sequences in the input (keys are sequence ids, 15 values are protein sequences16 emb_dir - STRING that defines the directory where embeddings are saved17 per_res - BOOL that determines whether per-residue embeddings are needed18 19 returns DICT with sequences that lack embeddings20 """21 seqs_wo_emb = {}22 for seq_id in list(sequences.keys()):23 seq_code = sha256(sequences[seq_id].encode('utf-8')).hexdigest()24 if(not path.exists(f"{emb_dir}/mean_{seq_code}.pt")):25 seqs_wo_emb[seq_id] = sequences[seq_id]26 if(per_res and not path.exists(f"{emb_dir}/per_res_{seq_code}.pt")):27 seqs_wo_emb[seq_id] = sequences[seq_id]28 return seqs_wo_emb29 30def collect_mean_embeddings(sequences, embeddings, emb_dir, input_size=1024):31 """32 Collecting mean embeddings into a dictionary.33 34 sequences - DICT of all sequences in the input (keys are sequence ids,35 values are protein sequences36 embeddings - DICT with generated embeddings. Keys are "mean_representations"37 and "per_res_representations", which have [DICT] values, which keys are 38 sequence ids and values are embeddings torch tensor39 emb_dir - STRING that determines the path to the embeddings 'cache' 40 directory41 input_size - INT that notes the dimension of each embeddings vector42 43 returns DICT with keys "x_test" (values are embeddings tensors) and 44 "y_test" (values are (irrelevant) temperature labels)45 """46 dataset = {}47 dataset['y_test'] = torch.tensor((), dtype=torch.int32)48 for i, seq_id in enumerate(sequences):49 if(emb_dir and path.exists(emb_dir)):50 # Loading sequences from cache51 embedding = torch.load("%s/mean_%s.pt" % (emb_dir,52 sha256(sequences[seq_id].encode('utf-8')).hexdigest()))["mean_representations"]53 else:54 # Taking freshly-generated embeddings55 embedding = torch.from_numpy(embeddings["mean_representations"][seq_id])56 if(i):57 dataset["x_test"] = torch.vstack((dataset["x_test"], torch.flatten(embedding)))58 else:59 dataset["x_test"] = torch.reshape(embedding, (1, input_size))60 dataset["y_test"] = torch.cat((dataset["y_test"], torch.tensor([999]).int()), 0)61 return dataset62 63def collect_per_res_embeddings(sequences, original_sequences, embeddings, emb_dir, 64 input_size=1024, smoothen=False, window_size=21):65 """66 Collecting per-residue embeddings into a dictionary.67 68 sequences - DICT of all sequences in the input (keys are sequence ids,69 values are protein sequences70 embeddings - DICT with generated embeddings. Keys are "mean_representations"71 and "per_res_representations", which have [DICT] values, which keys are 72 sequence ids and values are embeddings torch tensor73 emb_dir - STRING that determines the path to the embeddings 'cache' 74 directory75 input_size - INT that notes the dimension of each embeddings vector76 smoothen - BOOL indicates whether to make average smoothing of embeddings77 78 returns DICT with keys "x_test" (values are embeddings tensors) and 79 "y_test" (values are fake temperature labels)80 """81 dataset = {}82 dataset['y_test'] = torch.tensor((), dtype=torch.int32)83 dataset['z_test'] = {}84 85 for i, seq_id in enumerate(sequences):86 87 iterations_for_seq = len(sequences[seq_id])88 89 if(emb_dir and path.exists(emb_dir)):90 embedding = torch.load("%s/per_res_%s.pt" % (emb_dir,91 sha256(sequences[seq_id].encode('utf-8')).hexdigest()))["per_res_representations"]92 else:93 # Taking freshly-generated embeddings94 embedding = torch.from_numpy(embeddings["per_res_representations"][seq_id])95 96 for j in range(iterations_for_seq):97 if(i == 0 and j == 0):98 dataset["x_test"] = torch.reshape(embedding[j], (1, input_size))99 else:100 dataset["x_test"] = torch.vstack((dataset["x_test"], torch.flatten(embedding[j])))101 if(not smoothen): dataset["y_test"] = torch.cat((dataset["y_test"], torch.tensor([999]).int()), 0)102 if(not smoothen): dataset["z_test"]['%s_%d' % (seq_id, j)] = original_sequences[seq_id][j]103 104 if(smoothen):105 WINDOW_SIZE = window_size106 smoothened_seqs = {}107 j = 0108 while(j < iterations_for_seq-WINDOW_SIZE+1):109 smoothened_embedding = dataset["x_test"][range(j, j+WINDOW_SIZE)].mean(dim=0)110 if(not j and not i):111 smoothened_embeddings = smoothened_embedding112 else:113 smoothened_embeddings = torch.vstack((smoothened_embeddings, smoothened_embedding))114 dataset['z_test']['%s_%d-%d' % (seq_id, j, j+WINDOW_SIZE)] = ''.join(original_sequences[seq_id][j:j+WINDOW_SIZE])115 dataset["y_test"] = torch.cat((dataset["y_test"], torch.tensor([999]).int()), 0)116 j += 1117 118 if(smoothen): dataset["x_test"] = smoothened_embeddings119 return dataset120 121def load_tensor_from_NPZ(NPZ_file, keywords):122 """123 Loading embeddings from file to dictionary.124 125 NPZ_file - STRING path to the NPZ file126 keywords - LIST with keywords to identify which subset of file to load127 128 returns DICT with keys as given keywords, values in tensors129 """130 dataset = {}131 with numpy.load(NPZ_file, allow_pickle=True) as data_loaded:132 for i in range(len(keywords)):133 dataset[keywords[i]] = torch.from_numpy(data_loaded[keywords[i]])134 return dataset135 