{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This notebook demonstrates how to use the Ankh  protein language model to extract embeddings from provided protein sequences. The following code calculates the embeddings for each protein sequence in the training and test FASTA files and saves them as parquet files. \nThis is slow to run! \n\nComputing the embeddings and subsequently reading in the `.pt` files can take a while. The resulting  arrays can be found in a dataset (to upload)\n\n* https://github.com/agemagician/Ankh\n\nEDIT: Due to memory issues: run only on train in this version","metadata":{}},{"cell_type":"code","source":"!pip install -q  ankh","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-30T10:14:37.987267Z","iopub.execute_input":"2023-04-30T10:14:37.987716Z","iopub.status.idle":"2023-04-30T10:15:00.420904Z","shell.execute_reply.started":"2023-04-30T10:14:37.987676Z","shell.execute_reply":"2023-04-30T10:15:00.419602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pathlib\nimport torch\nimport ankh\nimport numpy as np\nimport pandas as pd\nfrom Bio import SeqIO\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:15:00.423797Z","iopub.execute_input":"2023-04-30T10:15:00.424454Z","iopub.status.idle":"2023-04-30T10:15:05.142258Z","shell.execute_reply.started":"2023-04-30T10:15:00.424399Z","shell.execute_reply":"2023-04-30T10:15:05.14107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\nprint('Available device:', device)\n\ndef read_fasta(fastaPath):    \n    fasta_sequences = SeqIO.parse(open(fastaPath), 'fasta')\n    ids = []\n    sequences = []\n    for fasta in fasta_sequences:\n        ids.append(fasta.id)\n        sequences.append(str(fasta.seq))\n    return pd.DataFrame({'Id': ids, 'Sequence': sequences})\n\ndef embed_dataset(model, sequences, shift_left = 0, shift_right = -1):\n    \"\"\"copied from ankh examples. I do not know what special tokens are present, hopefully format is ok\n    Source: https://github.com/agemagician/Ankh/blob/main/examples/regression_fluorescence_task.ipynb\"\"\"\n    inputs_embedding = []\n    with torch.no_grad():\n        for sample in tqdm(sequences):\n            ids = tokenizer.batch_encode_plus([sample], add_special_tokens=True, \n                                              padding=True, is_split_into_words=True, \n                                              return_tensors=\"pt\",truncation=True,max_length=1010)\n            embedding = model(input_ids=ids['input_ids'].to(device))[0]\n            embedding = embedding[0].detach().cpu().numpy()[shift_left:shift_right]\n            inputs_embedding.append(embedding)\n    return inputs_embedding\n\ndef get_embed_cols(embeds,ID):\n    # Flatten the embeddings\n    flat_train_embed = [embedding[0].flatten() for embedding in embeds]\n    df = pd.DataFrame(flat_train_embed)\n    df.insert(0, \"EntryID\", ID)\n    df.columns = df.columns.astype(str)\n    return df","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:15:05.144026Z","iopub.execute_input":"2023-04-30T10:15:05.1444Z","iopub.status.idle":"2023-04-30T10:15:05.203332Z","shell.execute_reply.started":"2023-04-30T10:15:05.144356Z","shell.execute_reply":"2023-04-30T10:15:05.20223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# To load base model.\nmodel, tokenizer = ankh.load_base_model()","metadata":{"_kg_hide-output":true,"scrolled":true,"execution":{"iopub.status.busy":"2023-04-30T10:15:05.206887Z","iopub.execute_input":"2023-04-30T10:15:05.207653Z","iopub.status.idle":"2023-04-30T10:16:49.543703Z","shell.execute_reply.started":"2023-04-30T10:15:05.207612Z","shell.execute_reply":"2023-04-30T10:16:49.542636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model.half() # fails when using cpu - all nans\nmodel.eval()\nmodel.to(device=device)","metadata":{"scrolled":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-04-30T10:16:49.545121Z","iopub.execute_input":"2023-04-30T10:16:49.546138Z","iopub.status.idle":"2023-04-30T10:16:52.203811Z","shell.execute_reply.started":"2023-04-30T10:16:49.546093Z","shell.execute_reply":"2023-04-30T10:16:52.202792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_terms = pd.read_csv('/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv', sep='\\t')\ntrain_data = read_fasta('/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta')\n","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:18:29.421272Z","iopub.execute_input":"2023-04-30T10:18:29.42165Z","iopub.status.idle":"2023-04-30T10:18:32.770145Z","shell.execute_reply.started":"2023-04-30T10:18:29.421617Z","shell.execute_reply":"2023-04-30T10:18:32.768927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_data.shape[0])\ntrain_data.nunique()","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:18:32.77261Z","iopub.execute_input":"2023-04-30T10:18:32.773153Z","iopub.status.idle":"2023-04-30T10:18:33.030749Z","shell.execute_reply.started":"2023-04-30T10:18:32.773108Z","shell.execute_reply":"2023-04-30T10:18:33.028808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_data = read_fasta('/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset.fasta').drop_duplicates()\n# print(test_data.shape[0])\n# test_data.nunique() # there are duplicate sequences. We could drop them then join by ID, but do redundnat calcs for now","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:18:33.034162Z","iopub.execute_input":"2023-04-30T10:18:33.034522Z","iopub.status.idle":"2023-04-30T10:18:33.267321Z","shell.execute_reply.started":"2023-04-30T10:18:33.034495Z","shell.execute_reply":"2023-04-30T10:18:33.26638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:18:33.596694Z","iopub.execute_input":"2023-04-30T10:18:33.597407Z","iopub.status.idle":"2023-04-30T10:18:33.610224Z","shell.execute_reply.started":"2023-04-30T10:18:33.597368Z","shell.execute_reply":"2023-04-30T10:18:33.609039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(\"# IDs in test but not in train:\",len([i for i in test_data.Id if i not in train_data.Id]))","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:18:33.827015Z","iopub.execute_input":"2023-04-30T10:18:33.82772Z","iopub.status.idle":"2023-04-30T10:18:35.035288Z","shell.execute_reply.started":"2023-04-30T10:18:33.827682Z","shell.execute_reply":"2023-04-30T10:18:35.034063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## split seqs into character level; Ankh doesn't use multi aa tokenization?\nprotein_sequences_train = [list(seq) for seq in train_data.Sequence]\n# protein_sequences_test = [list(seq) for seq in test_data.Sequence]\nprint(protein_sequences_train[0][0:5])","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-04-30T10:18:35.03755Z","iopub.execute_input":"2023-04-30T10:18:35.037946Z","iopub.status.idle":"2023-04-30T10:18:39.060232Z","shell.execute_reply.started":"2023-04-30T10:18:35.037905Z","shell.execute_reply":"2023-04-30T10:18:39.058908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrain_embed = embed_dataset(model, protein_sequences_train)","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:18:39.062101Z","iopub.execute_input":"2023-04-30T10:18:39.062866Z","iopub.status.idle":"2023-04-30T10:19:23.255194Z","shell.execute_reply.started":"2023-04-30T10:18:39.062825Z","shell.execute_reply":"2023-04-30T10:19:23.254036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrain_df = get_embed_cols(embeds=train_embed,ID=train_data.Id)","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:19:23.257872Z","iopub.execute_input":"2023-04-30T10:19:23.258542Z","iopub.status.idle":"2023-04-30T10:19:23.279017Z","shell.execute_reply.started":"2023-04-30T10:19:23.2585Z","shell.execute_reply":"2023-04-30T10:19:23.278055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.to_parquet(\"cafa5_train_ankh_base_embed.parquet\")","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:19:23.280471Z","iopub.execute_input":"2023-04-30T10:19:23.28081Z","iopub.status.idle":"2023-04-30T10:19:23.298854Z","shell.execute_reply.started":"2023-04-30T10:19:23.280774Z","shell.execute_reply":"2023-04-30T10:19:23.297308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%time\n# del train_embed,train_df\n\n# test_embed = embed_dataset(model, protein_sequences_test)","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:19:23.300347Z","iopub.status.idle":"2023-04-30T10:19:23.300853Z","shell.execute_reply.started":"2023-04-30T10:19:23.300595Z","shell.execute_reply":"2023-04-30T10:19:23.300621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%time\n# test_df = get_embed_cols(embeds=test_embed,ID=test_data.Id)\n# test_df.to_parquet(\"cafa5_test_ankh_base_embed.parquet\")","metadata":{"execution":{"iopub.status.busy":"2023-04-30T10:19:23.302657Z","iopub.status.idle":"2023-04-30T10:19:23.303185Z","shell.execute_reply.started":"2023-04-30T10:19:23.302914Z","shell.execute_reply":"2023-04-30T10:19:23.302941Z"},"trusted":true},"execution_count":null,"outputs":[]}]}