{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":41875,"databundleVersionId":5521661,"sourceType":"competition"}],"dockerImageVersionId":30527,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n# Introduction\n- The main focus of this notebook is to combine all available data sources into one NPZ file (https://numpy.org/doc/stable/reference/generated/numpy.savez.html)\n- the file created is later used for training ","metadata":{}},{"cell_type":"markdown","source":"# Configurations \n- one can specify **min** and **max length** of the protein sequences \n- keep in mind the notebook runs **max 9h**\n- very long protein sequences might result in an **out of memory error**, if you feed them into ProtBert","metadata":{}},{"cell_type":"code","source":"#filters protein sequences shorter than n out (0 == no min length)\nmin_sq_length_config = 1000\n\n#filters protein sequences longer than n out\nmax_sq_length_config = 1100\n\n#how much protein sequences are pushed through protbert at once #(higher values lead to out of memory error)\nprotbert_batch_size = 10\n\n#the name of your resulting file\nsave_file_name = str(min_sq_length_config) + \"-\" + str(max_sq_length_config) + \"length_NEW_FORMAT\"\n\n#use only the most common gene ontologies, in order to reduce training time\nmost_common_gene_ontologies = 500\n\nsave_file_name","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd \nimport os\n\nfrom transformers import BertModel, BertTokenizer\nimport torch\nfrom Bio import SeqIO\nimport plotly.graph_objects as go\nfrom collections import Counter\n\nimport re\nimport sys\n\nfrom timeit import default_timer as timer\nfrom datetime import timedelta\nimport time \n\nfrom scipy.special import softmax\nnp.set_printoptions(precision=5)\n\ntorch.set_printoptions(threshold=5)\ntorch.manual_seed(0)\n\n!python --version","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-03T23:01:14.847098Z","iopub.execute_input":"2023-08-03T23:01:14.847602Z","iopub.status.idle":"2023-08-03T23:01:29.073937Z","shell.execute_reply.started":"2023-08-03T23:01:14.84756Z","shell.execute_reply":"2023-08-03T23:01:29.072693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# desc\n- load the necessary data\n> 1. **Protein Sequences**   \n> 1. **EntryIDs** (ids of the protein sequences)  \n> 1. **Gene Ontologies** (GO) [these are the attributes of the protein] \n- filter prefered protein sequence lengths out\n- get the 500 most common GO\n\n\n","metadata":{}},{"cell_type":"code","source":"#check gpu is available\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-08-03T23:01:29.077684Z","iopub.execute_input":"2023-08-03T23:01:29.078147Z","iopub.status.idle":"2023-08-03T23:01:29.117017Z","shell.execute_reply.started":"2023-08-03T23:01:29.078111Z","shell.execute_reply":"2023-08-03T23:01:29.116095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# load and filter data","metadata":{}},{"cell_type":"code","source":"#load training data\ntrain_sequences_fasta    = \"/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta\"\ntrain_sequences_full     = SeqIO.parse(train_sequences_fasta, 'fasta')\ntrain_sequences          = np.array([str(seq.seq) for seq in SeqIO.parse(train_sequences_fasta, 'fasta')], dtype=object)\ntrain_sequences_entry_id = np.array([id.id for id in SeqIO.parse(train_sequences_fasta, 'fasta')], dtype=object)\n\n#example of an entry_id with related protein sequence\ntrain_sequences_entry_id[1], train_sequences[1]","metadata":{"execution":{"iopub.status.busy":"2023-08-03T23:01:29.118799Z","iopub.execute_input":"2023-08-03T23:01:29.121688Z","iopub.status.idle":"2023-08-03T23:01:33.447138Z","shell.execute_reply.started":"2023-08-03T23:01:29.121662Z","shell.execute_reply":"2023-08-03T23:01:33.44621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#filter out too long and too short protein sequences (see (2) config)\n\nreduced_training_sequences = []\nfor i, ts in enumerate(train_sequences):\n     if len(ts) >= min_sq_length_config and len(ts) < max_sq_length_config:\n            reduced_training_sequences.append((train_sequences_entry_id[i], ts))\n            \nprint(\"protein sequences with the length of \", min_sq_length_config ,\"to\" , max_sq_length_config,\"->\",  len(reduced_training_sequences))","metadata":{"execution":{"iopub.status.busy":"2023-08-03T23:05:44.702656Z","iopub.execute_input":"2023-08-03T23:05:44.703027Z","iopub.status.idle":"2023-08-03T23:05:44.790221Z","shell.execute_reply.started":"2023-08-03T23:05:44.702997Z","shell.execute_reply":"2023-08-03T23:05:44.789233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Gene Ontologies","metadata":{}},{"cell_type":"code","source":"#this methods will help with finding the most common gene ontologies (GO)\ndef count_elements(lst):\n    counts = {}\n    for element in lst:\n        if element in counts:\n            counts[element] += 1\n        else:\n            counts[element] = 1\n    return counts\n\ndef sort_by_count(counts):\n    sorted_counts = sorted(counts.items(), key=lambda x: x[1], reverse=True)\n    return sorted_counts","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:04.079246Z","iopub.execute_input":"2023-08-03T21:51:04.079594Z","iopub.status.idle":"2023-08-03T21:51:04.087168Z","shell.execute_reply.started":"2023-08-03T21:51:04.079562Z","shell.execute_reply":"2023-08-03T21:51:04.084963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#filter unique gene ontologies (GOs) and create one hot encodings (target vecs)\ngeneOntologyTerms = pd.read_csv(\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv\",sep=\"\\t\")\ngeneOntologyTerms =  geneOntologyTerms.sort_values(by=[\"EntryID\"])\nallGO = geneOntologyTerms[\"term\"]\nuniqueGO = allGO.unique()\nprint(\"all go:    \" + str(len(allGO)))\nprint(\"unique go: \" + str(len(uniqueGO)))\n#uniqueGO.sort()\n\n#find the 500 most frequent GOs\ncommon_go_with_count = sort_by_count(count_elements(allGO))\ncommon_go_with_count = common_go_with_count[:most_common_gene_ontologies]\n\ncommon_go = []\nfor cg in common_go_with_count:\n    common_go.append(cg[0])\n\n#sort by name\ncommon_go.sort()\n#oneHotGo = pd.get_dummies(common_go)","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:04.088811Z","iopub.execute_input":"2023-08-03T21:51:04.08917Z","iopub.status.idle":"2023-08-03T21:51:17.592867Z","shell.execute_reply.started":"2023-08-03T21:51:04.089114Z","shell.execute_reply":"2023-08-03T21:51:17.589792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# one hot encodings from gene ontologies","metadata":{}},{"cell_type":"code","source":"#create one hots\none_hot = torch.nn.functional.one_hot(torch.arange(0, len(common_go)))\none_hot_encode_gos = {}\n\n#assign GOs to one hots\nfor i, key in enumerate(common_go):\n    one_hot_encode_gos[key] = one_hot[i]\n\n#also create a invertet list\ninverted_one_hot_encode_gos = {v: k for k, v in one_hot_encode_gos.items()}\n    ","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.594687Z","iopub.status.idle":"2023-08-03T21:51:17.595197Z","shell.execute_reply.started":"2023-08-03T21:51:17.594933Z","shell.execute_reply":"2023-08-03T21:51:17.594956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#resulting dict, key==GO and value==one_hot\nfor i, key in enumerate(one_hot_encode_gos):\n    print(key + \" ->\", one_hot_encode_gos[key])\n    if i == 5: break","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.597221Z","iopub.status.idle":"2023-08-03T21:51:17.597686Z","shell.execute_reply.started":"2023-08-03T21:51:17.59745Z","shell.execute_reply":"2023-08-03T21:51:17.597473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ProtBert (last hidden state)","metadata":{}},{"cell_type":"code","source":"#init ProtBert\ntokenizer = BertTokenizer.from_pretrained(\"Rostlab/prot_bert\", do_lower_case=False)\nbert_model = BertModel.from_pretrained(\"Rostlab/prot_bert\").to(device)","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.599012Z","iopub.status.idle":"2023-08-03T21:51:17.599813Z","shell.execute_reply.started":"2023-08-03T21:51:17.599567Z","shell.execute_reply":"2023-08-03T21:51:17.599594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#returns the last hidden states of ProtBert, which is later fed into a NN  \ndef getLastHiddenStates(t_sequences: list):\n    \n    if not isinstance(t_sequences, list):\n        print(\"wrong type! must be list of String(s)!\")\n        return\n\n    spaced_sequences = []\n    for s in t_sequences:\n         spaced_sequences.append(re.sub(r\"[UZOB]\", \"X\", \" \".join([*str(s)])))\n    encoded_input = tokenizer(spaced_sequences, return_tensors='pt',padding=True).to(device)\n    bert_model_output = bert_model(**encoded_input)\n    return bert_model_output","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.601565Z","iopub.status.idle":"2023-08-03T21:51:17.602021Z","shell.execute_reply.started":"2023-08-03T21:51:17.601784Z","shell.execute_reply":"2023-08-03T21:51:17.601805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#collect ProtBert embeddings (== last hidden state of the model)\n\nwith torch.no_grad():\n    \n    start = timer()\n    protbert_last_hidden_state = []\n    \n    for i in range(0,len(reduced_training_sequences), protbert_batch_size ):\n        \n        #measure passed time\n        if (i%100) == 0 and i > 0:\n            end = timer()\n            print(i, timedelta(seconds=end-start))\n            start = timer()\n            \n        batch = []\n        for protein_sqnc in reduced_training_sequences[i:i+protbert_batch_size]:\n            #print(len(protein_sqnc[1]))\n            batch.append(protein_sqnc[1])\n\n        #print(batch)\n        last_hidden_states = getLastHiddenStates(batch)\n\n        #batch size\n        #print((last_hidden_states.last_hidden_state))\n        for x in last_hidden_states.last_hidden_state:\n            protbert_last_hidden_state.append(x[-1].detach().cpu())","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.603665Z","iopub.status.idle":"2023-08-03T21:51:17.604111Z","shell.execute_reply.started":"2023-08-03T21:51:17.603883Z","shell.execute_reply":"2023-08-03T21:51:17.603904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# collect data and save","metadata":{}},{"cell_type":"code","source":"geneOntologyTerms =  geneOntologyTerms.sort_values(by=[\"EntryID\"])\ngos_numpy = geneOntologyTerms.to_numpy()\n\n\nindex_search_go = geneOntologyTerms[\"EntryID\"].to_numpy()\nindex_gene_ontologys = geneOntologyTerms[\"term\"].to_numpy()","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.605781Z","iopub.status.idle":"2023-08-03T21:51:17.606777Z","shell.execute_reply.started":"2023-08-03T21:51:17.606537Z","shell.execute_reply":"2023-08-03T21:51:17.606561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#buld training dataframe and save to disk\n\n\n# geneOntologyTerms =  geneOntologyTerms.sort_values(by=[\"EntryID\"])\n# gos_numpy = geneOntologyTerms.to_numpy()\n\n\n\ntarget_vecs_list = [] \n\n\n\nstart = timer()\nfor i in range(len(reduced_training_sequences)):\n    \n    if (i%100) == 0 and i > 0:\n        end = timer()\n        print(i,timedelta(seconds=end-start))\n        start = timer()\n\n\n   \n    \n    #get entry id in order to fetch the go terms\n    entry_id = reduced_training_sequences[i][0]\n    \n    #get all related go terms\n    idx = np.where(index_search_go == entry_id)[0]\n    relatedGOs = index_gene_ontologys[idx]\n    \n    target_vec = torch.zeros([1, len(common_go)], dtype=torch.int)\n\n    #iterate over go terms to create a target vector for a specific protein\n    for entry in relatedGOs:\n        go_key = entry\n        \n \n        \n        target_vec_tmp_list = []\n        \n        #is go in most_common_gos?\n        if go_key in one_hot_encode_gos: \n            target_vec = target_vec +  one_hot_encode_gos[go_key]\n\n    target_vecs_list.append(target_vec)        \n\nprint(len(target_vecs_list))\n\ntarget_vecs_list[:5]","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.608042Z","iopub.status.idle":"2023-08-03T21:51:17.608879Z","shell.execute_reply.started":"2023-08-03T21:51:17.608643Z","shell.execute_reply":"2023-08-03T21:51:17.608666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#print(reduced_training_sequences[0] ,protbert_last_hidden_state[0], target_vec[0])\nlen(reduced_training_sequences) ,len(protbert_last_hidden_state), len(target_vecs_list)","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.610543Z","iopub.status.idle":"2023-08-03T21:51:17.611342Z","shell.execute_reply.started":"2023-08-03T21:51:17.611079Z","shell.execute_reply":"2023-08-03T21:51:17.611101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"reduced_training_sequences[0] ,protbert_last_hidden_state[0], target_vecs_list[0]","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.613353Z","iopub.status.idle":"2023-08-03T21:51:17.614022Z","shell.execute_reply.started":"2023-08-03T21:51:17.613774Z","shell.execute_reply":"2023-08-03T21:51:17.613796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#TODO save\nnp.savez(\"/kaggle/working/\" + save_file_name ,entry_id=[x[0] for x in reduced_training_sequences],  train_sequence=[x[1] for x in reduced_training_sequences],target_vec=target_vecs_list, last_hidden_state=protbert_last_hidden_state, allow_pickle=True )","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.615563Z","iopub.status.idle":"2023-08-03T21:51:17.61646Z","shell.execute_reply.started":"2023-08-03T21:51:17.616189Z","shell.execute_reply":"2023-08-03T21:51:17.616215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_loaded = np.load(\"/kaggle/working/\" + save_file_name+ \".npz\", allow_pickle=True)\ndata_loaded[\"target_vec\"]\ndata_loaded.files\n\ndata_loaded[\"entry_id\"][10], data_loaded[\"train_sequence\"][10], data_loaded[\"target_vec\"][10], data_loaded[\"last_hidden_state\"][10]","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.618102Z","iopub.status.idle":"2023-08-03T21:51:17.618976Z","shell.execute_reply.started":"2023-08-03T21:51:17.618728Z","shell.execute_reply":"2023-08-03T21:51:17.618752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#check positions right\ndef target_vec_to_gos(t_vec):\n    elms = []\n    indcs = np.where(t_vec.numpy() == 1)\n\n    for i, el in enumerate(one_hot_encode_gos.items()):\n        if i  in indcs[0]:\n            elms.append(el)\n            \n    return elms","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.620387Z","iopub.status.idle":"2023-08-03T21:51:17.621196Z","shell.execute_reply.started":"2023-08-03T21:51:17.620935Z","shell.execute_reply":"2023-08-03T21:51:17.620958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# reduced_training_sequences[1][0]","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.622849Z","iopub.status.idle":"2023-08-03T21:51:17.623923Z","shell.execute_reply.started":"2023-08-03T21:51:17.623674Z","shell.execute_reply":"2023-08-03T21:51:17.623697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# first_target_vec = sorted(target_vec_to_gos(target_vecs_list[1][0].detach().cpu()))\n# all_goes_in_ftv = []\n# for ftv in first_target_vec:\n#     all_goes_in_ftv.append(ftv[0])  \n\n# all_goes_in_ftv","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.625199Z","iopub.status.idle":"2023-08-03T21:51:17.6261Z","shell.execute_reply.started":"2023-08-03T21:51:17.625819Z","shell.execute_reply":"2023-08-03T21:51:17.625845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# idx = np.where(gos_numpy == 'O94457')[0]\n# entry_id_goes = gos_numpy[idx]\n\n# all_goes_via_entry_id = []\n# for el in entry_id_goes:\n#     all_goes_via_entry_id.append(el[1])\n    \n# all_goes_in_eid = sorted(all_goes_via_entry_id)\n\n\n","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.62777Z","iopub.status.idle":"2023-08-03T21:51:17.628629Z","shell.execute_reply.started":"2023-08-03T21:51:17.628392Z","shell.execute_reply":"2023-08-03T21:51:17.628415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for i, a in enumerate(all_goes_in_ftv):\n#     if a in all_goes_in_eid:\n#         print(i , a)\n#     else:\n#         print(\"not in\")","metadata":{"execution":{"iopub.status.busy":"2023-08-03T21:51:17.630222Z","iopub.status.idle":"2023-08-03T21:51:17.631017Z","shell.execute_reply.started":"2023-08-03T21:51:17.630768Z","shell.execute_reply":"2023-08-03T21:51:17.630791Z"},"trusted":true},"execution_count":null,"outputs":[]}]}