Team Ai
Modelpublic

OneScience-Group/TemStaPro-main

sourceHugging Facemitupdated 1mo agoView on Hugging Face
0likes17downloads
data_process.py135 linesDownload Raw Back to scripts
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