{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","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"},{"sourceId":5499219,"sourceType":"datasetVersion","datasetId":3167603},{"sourceId":5549164,"sourceType":"datasetVersion","datasetId":3197305},{"sourceId":5607816,"sourceType":"datasetVersion","datasetId":3225525}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **Protein Function Prediction Part I: Data Preparation**","metadata":{}},{"cell_type":"code","source":"%%bash\n\npip3 install obonet pyvis networkx transformers torchmetrics torchsummary sentencepiece psutil biopython scikit-multilearn seaborn","metadata":{"execution":{"iopub.status.busy":"2024-03-08T16:09:29.075199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nimport io\nimport joblib\nimport json\nimport os\nimport pickle\nimport re\nimport shutil\nimport time\nimport torch\nimport typing\n\nimport matplotlib.pyplot as plt\nimport networkx\nimport numpy as np\nimport obonet\nimport pandas as pd\nimport seaborn as sns\nimport tqdm\n\nfrom Bio import SeqIO\nfrom pyvis.network import Network","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Directory Setup**","metadata":{}},{"cell_type":"code","source":"base_dir = \"/kaggle/working\"\ndata_dir = \"/dataset\"\nsub_data_dirs = [\n    \"processed/pb\",\n    \"processed/pb/train\",\n    \"processed/pb/test\",\n    \"processed/esm2\",\n    \"processed/esm2/train\",\n    \"processed/esm2/test\",\n    \"processed/t5\",\n    \"processed/t5/train\",\n    \"processed/t5/test\",\n    \"prepared/multi_target\",\n    \"prepared/multi_target/esm2\",\n    \"prepared/multi_target/pb\",\n    \"prepared/multi_target/t5\",\n]\n\nfor sub_dir in sub_data_dirs:\n    dir_path = os.path.join(base_dir+data_dir, sub_dir)\n    if not os.path.exists(dir_path):\n        os.makedirs(dir_path)\n        print(f\"Directory '{dir_path}' created successfully.\")\n    else:\n        print(f\"Directory '{dir_path}' already exists.\")","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:14:52.761613Z","iopub.execute_input":"2024-03-08T10:14:52.762322Z","iopub.status.idle":"2024-03-08T10:14:52.773462Z","shell.execute_reply.started":"2024-03-08T10:14:52.76228Z","shell.execute_reply":"2024-03-08T10:14:52.771994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Directory Path Configuration**","metadata":{}},{"cell_type":"code","source":"class Config:\n    def __init__(self):\n        self.num_labels = 1500\n\n        # Main Raw Dataset\n        self.go = \"/kaggle/input/cafa-5-protein-function-prediction/Train/go-basic.obo\"\n        self.test_prot_seq = \"/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset.fasta\"\n        self.test_taxon = \"/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset-taxon-list.tsv\"\n        self.train_prot_seq = \"/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta\"\n        self.train_terms = \"/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv\"\n        self.train_taxon = \"/kaggle/input/cafa-5-protein-function-prediction/Train/train_taxonomy.tsv\"\n\n        # Extras\n        self.valuation_weight = \"/kaggle/input/cafa-5-protein-function-prediction/IA.txt\"\n        self.ssub = \"/kaggle/input/cafa-5-protein-function-prediction/sample_submission.tsv\"\n\n        # Pre-Embedded Datasets\n        self.esm2_train_seq = \"/kaggle/input/cafa-5-ems-2-embeddings-numpy/train_embeddings.npy\"\n        self.esm2_train_id = \"/kaggle/input/cafa-5-ems-2-embeddings-numpy/train_ids.npy\"\n        self.esm2_test_seq = \"/kaggle/input/cafa-5-ems-2-embeddings-numpy/test_embeddings.npy\"\n        self.esm2_test_id = \"/kaggle/input/cafa-5-ems-2-embeddings-numpy/test_ids.npy\"\n        self.pb_train_seq = \"/kaggle/input/protbert-embeddings-for-cafa5/train_embeddings.npy\"\n        self.pb_train_id = \"/kaggle/input/protbert-embeddings-for-cafa5/train_ids.npy\"\n        self.pb_test_seq = \"/kaggle/input/protbert-embeddings-for-cafa5/test_embeddings.npy\"\n        self.pb_test_id = \"/kaggle/input/protbert-embeddings-for-cafa5/test_ids.npy\"\n        self.t5_train_seq = \"/kaggle/input/t5embeds/train_embeds.npy\"\n        self.t5_train_id = \"/kaggle/input/t5embeds/train_ids.npy\"\n        self.t5_test_seq = \"/kaggle/input/t5embeds/test_embeds.npy\"\n        self.t5_test_id = \"/kaggle/input/t5embeds/test_ids.npy\"\n\n        # Self-Processed Datasets\n        self.pro_pb_train = \"processed/pb/train\"\n        self.pro_pb_test = \"processed/pb/test\"\n        self.pro_pb_train = \"processed/esm2/train\"\n        self.pro_pb_test = \"processed/esm2/test\"\n        self.pro_t5_train = \"processed/t5/train\"\n        self.pro_t5_test = \"processed/t5/test\"\n        self.pre_esm2 = \"prepared/multi_target/esm2\"\n        self.pre_pb = \"prepared/multi_target/pb\"\n        self.pre_t5 = \"prepared/multi_target/t5\"\n\nCFG = Config()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:14:52.776676Z","iopub.execute_input":"2024-03-08T10:14:52.777595Z","iopub.status.idle":"2024-03-08T10:14:52.789161Z","shell.execute_reply.started":"2024-03-08T10:14:52.777551Z","shell.execute_reply":"2024-03-08T10:14:52.78811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **1. Data Understanding**","metadata":{}},{"cell_type":"markdown","source":"## **1.1 Data Collection**\n\nAll datasets are provided by the CAFA5 challenge which is organized by the Function Community of Special Interest (Function-COSI) which you can find [here](https://www.kaggle.com/competitions/cafa-5-protein-function-prediction).","metadata":{}},{"cell_type":"markdown","source":"## **1.2 CAFA5 Dataset**\n\nAs a starting point, CAFA5 dataset consists of 7 modules:\n\n* `go-basic`.obo $\\rightarrow$ A dataset full of nodes of protein functions in Direct Acyclic Graph structure.\n* `train_sequences.fasta` $\\rightarrow$ A train dataset with amino acid sequences for model training.\n* `train_terms.tsv` $\\rightarrow$ A train dataset with labels (ground truth) for the amino acid sequences in `train_sequences.fasta`.\n* `testsuperset.fasta` $\\rightarrow$ A test dataset for the model to infer. This dataset consists of protein ID and protein function.\n\n* `IA.txt` $\\rightarrow$ A dataset that contains evaluation weights for each protein function in `testsuperset.fasta`.\n* `sample_submission.tsv` $\\rightarrow$ An example structure of the model inference for the submission which is insignificant for my bachelor's thesis.\n\nThis two datasets are optional since it holds 1 extra feature for possible further prediction which ties the GO annotated protein functions to the respective species.\n* `train_taxonomy.tsv` $\\rightarrow$ A train dataset with species feature.\n* `testsuperset-taxon-list.tsv` $\\rightarrow$ A test dataset for mode inference.","metadata":{}},{"cell_type":"markdown","source":"#### **Glossary**\n* **The `term` refers to the protein function name after the annotation by the Gene Ontology standard.**\n* **The `aspect` refers to the sub-ontology: BP, CC, and MF**","metadata":{}},{"cell_type":"markdown","source":"## **1.3 Goal**\n\nThe goal is to predict the function of a protein from the amino acid sequence of the protein. This will be clearer at the end of this Notebook, once we finish exploring and preparing the dataset.","metadata":{}},{"cell_type":"markdown","source":"# **2. EDA**","metadata":{}},{"cell_type":"code","source":"def read_data(\n    dataset: str,\n    dtype: typing.Literal[\"csv\", \"fasta\", \"obo\", \"tsv\"],\n    encoding: typing.Literal[\"utf-8\", \"ISO-8859-1\"] = \"utf-8\"\n) -> pd.DataFrame:\n    if dtype == \"csv\":\n        df = pd.read_csv(dataset)\n    elif dtype == \"fasta\":\n        sequences = []\n        for record in SeqIO.parse(dataset, \"fasta\"):\n            sequences.append(\n                {\n                    \"protein_id\": record.id,\n                    \"description\": record.description,\n                    \"amino_acid_sequence\": str(record.seq)\n                }\n            )\n        df = pd.DataFrame(sequences)\n    elif dtype == \"obo\":\n        graph = obonet.read_obo(dataset)\n        terms = []\n        for node_id, data in graph.nodes(data=True):\n            terms.append({\n                \"term\": node_id,\n                \"name\": data.get(\"name\"),\n                \"aspect\": data.get(\"namespace\"),\n                \"description\": data.get(\"def\")\n            })\n\n        return pd.DataFrame(terms)\n    elif dtype == \"tsv\":\n        try:\n            df = pd.read_csv(dataset, delimiter=\"\\t\", encoding=encoding)\n        except UnicodeDecodeError as _err:\n            print(_err)\n    return df","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:14:52.790588Z","iopub.execute_input":"2024-03-08T10:14:52.791933Z","iopub.status.idle":"2024-03-08T10:14:52.80551Z","shell.execute_reply.started":"2024-03-08T10:14:52.791888Z","shell.execute_reply":"2024-03-08T10:14:52.804374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **2.1 Train Samples**","metadata":{}},{"cell_type":"code","source":"train_sample_df = read_data(\n    dataset=CFG.train_prot_seq,\n    dtype=\"fasta\"\n)\n\ntrain_sample_df","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:14:52.806907Z","iopub.execute_input":"2024-03-08T10:14:52.807287Z","iopub.status.idle":"2024-03-08T10:14:56.553572Z","shell.execute_reply.started":"2024-03-08T10:14:52.807256Z","shell.execute_reply":"2024-03-08T10:14:56.552323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_sample_df.info()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:14:56.554851Z","iopub.execute_input":"2024-03-08T10:14:56.555222Z","iopub.status.idle":"2024-03-08T10:14:56.650125Z","shell.execute_reply.started":"2024-03-08T10:14:56.555191Z","shell.execute_reply":"2024-03-08T10:14:56.648954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_sample_df.nunique()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:14:56.651687Z","iopub.execute_input":"2024-03-08T10:14:56.652095Z","iopub.status.idle":"2024-03-08T10:14:57.123221Z","shell.execute_reply.started":"2024-03-08T10:14:56.652047Z","shell.execute_reply":"2024-03-08T10:14:57.122011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\nHere we observethat the`train_sequences.fasta` contains:\n* 142246 rows\n* All rows contain unique value, hence\n* protein id - amino acid sequence are a unique pair","metadata":{}},{"cell_type":"code","source":"train_sample_df[\"sequence_length\"] = [len(seq) for seq in train_sample_df[\"amino_acid_sequence\"].values]\ntrain_sample_df","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:14:57.124532Z","iopub.execute_input":"2024-03-08T10:14:57.124935Z","iopub.status.idle":"2024-03-08T10:14:57.269886Z","shell.execute_reply.started":"2024-03-08T10:14:57.124904Z","shell.execute_reply":"2024-03-08T10:14:57.268989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp_df = train_sample_df.copy(deep=True)\ntop_50_most_lengthy_prot_seq = temp_df.sort_values(\"sequence_length\", ascending=False)[[\"protein_id\", \"sequence_length\"]].reset_index(drop=True).iloc[:50]\ntop_50_most_lengthy_prot_seq.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:14:57.2741Z","iopub.execute_input":"2024-03-08T10:14:57.275177Z","iopub.status.idle":"2024-03-08T10:14:57.368658Z","shell.execute_reply.started":"2024-03-08T10:14:57.275137Z","shell.execute_reply":"2024-03-08T10:14:57.367554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"figure, axis = plt.subplots(1, 1, figsize=(15, 8))\n\nbp = sns.barplot(ax=axis, x=top_50_most_lengthy_prot_seq[\"protein_id\"], y=top_50_most_lengthy_prot_seq[\"sequence_length\"])\nbp.set_xticklabels(bp.get_xticklabels(), rotation=75, size=8)\naxis.set_title(\"Top 50 Longest Amino Acid Sequences\", fontsize=15)\nbp.set_xlabel(\"Protein IDs\", fontsize=12)\nbp.set_ylabel(\"Sequence Length\", fontsize=12)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:14:57.370025Z","iopub.execute_input":"2024-03-08T10:14:57.370381Z","iopub.status.idle":"2024-03-08T10:14:58.240101Z","shell.execute_reply.started":"2024-03-08T10:14:57.370351Z","shell.execute_reply":"2024-03-08T10:14:58.23884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **2.2 Train Labels**","metadata":{}},{"cell_type":"code","source":"train_label_df = read_data(\n    dataset=CFG.train_terms,\n    dtype=\"tsv\"\n)\n\ntrain_label_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:14:58.241712Z","iopub.execute_input":"2024-03-08T10:14:58.242108Z","iopub.status.idle":"2024-03-08T10:15:02.556984Z","shell.execute_reply.started":"2024-03-08T10:14:58.242048Z","shell.execute_reply":"2024-03-08T10:15:02.555512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_label_df.columns = [\"protein_id\", \"protein_function\", \"sub_ontology\"]\ntrain_label_df","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:02.558819Z","iopub.execute_input":"2024-03-08T10:15:02.560054Z","iopub.status.idle":"2024-03-08T10:15:02.577135Z","shell.execute_reply.started":"2024-03-08T10:15:02.560007Z","shell.execute_reply":"2024-03-08T10:15:02.575799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Why is this a multi-label problem?**\n\nThis question will be answered with the following visualizations.","metadata":{}},{"cell_type":"markdown","source":"#### **How many functions are owned by the top 50 proteins based on its IDs?**","metadata":{}},{"cell_type":"code","source":"plot_df = train_label_df.groupby(\"protein_id\")[\"protein_function\"].count().sort_values(ascending=False).iloc[:50]\nfigure, axis = plt.subplots(1, 1, figsize=(15, 8))\n\nbp = sns.barplot(ax=axis, x=np.array(plot_df.index), y=plot_df.values)\nbp.set_xticklabels(bp.get_xticklabels(), rotation=75, size=8)\naxis.set_title(\"Top 50 Protein with Most Functions\", fontsize=15)\nbp.set_xlabel(\"Protein IDs\", fontsize=12)\nbp.set_ylabel(\"Function Count\", fontsize=12)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:02.578828Z","iopub.execute_input":"2024-03-08T10:15:02.580033Z","iopub.status.idle":"2024-03-08T10:15:05.274588Z","shell.execute_reply.started":"2024-03-08T10:15:02.579993Z","shell.execute_reply":"2024-03-08T10:15:05.27316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### **What are the top 50 most owned protein functions?**","metadata":{}},{"cell_type":"code","source":"plot_df = train_label_df.groupby(\"protein_function\")[\"protein_id\"].count().sort_values(ascending=False).iloc[:50]\nfigure, axis = plt.subplots(1, 1, figsize=(15, 8))\n\nbp = sns.barplot(ax=axis, x=np.array(plot_df.index), y=plot_df.values)\nbp.set_xticklabels(bp.get_xticklabels(), rotation=75, size=8)\naxis.set_title(\"Top 50 The Most Owned Protein Function\", fontsize=15)\nbp.set_xlabel(\"GO Annotated Protein Functions\", fontsize=12)\nbp.set_ylabel(\"Protein ID Count\", fontsize=12)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:05.276457Z","iopub.execute_input":"2024-03-08T10:15:05.276956Z","iopub.status.idle":"2024-03-08T10:15:07.405426Z","shell.execute_reply.started":"2024-03-08T10:15:05.276911Z","shell.execute_reply":"2024-03-08T10:15:07.404106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### **This is a \"Multi-Label\"problem because...**\nWe can observe from both visualizations above that Protein Functions and Proteins, represented by their unique ID, have the \"Many-to-Many\" relationship where 1 protein can own up to over 800 functions (see visualization 1) and 1 function can be owned by over 8000 proteins.","metadata":{}},{"cell_type":"markdown","source":"Let's see how many protein functions owned by each sub-ontology.","metadata":{}},{"cell_type":"code","source":"pie_df = train_label_df[\"sub_ontology\"].value_counts()\npalette_color = sns.color_palette(\"bright\")\n\nplt.pie(pie_df.values, labels=np.array(pie_df.index), colors=palette_color, autopct=\"%.0f%%\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:07.407165Z","iopub.execute_input":"2024-03-08T10:15:07.40757Z","iopub.status.idle":"2024-03-08T10:15:08.51196Z","shell.execute_reply.started":"2024-03-08T10:15:07.407537Z","shell.execute_reply":"2024-03-08T10:15:08.510709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"There are:\n* 65% of the protein functions belong to Biological Process Ontology (BPO),\n* 12% belong to Molecular Function Ontology (MFO), and\n* 22% belong to Cellular Component Ontology (CCO).","metadata":{}},{"cell_type":"markdown","source":"## **2.3 Test Dataset**","metadata":{}},{"cell_type":"code","source":"test_sample_df = read_data(\n    dataset=CFG.test_prot_seq,\n    dtype=\"fasta\"\n)\n\ntest_sample_df","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:08.513694Z","iopub.execute_input":"2024-03-08T10:15:08.515283Z","iopub.status.idle":"2024-03-08T10:15:11.638915Z","shell.execute_reply.started":"2024-03-08T10:15:08.515204Z","shell.execute_reply":"2024-03-08T10:15:11.637833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The test dataset consists of 3 features where `protein_id` and `amino_acid_sequence` exist but with a different description that resembles the `protein_id/taxon_id` which won't be used in this experimentation.","metadata":{}},{"cell_type":"markdown","source":"## **2.4 Direct Acyclic Graph**\n\nThis dataset will visualize the structure of the Gene Ontology data. The nodes in this graph (DAG) are indexed by the GO annotated protein function name. As a refresher, these are the 3 sub-ontologies:\n* Biological Process (BP)\n* Molecular Function (MF)\n* Cellular Component (CC)\n\nExample of the root graph in ontology data:","metadata":{}},{"cell_type":"code","source":"subontology_roots = {\n    \"BPO\": \"GO:0008150\",\n    \"CCO\": \"GO:0005575\",\n    \"MFO\": \"GO:0003674\"\n}","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:11.640629Z","iopub.execute_input":"2024-03-08T10:15:11.64097Z","iopub.status.idle":"2024-03-08T10:15:11.647268Z","shell.execute_reply.started":"2024-03-08T10:15:11.64094Z","shell.execute_reply":"2024-03-08T10:15:11.645998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_dag(graph, term, radius=1):\n    # create smaller subgraph\n    # radius - include all neighbors of distance<=radius from n (increse it to add further parent's branches).\n    ng_graph = networkx.ego_graph(graph, term, radius=radius)\n\n    for n in ng_graph.nodes(data=True):\n        # concatenate label of the node with its attribute\n        n[1][\"label\"] = n[0] + \" \" +n[1][\"name\"]\n\n    nt = Network(directed=True, notebook=True, neighborhood_highlight=True, cdn_resources=\"in_line\")\n    Network()\n\n    nt.from_nx(ng_graph)\n    return nt.show(\"network.html\")","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:11.64864Z","iopub.execute_input":"2024-03-08T10:15:11.649011Z","iopub.status.idle":"2024-03-08T10:15:11.658935Z","shell.execute_reply.started":"2024-03-08T10:15:11.648981Z","shell.execute_reply":"2024-03-08T10:15:11.657737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"graph = obonet.read_obo(CFG.go)\n\nprint(f\"Number of nodes: {len(graph)}\")\nprint(f\"Number of edges: {graph.number_of_edges()}\")","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:11.660659Z","iopub.execute_input":"2024-03-08T10:15:11.661223Z","iopub.status.idle":"2024-03-08T10:15:32.505337Z","shell.execute_reply.started":"2024-03-08T10:15:11.661182Z","shell.execute_reply":"2024-03-08T10:15:32.504103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's check the details from one random protein function. Since my lucky number is 3, Let's explore: \"GO:0003333\"","metadata":{}},{"cell_type":"code","source":"go_term = \"GO:0003333\"\ngraph.nodes[go_term]","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:32.506796Z","iopub.execute_input":"2024-03-08T10:15:32.507164Z","iopub.status.idle":"2024-03-08T10:15:32.515283Z","shell.execute_reply.started":"2024-03-08T10:15:32.507132Z","shell.execute_reply":"2024-03-08T10:15:32.514132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_dag(graph=graph, term=go_term, radius=1)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:32.51669Z","iopub.execute_input":"2024-03-08T10:15:32.517038Z","iopub.status.idle":"2024-03-08T10:15:32.845583Z","shell.execute_reply.started":"2024-03-08T10:15:32.516984Z","shell.execute_reply":"2024-03-08T10:15:32.844339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"In the graph above, the 2 nodes \"GO:0006865\" and \"GO:1905039\" are connected by relational ties between them through the property `is_a`, shown as the key in the dictionary of the go data. Let's increase the radius to 1000","metadata":{}},{"cell_type":"code","source":"plot_dag(graph=graph, term=go_term, radius=1000)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:32.847014Z","iopub.execute_input":"2024-03-08T10:15:32.847372Z","iopub.status.idle":"2024-03-08T10:15:33.170683Z","shell.execute_reply.started":"2024-03-08T10:15:32.847343Z","shell.execute_reply":"2024-03-08T10:15:33.169357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This graph shows a complete graph in which the connection is visualized with a non-peripheral relationship between all the nodes in the graph.","metadata":{}},{"cell_type":"markdown","source":"# **3. Data Processing**","metadata":{}},{"cell_type":"code","source":"model_name_and_dim = [\n    (\"facebook/esm2_t33_650M_UR50D\", 1280),\n    (\"Rostlab/prot_bert\", 1024),\n    (\"Rostlab/prot_t5_xl_half_uniref50-enc\", 1024)\n]","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:33.172323Z","iopub.execute_input":"2024-03-08T10:15:33.173037Z","iopub.status.idle":"2024-03-08T10:15:33.179567Z","shell.execute_reply.started":"2024-03-08T10:15:33.172999Z","shell.execute_reply":"2024-03-08T10:15:33.177929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ndevice","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:33.182036Z","iopub.execute_input":"2024-03-08T10:15:33.182832Z","iopub.status.idle":"2024-03-08T10:15:33.195877Z","shell.execute_reply.started":"2024-03-08T10:15:33.182778Z","shell.execute_reply":"2024-03-08T10:15:33.194329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_tokenizer_model(model, model_name: str):\n    return model.from_pretrained(model_name, legacy=False, do_lower_case=False)\n\ndef get_embedding_model(model, model_name: str):\n    embedding_model = model.from_pretrained(model_name, add_cross_attention=False, is_decoder=False).to(device)\n    embedding_model.eval()\n    return embedding_model\n\ndef get_embedding(tokenizer, model, sequence: str):\n    \"\"\"\n    Create embedding vector from pre-trained ESM2, ProtBert, or ProtT5.\n\n    Parameters\n    ----------\n    * sequence (str) : protein sequence (ex : AAABBB) from the training sequence dataset.\n    * model (Transformer Object) : pre-trained ESM2, ProtBert, or T5 embedding.\n    * tokenizer (Transformer Object) : pre-trained ESM2, ProtBert, or T5 tokenizer.\n\n    Returns\n    -------\n    * output_hidden : last hidden state embedding vector for input sequence of a certain length\n    \"\"\"\n    print(\"Starting embedding process. . .\")\n\n    sequence_examples = [\" \".join(list(re.sub(r\"[UZOB]\", \"X\", sequence)))]\n\n    ids = tokenizer(sequence_examples, add_special_tokens=True, padding=\"longest\")\n\n    input_ids = torch.tensor(ids[\"input_ids\"]).to(device)\n    attention_mask = torch.tensor(ids[\"attention_mask\"]).to(device)\n\n    print(\"Generating embedding. . .\")\n    with torch.no_grad():\n        embedding_repr = model(\n            input_ids=input_ids,\n            attention_mask=attention_mask\n        )\n\n    # Extract residue embeddings for the first ([0,:]) sequence in the batch and remove padded & special tokens ([0,:7])\n    emb_0 = embedding_repr.last_hidden_state[0]\n    emb_0_per_protein = emb_0.mean(dim=0)\n\n    print(\"Embedding process is successfully executed!\")\n    return emb_0_per_protein\n\ndef embed_dataset(\n    tokenizer,\n    model,\n    sequences,\n    sequence_size,\n    dir_path: str,\n    is_train: bool = True,\n    save_before_finish: bool = False,\n    save_at_iteration: int = 1\n):\n    file_prefix = \"train\" if is_train else \"test\"\n    ids = []\n    num_sequences = sum(1 for seq in sequences)\n    embeds = np.zeros((num_sequences, sequence_size))\n    i = 0\n\n    for seq in tqdm.tqdm(sequences):\n        ids.append(seq.id)\n        embeds[i] = get_embedding(tokenizer=tokenizer, model=model, sequence=str(seq.seq)).detach().cpu().numpy()\n        i += 1\n        if save_before_finish:\n            if i == save_at_iteration:\n                np.save(os.path.join(dir_path, f\"{file_prefix}_embeds.npy\"), np.array(embeds))\n                np.save(os.path.join(dir_path, f\"{file_prefix}_ids.npy\"), np.array(ids))\n    np.save(os.path.join(dir_path, f\"{file_prefix}_embeds.npy\"), np.array(embeds))\n    np.save(os.path.join(dir_path, f\"{file_prefix}_ids.npy\"), np.array(ids))\n\n    print(\"Dataset is successfully embedded!\")","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:33.197657Z","iopub.execute_input":"2024-03-08T10:15:33.198177Z","iopub.status.idle":"2024-03-08T10:15:33.21969Z","shell.execute_reply.started":"2024-03-08T10:15:33.198133Z","shell.execute_reply":"2024-03-08T10:15:33.2181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_sequences = SeqIO.parse(\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta\", \"fasta\")\nprint(\"Number of Sequences in Train:\", sum(1 for seq in train_sequences))\n\ntest_sequences = SeqIO.parse(\"/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset.fasta\", \"fasta\")\nprint(\"Number of Sequences in Test:\", sum(1 for seq in test_sequences))","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:33.22163Z","iopub.execute_input":"2024-03-08T10:15:33.2222Z","iopub.status.idle":"2024-03-08T10:15:37.570229Z","shell.execute_reply.started":"2024-03-08T10:15:33.222153Z","shell.execute_reply":"2024-03-08T10:15:37.567432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import BertModel, BertTokenizer, T5Tokenizer, T5EncoderModel, EsmModel, EsmTokenizer","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:37.577Z","iopub.execute_input":"2024-03-08T10:15:37.577565Z","iopub.status.idle":"2024-03-08T10:15:42.462518Z","shell.execute_reply.started":"2024-03-08T10:15:37.577526Z","shell.execute_reply":"2024-03-08T10:15:42.461181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#esm2_tokenizer = get_tokenizer_model(model=EsmTokenizer, model_name=model_name_and_dim[0][0])\n#esm2_embedding = get_embedding_model(model=EsmEncoderModel, model_name=model_name_and_dim[0][0])\n#bert_tokenizer = get_tokenizer_model(model=BertTokenizer, model_name=model_name_and_dim[1][0])\n#bert_embedding = get_embedding_model(model=BertModel, model_name=model_name_and_dim[1][0])\n#t5_tokenizer = get_tokenizer_model(model=T5Tokenizer, model_name=model_name_and_dim[2][0])\n#t5_embedding = get_embedding_model(model=T5EncoderModel, model_name=model_name_and_dim[2][0])","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:42.464238Z","iopub.execute_input":"2024-03-08T10:15:42.464962Z","iopub.status.idle":"2024-03-08T10:15:42.470704Z","shell.execute_reply.started":"2024-03-08T10:15:42.464923Z","shell.execute_reply":"2024-03-08T10:15:42.469815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **ATTENTION! Take 30min ++ to Finish!**\n\nIf you don't want to wait 30min ++ for the data processing, you can just download the precalculated datasets here:\n\n* ESM2 Embdedded Dataset: https://www.kaggle.com/datasets/viktorfairuschin/cafa-5-ems-2-embeddings-numpy\n* ProtBert Embedded Dataset: https://www.kaggle.com/datasets/henriupton/protbert-embeddings-for-cafa5\n* T5 EMbedded Dataset: https://www.kaggle.com/datasets/sergeifironov/t5embeds","metadata":{}},{"cell_type":"code","source":"train_sequences = list(SeqIO.parse(\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta\", \"fasta\"))","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:42.474879Z","iopub.execute_input":"2024-03-08T10:15:42.475319Z","iopub.status.idle":"2024-03-08T10:15:47.154677Z","shell.execute_reply.started":"2024-03-08T10:15:42.47528Z","shell.execute_reply":"2024-03-08T10:15:47.153204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#embed_dataset(\n#    tokenizer=bert_tokenizer,\n#    model=bert_embedding,\n#    sequences=train_sequences,\n#    sequence_size=model_name_and_dim[1][1],\n#    dir_path=base_dir+\"/\"+sub_dirs[1],\n#    is_train=True,\n#)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:47.156366Z","iopub.execute_input":"2024-03-08T10:15:47.157389Z","iopub.status.idle":"2024-03-08T10:15:47.162547Z","shell.execute_reply.started":"2024-03-08T10:15:47.157348Z","shell.execute_reply":"2024-03-08T10:15:47.161375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_sequences = list(SeqIO.parse(\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta\", \"fasta\"))","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:47.163866Z","iopub.execute_input":"2024-03-08T10:15:47.164267Z","iopub.status.idle":"2024-03-08T10:15:51.545419Z","shell.execute_reply.started":"2024-03-08T10:15:47.164224Z","shell.execute_reply":"2024-03-08T10:15:51.544193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#embed_dataset(\n#    tokenizer=bert_tokenizer,\n#    model=bert_embedding,\n#    sequences=test_sequences,\n#    sequence_size=model_name_and_dim[1][1],\n#    dir_path=base_dir+\"/\"+sub_dirs[1],\n#    is_train=True,\n#)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:51.546914Z","iopub.execute_input":"2024-03-08T10:15:51.547314Z","iopub.status.idle":"2024-03-08T10:15:51.552574Z","shell.execute_reply.started":"2024-03-08T10:15:51.547277Z","shell.execute_reply":"2024-03-08T10:15:51.55159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **4. Data Preparation**","metadata":{}},{"cell_type":"markdown","source":"## **4.1 Conventional Machine Learning**","metadata":{}},{"cell_type":"code","source":"labels = train_label_df[\"protein_function\"].value_counts().index[:CFG.num_labels].tolist()\nlabels[:3]","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:51.554091Z","iopub.execute_input":"2024-03-08T10:15:51.554698Z","iopub.status.idle":"2024-03-08T10:15:52.743213Z","shell.execute_reply.started":"2024-03-08T10:15:51.554664Z","shell.execute_reply":"2024-03-08T10:15:52.741716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp_df = train_label_df.copy(deep=True)\ntrain_label_1500_df = temp_df.loc[temp_df[\"protein_function\"].isin(labels)]\ntrain_label_1500_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:52.744769Z","iopub.execute_input":"2024-03-08T10:15:52.745221Z","iopub.status.idle":"2024-03-08T10:15:53.919641Z","shell.execute_reply.started":"2024-03-08T10:15:52.745175Z","shell.execute_reply":"2024-03-08T10:15:53.918194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def np_to_series(dataset: np.ndarray) -> pd.Series:\n    return pd.Series(dataset)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:53.921426Z","iopub.execute_input":"2024-03-08T10:15:53.921935Z","iopub.status.idle":"2024-03-08T10:15:53.928095Z","shell.execute_reply.started":"2024-03-08T10:15:53.921891Z","shell.execute_reply":"2024-03-08T10:15:53.92673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import progressbar","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:53.929657Z","iopub.execute_input":"2024-03-08T10:15:53.930158Z","iopub.status.idle":"2024-03-08T10:15:53.987617Z","shell.execute_reply.started":"2024-03-08T10:15:53.930117Z","shell.execute_reply":"2024-03-08T10:15:53.986488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_train_label_df(dataset) -> pd.DataFrame:\n    bar = progressbar.ProgressBar(\n        maxval=CFG.num_labels,\n        widgets=[\n            progressbar.Bar(\"=\", \"[\", \"]\"), \" \",progressbar.Percentage()\n        ]\n    )\n    train_size = dataset.shape[0]\n    train_labels = np.zeros((train_size, CFG.num_labels))\n    series = np_to_series(dataset=dataset)\n\n    for i in range(CFG.num_labels):\n        n_train_labels = train_label_1500_df[train_label_1500_df[\"protein_function\"] == labels[i]]\n        label_related_proteins = n_train_labels[\"protein_id\"].unique()\n        train_labels[:,i] = series.isin(label_related_proteins).astype(float)\n        bar.update(i+1)\n\n    bar.finish()\n    return pd.DataFrame(data=train_labels, columns=labels)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:53.988976Z","iopub.execute_input":"2024-03-08T10:15:53.990014Z","iopub.status.idle":"2024-03-08T10:15:54.001332Z","shell.execute_reply.started":"2024-03-08T10:15:53.989966Z","shell.execute_reply":"2024-03-08T10:15:54.000109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"esm2_embed_train_ids = np.load(CFG.esm2_train_id)\npb_embed_train_ids = np.load(CFG.pb_train_id)\nt5_embed_train_ids = np.load(CFG.t5_train_id)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:15:54.00333Z","iopub.execute_input":"2024-03-08T10:15:54.004078Z","iopub.status.idle":"2024-03-08T10:15:54.25619Z","shell.execute_reply.started":"2024-03-08T10:15:54.004025Z","shell.execute_reply":"2024-03-08T10:15:54.25442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def store_dataset(dataset: pd.DataFrame, model_name: str, stype: typing.Literal[\"train_label\", \"train_sample\", \"test_label\", \"test_sample\"]):\n    data_dirs = {\"esm2\": -3, \"pb\": -2, \"t5\": -1}\n    fn = f\"{model_name}_{stype}_row_{str(dataset.shape[0])}_feat_{str(dataset.shape[1])}.csv\"\n    fp = base_dir + data_dir + \"/\" + sub_data_dirs[data_dirs[model_name]] + \"/\" + fn\n    dataset.to_csv(path_or_buf=fp)\n    print(f\"Dataset is successfully stored in {fp}\")","metadata":{"execution":{"iopub.status.busy":"2024-03-08T09:03:49.699344Z","iopub.execute_input":"2024-03-08T09:03:49.699848Z","iopub.status.idle":"2024-03-08T09:03:49.708738Z","shell.execute_reply.started":"2024-03-08T09:03:49.699811Z","shell.execute_reply":"2024-03-08T09:03:49.707103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### **IMPORTANT**\n\nIn each embedding dataset, you need to either create the `train_label_df` or upload them depending whether you have run this notebook or not. One df creation can take up to 10 minutes.","metadata":{}},{"cell_type":"markdown","source":"#### **ESM2**","metadata":{}},{"cell_type":"markdown","source":"**Create DF**","metadata":{}},{"cell_type":"code","source":"esm2_train_label_df = create_train_label_df(dataset=esm2_embed_train_ids)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T08:18:02.41315Z","iopub.execute_input":"2024-03-08T08:18:02.413543Z","iopub.status.idle":"2024-03-08T08:40:52.804198Z","shell.execute_reply.started":"2024-03-08T08:18:02.413513Z","shell.execute_reply":"2024-03-08T08:40:52.802135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"esm2_train_label_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T08:49:52.25391Z","iopub.execute_input":"2024-03-08T08:49:52.254454Z","iopub.status.idle":"2024-03-08T08:49:52.302638Z","shell.execute_reply.started":"2024-03-08T08:49:52.254419Z","shell.execute_reply":"2024-03-08T08:49:52.300733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"store_dataset(\n    dataset=esm2_train_label_df,\n    model_name=\"esm2\",\n    stype=\"train_label\"\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"memory_usage: {esm2_train_label_df.memory_usage(index=True).sum()}\")\nesm2_train_label_df.describe()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T08:50:26.618409Z","iopub.execute_input":"2024-03-08T08:50:26.619618Z","iopub.status.idle":"2024-03-08T08:50:46.40611Z","shell.execute_reply.started":"2024-03-08T08:50:26.619548Z","shell.execute_reply":"2024-03-08T08:50:46.404892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Uplod DF**","metadata":{}},{"cell_type":"code","source":"esm2_train_label_df = pd.read_csv(base_dir + data_dir + \"/\" + sub_data_dirs[-3] + \"/\" + \"esm2_train_label_row_142246_feat_1500.csv\", index_col=0)\nesm2_train_label_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:39:56.08047Z","iopub.execute_input":"2024-03-08T10:39:56.080958Z","iopub.status.idle":"2024-03-08T10:40:29.997151Z","shell.execute_reply.started":"2024-03-08T10:39:56.080923Z","shell.execute_reply":"2024-03-08T10:40:29.99598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### **ProtBert**","metadata":{}},{"cell_type":"markdown","source":"**Create DF**","metadata":{}},{"cell_type":"code","source":"pb_train_label_df = create_train_label_df(dataset=pb_embed_train_ids)\npb_train_label_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T01:34:45.709833Z","iopub.status.idle":"2024-03-08T01:34:45.710187Z","shell.execute_reply.started":"2024-03-08T01:34:45.710013Z","shell.execute_reply":"2024-03-08T01:34:45.710027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"store_dataset(\n    dataset=pb_train_label_df,\n    model_name=\"pb\",\n    stype=\"train_label\"\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Uplod DF**","metadata":{}},{"cell_type":"code","source":"pb_train_label_df = pd.read_csv(base_dir + data_dir + \"/\" + sub_data_dirs[-2] + \"/\" + \"pb_train_label_row_142246_feat_1500.csv\", index_col=0)\npb_train_label_df.head(3)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### **ProtT5**","metadata":{}},{"cell_type":"markdown","source":"**Create DF**","metadata":{}},{"cell_type":"code","source":"t5_train_label_df = create_train_label_df(dataset=t5_embed_train_ids)\n\nt5_train_label_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T01:34:45.711486Z","iopub.status.idle":"2024-03-08T01:34:45.711904Z","shell.execute_reply.started":"2024-03-08T01:34:45.711688Z","shell.execute_reply":"2024-03-08T01:34:45.711706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"store_dataset(\n    dataset=t5_train_label_df,\n    model_name=\"t5\",\n    stype=\"train_label\"\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Uplod DF**","metadata":{}},{"cell_type":"code","source":"t5_train_label_df = pd.read_csv(base_dir + data_dir + \"/\" + sub_data_dirs[-1] + \"/\" + \"t5_train_label_row_142246_feat_1500.csv\", index_col=0)\nt5_train_label_df.head(3)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Deep Learning**","metadata":{}},{"cell_type":"markdown","source":"# **5. Model Development**\n\n* Logistic Regression\n* DecisionTree\n* KNeighbor\n* SVM\n* XGBoost\n* ANN\n* LSTM","metadata":{}},{"cell_type":"markdown","source":"## **3.2 Prepare X Train**","metadata":{}},{"cell_type":"markdown","source":"#### **ESM2**","metadata":{}},{"cell_type":"code","source":"esm2_train_prot_seq = np.load(CFG.esm2_train_seq)\nesm2_train_prot_seq","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:44:39.128182Z","iopub.execute_input":"2024-03-08T10:44:39.128709Z","iopub.status.idle":"2024-03-08T10:44:49.616255Z","shell.execute_reply.started":"2024-03-08T10:44:39.128668Z","shell.execute_reply":"2024-03-08T10:44:49.614983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"column_num = esm2_train_prot_seq.shape[1]\nesm2_train_samples_df = pd.DataFrame(esm2_train_prot_seq, columns = [\"Column_\" + str(i) for i in range(1, column_num+1)])\n\nesm2_train_samples_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:45:28.81885Z","iopub.execute_input":"2024-03-08T10:45:28.819941Z","iopub.status.idle":"2024-03-08T10:45:28.851002Z","shell.execute_reply.started":"2024-03-08T10:45:28.819895Z","shell.execute_reply":"2024-03-08T10:45:28.849648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"esm2_train_samples_df.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:45:39.433384Z","iopub.execute_input":"2024-03-08T10:45:39.433839Z","iopub.status.idle":"2024-03-08T10:45:39.442302Z","shell.execute_reply.started":"2024-03-08T10:45:39.433808Z","shell.execute_reply":"2024-03-08T10:45:39.440997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"esm2_train_samples_df.shape[0] == esm2_train_label_df.shape[0]","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:46:06.449883Z","iopub.execute_input":"2024-03-08T10:46:06.450393Z","iopub.status.idle":"2024-03-08T10:46:06.458622Z","shell.execute_reply.started":"2024-03-08T10:46:06.450357Z","shell.execute_reply":"2024-03-08T10:46:06.457148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**The above result confirm that each row in train samples point to each row in train labels.**","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\n\nINPUT_SHAPE = [esm2_train_samples_df.shape[1]]\nBATCH_SIZE = 5120\n\nmodel = tf.keras.Sequential([\n    tf.keras.layers.BatchNormalization(input_shape=INPUT_SHAPE),    \n    tf.keras.layers.Dense(units=512, activation='relu'),\n    tf.keras.layers.Dense(units=512, activation='relu'),\n    tf.keras.layers.Dense(units=512, activation='relu'),\n    tf.keras.layers.Dense(units=CFG.num_labels,activation='sigmoid')\n])\n\n\n# Compile model\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(learning_rate=0.001),\n    loss='binary_crossentropy',\n    metrics=['binary_accuracy', tf.keras.metrics.AUC()],\n)\n\nhistory = model.fit(\n    esm2_train_samples_df, esm2_train_label_df,\n    batch_size=BATCH_SIZE,\n    epochs=5\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:48:49.589188Z","iopub.execute_input":"2024-03-08T10:48:49.58966Z","iopub.status.idle":"2024-03-08T10:51:54.223756Z","shell.execute_reply.started":"2024-03-08T10:48:49.589625Z","shell.execute_reply":"2024-03-08T10:51:54.222528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history_df = pd.DataFrame(history.history)\nhistory_df.loc[:, [\"loss\"]].plot(title=\"Cross-entropy\")\nhistory_df.loc[:, [\"binary_accuracy\"]].plot(title=\"Accuracy\")","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:51:54.226389Z","iopub.execute_input":"2024-03-08T10:51:54.22694Z","iopub.status.idle":"2024-03-08T10:51:55.021724Z","shell.execute_reply.started":"2024-03-08T10:51:54.226903Z","shell.execute_reply":"2024-03-08T10:51:55.02015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"esm2_test_prot_seq = np.load(CFG.esm2_test_seq)\npb_test_prot_seq = np.load(CFG.pb_test_seq)\nt5_test_prot_seq = np.load(CFG.t5_test_seq)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T10:55:47.922971Z","iopub.execute_input":"2024-03-08T10:55:47.923489Z","iopub.status.idle":"2024-03-08T10:56:13.244187Z","shell.execute_reply.started":"2024-03-08T10:55:47.923451Z","shell.execute_reply":"2024-03-08T10:56:13.242809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"esm2_test_prot_id = np.load(CFG.esm2_test_id)\npb_test_prot_id = np.load(CFG.pb_test_id)\nt5_test_prot_id = np.load(CFG.t5_test_id)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T11:06:04.515346Z","iopub.execute_input":"2024-03-08T11:06:04.517466Z","iopub.status.idle":"2024-03-08T11:06:04.830462Z","shell.execute_reply.started":"2024-03-08T11:06:04.517405Z","shell.execute_reply":"2024-03-08T11:06:04.829404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"column_num = esm2_test_prot_seq.shape[1]\nesm2_test_sample_df = pd.DataFrame(esm2_test_prot_seq, columns = [\"Column_\" + str(i) for i in range(1, column_num+1)])\n\nesm2_test_sample_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-08T11:03:06.929671Z","iopub.execute_input":"2024-03-08T11:03:06.930462Z","iopub.status.idle":"2024-03-08T11:03:06.967465Z","shell.execute_reply.started":"2024-03-08T11:03:06.930384Z","shell.execute_reply":"2024-03-08T11:03:06.966106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"esm2_test_sample_df.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-08T11:03:25.288318Z","iopub.execute_input":"2024-03-08T11:03:25.288792Z","iopub.status.idle":"2024-03-08T11:03:25.295991Z","shell.execute_reply.started":"2024-03-08T11:03:25.28875Z","shell.execute_reply":"2024-03-08T11:03:25.294882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"esm2_dataset_pediction = model.predict(esm2_test_sample_df)","metadata":{"execution":{"iopub.status.busy":"2024-03-08T11:04:12.327904Z","iopub.execute_input":"2024-03-08T11:04:12.329209Z","iopub.status.idle":"2024-03-08T11:04:41.998703Z","shell.execute_reply.started":"2024-03-08T11:04:12.329156Z","shell.execute_reply":"2024-03-08T11:04:41.99762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"esm2_model_inference_df = pd.DataFrame(columns=[\"Protein ID\", \"Protein Function\",'Prediction'])\n\nl = []\nfor k in list(esm2_test_prot_id):\n    l += [k] * esm2_dataset_pediction.shape[1]   \n\nesm2_model_inference_df[\"Protein ID\"] = l\nesm2_model_inference_df[\"Protein Function\"] = labels * esm2_dataset_pediction.shape[0]\nesm2_model_inference_df[\"Prediction\"] = esm2_dataset_pediction.ravel()\nesm2_model_inference_df.head()\n#esm2_model_inference_df.to_csv(\"submission.tsv\",header=False, index=False, sep=\"\\t\")","metadata":{"execution":{"iopub.status.busy":"2024-03-08T11:07:29.230817Z","iopub.execute_input":"2024-03-08T11:07:29.232009Z","iopub.status.idle":"2024-03-08T11:08:40.896771Z","shell.execute_reply.started":"2024-03-08T11:07:29.231954Z","shell.execute_reply":"2024-03-08T11:08:40.895504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"esm2_model_inference_df[esm2_model_inference_df[\"Prediction\"].values > 0.7]","metadata":{"execution":{"iopub.status.busy":"2024-03-08T11:10:28.093564Z","iopub.execute_input":"2024-03-08T11:10:28.09405Z","iopub.status.idle":"2024-03-08T11:10:28.233846Z","shell.execute_reply.started":"2024-03-08T11:10:28.094011Z","shell.execute_reply":"2024-03-08T11:10:28.232371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### **ProtBert**","metadata":{}},{"cell_type":"code","source":"X_pb_train_protein_seq = np.load(CFG.pb_train_seq)\nX_pb_train_protein_seq","metadata":{"execution":{"iopub.status.busy":"2024-03-08T00:51:04.811892Z","iopub.execute_input":"2024-03-08T00:51:04.812379Z","iopub.status.idle":"2024-03-08T00:51:12.581827Z","shell.execute_reply.started":"2024-03-08T00:51:04.812344Z","shell.execute_reply":"2024-03-08T00:51:12.580511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_pb_train_protein_ids = np.load(CFG.pb_train_id)\nX_pb_train_protein_ids","metadata":{"execution":{"iopub.status.busy":"2024-03-08T00:51:17.888595Z","iopub.execute_input":"2024-03-08T00:51:17.889252Z","iopub.status.idle":"2024-03-08T00:51:17.985102Z","shell.execute_reply.started":"2024-03-08T00:51:17.889218Z","shell.execute_reply":"2024-03-08T00:51:17.983847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(set(X_pb_train_protein_ids) & set(train_label_df[\"protein_id\"])), len(X_pb_train_protein_ids) ) ","metadata":{"execution":{"iopub.status.busy":"2024-03-08T00:51:30.091489Z","iopub.execute_input":"2024-03-08T00:51:30.091909Z","iopub.status.idle":"2024-03-08T00:51:30.748524Z","shell.execute_reply.started":"2024-03-08T00:51:30.091881Z","shell.execute_reply":"2024-03-08T00:51:30.74759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### **ProtT5**","metadata":{}},{"cell_type":"code","source":"X_t5_train_protein_seq = np.load(CFG.t5_train_seq)\nX_t5_train_protein_seq","metadata":{"execution":{"iopub.status.busy":"2024-03-08T00:51:47.421008Z","iopub.execute_input":"2024-03-08T00:51:47.421502Z","iopub.status.idle":"2024-03-08T00:51:59.793061Z","shell.execute_reply.started":"2024-03-08T00:51:47.421465Z","shell.execute_reply":"2024-03-08T00:51:59.791773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_t5_train_protein_ids = np.load(CFG.t5_train_id)\nX_t5_train_protein_ids","metadata":{"execution":{"iopub.status.busy":"2024-03-08T00:51:59.81297Z","iopub.execute_input":"2024-03-08T00:51:59.813502Z","iopub.status.idle":"2024-03-08T00:51:59.882105Z","shell.execute_reply.started":"2024-03-08T00:51:59.813424Z","shell.execute_reply":"2024-03-08T00:51:59.880746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(set(X_t5_train_protein_ids) & set(train_label_df[\"protein_id\"])), len(X_t5_train_protein_ids) ) ","metadata":{"execution":{"iopub.status.busy":"2024-03-08T00:52:11.930462Z","iopub.execute_input":"2024-03-08T00:52:11.930926Z","iopub.status.idle":"2024-03-08T00:52:12.592796Z","shell.execute_reply.started":"2024-03-08T00:52:11.930895Z","shell.execute_reply":"2024-03-08T00:52:12.591507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"markdown","source":"# **Conclusion**\n\nNext step is to develop the models in Notebook Part II.","metadata":{}}]}