{"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":"# CAFA 5 protein function Prediction with TensorFlow\n\nThis notebook walks you through how to train a DNN model using TensorFlow on the CAFA 5 protein function Prediction dataset made available for this competition. \n\nThe objective of the model is to predict the function(aka **GO term ID**) of a set of proteins based on their amino acid sequences and other data.\n\n\n**Note** : This notebook runs without any GPU. This is because enabling GPUs leaves less RAM memory on the VM and the submission step needs a lot of memory. One point where this would impact is when training the model. With CPU it will take around 2 minutes while on GPU it would take around 30 seconds.","metadata":{"papermill":{"duration":0.008875,"end_time":"2023-05-09T08:30:09.052999","exception":false,"start_time":"2023-05-09T08:30:09.044124","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## About the Data\n\n### Protein Sequence\n\nEach protein is composed of dozens or hundreds of amino acids that are linked sequentially. Each amino acid in the sequence may be represented by a one-letter or three-letter code. Thus the sequence of a protein is often notated as a string of letters. \n\n<img src=\"https://cityu-bioinformatics.netlify.app/img/tools/protein/pro_seq.png\" alt =\"Sequence.png\" style='width: 800px;' >\n\nImage source - [https://cityu-bioinformatics.netlify.app/](https://cityu-bioinformatics.netlify.app/too2/new_proteo/pro_seq/)\n\nThe `train_sequences.fasta` made available for this competitions, contains the sequences for proteins with annotations (labelled proteins).","metadata":{"papermill":{"duration":0.009266,"end_time":"2023-05-09T08:30:09.071132","exception":false,"start_time":"2023-05-09T08:30:09.061866","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Gene Ontology\n\nWe can define the functional properties of a proteins using Gene Ontology(GO). Gene Ontology (GO) describes our understanding of the biological domain with respect to three aspects:\n1. Molecular Function (MF)\n2. Biological Process (BP)\n3. Cellular Component (CC)\n\nRead more about Gene Ontology [here](http://geneontology.org/docs/ontology-documentation).\n\nFile `train_terms.tsv` contains the list of annotated terms (ground truth) for the proteins in `train_sequences.fasta`. In `train_terms.tsv` the first column indicates the protein's UniProt accession ID (unique protein id), the second is the `GO Term ID`, and the third indicates in which ontology the term appears. ","metadata":{"papermill":{"duration":0.008355,"end_time":"2023-05-09T08:30:09.08863","exception":false,"start_time":"2023-05-09T08:30:09.080275","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Labels of the dataset\n\nThe objective of our model is to predict the terms (functions) of a protein sequence. One protein sequence can have many functions and can thus be classified into any number of terms. Each term is uniquely identified by a `GO Term ID`. Thus our model has to predict all the `GO Term ID`s for a protein sequence. This means that the task at hand is a multi-label classification problem. ","metadata":{}},{"cell_type":"markdown","source":"\n","metadata":{"papermill":{"duration":0.008308,"end_time":"2023-05-09T08:30:09.105539","exception":false,"start_time":"2023-05-09T08:30:09.097231","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Protein embeddings for train and test data\n\nTo train a machine learning model we cannot use the alphabetical protein sequences in`train_sequences.fasta` directly. They have to be converted into a vector format. In this notebook, we will use embeddings of the protein sequences to train the model. You can think of protein embeddings to be similar to word embeddings used to train NLP models.\n<!-- Instead, to make calculations and data preparation easier we will use precalculated protein embeddings.\n -->\nProtein embeddings are a machine-friendly method of capturing the protein's structural and functional characteristics, mainly through its sequence. One approach is to train a custom ML model to learn the protein embeddings of the protein sequences in the dataset being used in this notebook. Since this dataset represents proteins using amino-acid sequences which is a standard approach, we can use any publicly available pre-trained protein embedding models to generate the embeddings.\n\nThere are a variety of protein embedding models. To make data preparation easier, we have used the precalculated protein embeddings created by [Sergei Fironov](https://www.kaggle.com/sergeifironov) using the Rost Lab's T5 protein language model in this notebook. The precalculated protein embeddings can be found [here](https://www.kaggle.com/datasets/sergeifironov/t5embeds). We have added this dataset to the notebook along with the dataset made available for the competition.\n\nTo add this to your enviroment, on the right side panel, click on `Add Data` and search for `t5embeds` (make sure that it's the correct [one](https://www.kaggle.com/datasets/sergeifironov/t5embeds)) and then click on the `+` beside it.\n\n","metadata":{"papermill":{"duration":0.008548,"end_time":"2023-05-09T08:30:09.122603","exception":false,"start_time":"2023-05-09T08:30:09.114055","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Set up FoldSeek","metadata":{}},{"cell_type":"code","source":"!conda install -c conda-forge -c bioconda foldseek -y\n","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:54:48.16141Z","iopub.execute_input":"2023-08-02T08:54:48.162188Z","iopub.status.idle":"2023-08-02T08:55:37.64637Z","shell.execute_reply.started":"2023-08-02T08:54:48.16211Z","shell.execute_reply":"2023-08-02T08:55:37.644494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install git+https://github.com/SamusRam/ProFun.git\n","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:55:37.649888Z","iopub.execute_input":"2023-08-02T08:55:37.650388Z","iopub.status.idle":"2023-08-02T08:55:59.315254Z","shell.execute_reply.started":"2023-08-02T08:55:37.650343Z","shell.execute_reply":"2023-08-02T08:55:59.312973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from profun.models import FoldseekMatching, FoldseekConfig\nfrom profun.utils.project_info import ExperimentInfo","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:55:59.317773Z","iopub.execute_input":"2023-08-02T08:55:59.31838Z","iopub.status.idle":"2023-08-02T08:56:00.552564Z","shell.execute_reply.started":"2023-08-02T08:55:59.318328Z","shell.execute_reply":"2023-08-02T08:56:00.551391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import the Required Libraries","metadata":{"papermill":{"duration":0.009086,"end_time":"2023-05-09T08:30:09.140473","exception":false,"start_time":"2023-05-09T08:30:09.131387","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras import callbacks\nimport pandas as pd\nimport numpy as np\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm\nimport time\n\n# TORCH MODULES FOR METRICS COMPUTATION :\nimport torch\nfrom torch.utils.data import Dataset\nfrom torch import nn\nfrom torch.utils.data import random_split\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torchmetrics.classification import MultilabelF1Score\nfrom torchmetrics.classification import MultilabelAccuracy\n\n\n\n# Required for progressbar widget\nimport progressbar","metadata":{"papermill":{"duration":9.85331,"end_time":"2023-05-09T08:30:19.002985","exception":false,"start_time":"2023-05-09T08:30:09.149675","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-08-02T08:33:44.612134Z","iopub.execute_input":"2023-08-02T08:33:44.612618Z","iopub.status.idle":"2023-08-02T08:34:02.135396Z","shell.execute_reply.started":"2023-08-02T08:33:44.612584Z","shell.execute_reply":"2023-08-02T08:34:02.13391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"TensorFlow v\" + tf.__version__)\nprint(\"Numpy v\" + np.__version__)","metadata":{"papermill":{"duration":0.018272,"end_time":"2023-05-09T08:30:19.030432","exception":false,"start_time":"2023-05-09T08:30:19.01216","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-08-02T08:34:02.137798Z","iopub.execute_input":"2023-08-02T08:34:02.138254Z","iopub.status.idle":"2023-08-02T08:34:02.144621Z","shell.execute_reply.started":"2023-08-02T08:34:02.138219Z","shell.execute_reply":"2023-08-02T08:34:02.143449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Configuration Class","metadata":{"papermill":{"duration":0.008429,"end_time":"2023-05-09T08:30:19.047756","exception":false,"start_time":"2023-05-09T08:30:19.039327","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class config:\n    train_sequences_path = \"/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta\"\n    train_labels_path = \"/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv\"\n    test_sequences_path = \"/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset.fasta\"\n    \n    num_labels = 500\n    n_epochs = 5\n    batch_size = 128\n    lr = 0.001\n    \n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:34:02.146157Z","iopub.execute_input":"2023-08-02T08:34:02.147652Z","iopub.status.idle":"2023-08-02T08:34:02.161401Z","shell.execute_reply.started":"2023-08-02T08:34:02.147598Z","shell.execute_reply":"2023-08-02T08:34:02.160174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(config.device)\n","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:34:02.165139Z","iopub.execute_input":"2023-08-02T08:34:02.165654Z","iopub.status.idle":"2023-08-02T08:34:02.181343Z","shell.execute_reply.started":"2023-08-02T08:34:02.165594Z","shell.execute_reply":"2023-08-02T08:34:02.18003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CONTINUE RUNNING - Collect labels vectors for train/test","metadata":{}},{"cell_type":"code","source":"print(\"GENERATE TARGETS FOR ENTRY IDS (\"+str(config.num_labels)+\" MOST COMMON GO TERMS)\")\nids = np.load(\"/kaggle/input/t5embeds/train_ids.npy\")\nlabels = pd.read_csv(config.train_labels_path, sep = \"\\t\")\n\ntop_terms = labels.groupby(\"term\")[\"EntryID\"].count().sort_values(ascending=False)\nlabels_names = top_terms[:config.num_labels].index.values\ntrain_labels_sub = labels[(labels.term.isin(labels_names)) & (labels.EntryID.isin(ids))]\nid_labels = train_labels_sub.groupby('EntryID')['term'].apply(list).to_dict()\n\n# TODO - get clarification on the lines below\ngo_terms_map = {label: i for i, label in enumerate(labels_names)}\nlabels_matrix = np.empty((len(ids), len(labels_names)))\n\nfor index, id in tqdm(enumerate(ids)):\n    id_gos_list = id_labels[id]\n    temp = [go_terms_map[go] for go in labels_names if go in id_gos_list]\n    labels_matrix[index, temp] = 1\n\nlabels_list = []\nfor l in range(labels_matrix.shape[0]):\n    labels_list.append(labels_matrix[l, :])\n\nlabels_df = pd.DataFrame(data={\"EntryID\":ids, \"labels_vect\":labels_list})\nlabels_df.to_pickle(\"/kaggle/working/train_targets_top\"+str(config.num_labels)+\".pkl\")\nprint(\"GENERATION FINISHED!\")","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:34:02.1831Z","iopub.execute_input":"2023-08-02T08:34:02.183541Z","iopub.status.idle":"2023-08-02T08:35:29.852547Z","shell.execute_reply.started":"2023-08-02T08:34:02.183508Z","shell.execute_reply":"2023-08-02T08:35:29.850828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pytorch Dataset Architecture","metadata":{}},{"cell_type":"code","source":"# Directories for the different embedding vectors : \nembeds_map = {\n    \"T5\" : \"t5embeds\",\n    \"ProtBERT\" : \"protbert-embeddings-for-cafa5\",\n    \"EMS2\" : \"cafa-5-ems-2-embeddings-numpy\"\n}\n# todo - check why we need these diffferent embedding vectors\n\n# Length of the different embedding vectors :\nembeds_dim = {\n    \"T5\" : 1024,\n    \"ProtBERT\" : 1024,\n    \"EMS2\" : 1280\n}","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:35:29.854192Z","iopub.execute_input":"2023-08-02T08:35:29.854653Z","iopub.status.idle":"2023-08-02T08:35:29.861402Z","shell.execute_reply.started":"2023-08-02T08:35:29.854619Z","shell.execute_reply":"2023-08-02T08:35:29.860102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ProteinSequenceDataset(Dataset):\n    def __init__(self, datatype, embeddings_source):\n        super(ProteinSequenceDataset).__init__()\n        self.datatype = datatype\n        \n        # todo - so do the two if statements below indicate that the model can take embedings from ether source? if so why\n        # only difference seems to be be slightly different nmpy file loaded \n        if embeddings_source in [\"ProtBERT\", \"EMS2\"]:\n            embeds = np.load(\"/kaggle/input/\"+embeds_map[embeddings_source]+\"/\"+datatype+\"_embeddings.npy\")\n            ids = np.load(\"/kaggle/input/\"+embeds_map[embeddings_source]+\"/\"+datatype+\"_ids.npy\")\n        \n        if embeddings_source == \"T5\":\n            embeds = np.load(\"/kaggle/input/\"+embeds_map[embeddings_source]+\"/\"+datatype+\"_embeds.npy\")\n            ids = np.load(\"/kaggle/input/\"+embeds_map[embeddings_source]+\"/\"+datatype+\"_ids.npy\")\n            \n        embeds_list = []\n        for l in range(embeds.shape[0]):\n            embeds_list.append(embeds[l,:])\n        self.df = pd.DataFrame(data={\"EntryID\": ids, \"embed\" : embeds_list})\n        \n        if datatype==\"train\":\n            df_labels = pd.read_pickle(\n                \"/kaggle/working/train_targets_top\"+str(config.num_labels)+\".pkl\")\n            self.df = self.df.merge(df_labels, on=\"EntryID\")\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        embed = torch.tensor(self.df.iloc[index][\"embed\"] , dtype = torch.float32)\n        if self.datatype==\"train\":\n            targets = torch.tensor(self.df.iloc[index][\"labels_vect\"], dtype = torch.float32)\n            return embed, targets\n        if self.datatype==\"test\":\n            id = self.df.iloc[index][\"EntryID\"]\n            return embed, id\n     ","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:35:29.863614Z","iopub.execute_input":"2023-08-02T08:35:29.863986Z","iopub.status.idle":"2023-08-02T08:35:29.884693Z","shell.execute_reply.started":"2023-08-02T08:35:29.863956Z","shell.execute_reply":"2023-08-02T08:35:29.883641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = ProteinSequenceDataset(datatype=\"train\",embeddings_source=\"T5\")","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:35:29.886508Z","iopub.execute_input":"2023-08-02T08:35:29.888038Z","iopub.status.idle":"2023-08-02T08:35:42.806813Z","shell.execute_reply.started":"2023-08-02T08:35:29.887992Z","shell.execute_reply":"2023-08-02T08:35:42.805328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"embeddings, labels = dataset.__getitem__(0)\nprint(\"COMPONENTS FOR FIRST PROTEIN : \")\nprint(\"EMBEDDINGS VECTOR : \\n \", embeddings, \"\\n\")\nprint(\"TARGETS LABELS VECTOR : \\n \", labels, \"\\n\")","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:35:42.808579Z","iopub.execute_input":"2023-08-02T08:35:42.808953Z","iopub.status.idle":"2023-08-02T08:35:42.985229Z","shell.execute_reply.started":"2023-08-02T08:35:42.808924Z","shell.execute_reply":"2023-08-02T08:35:42.984277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pytorch Models Architectures","metadata":{}},{"cell_type":"code","source":"class MultiLayerPerceptron(torch.nn.Module):\n\n    def __init__(self, input_dim, num_classes):\n        super(MultiLayerPerceptron, self).__init__()\n\n        self.linear1 = torch.nn.Linear(input_dim, 1012)\n        self.activation1 = torch.nn.ReLU()\n        self.linear2 = torch.nn.Linear(1012, 712)\n        self.activation2 = torch.nn.ReLU()\n        self.linear3 = torch.nn.Linear(712, num_classes)\n\n    def forward(self, x):\n        x = self.linear1(x)\n        x = self.activation1(x)\n        x = self.linear2(x)\n        x = self.activation2(x)\n        x = self.linear3(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:35:42.988962Z","iopub.execute_input":"2023-08-02T08:35:42.989402Z","iopub.status.idle":"2023-08-02T08:35:42.998973Z","shell.execute_reply.started":"2023-08-02T08:35:42.989368Z","shell.execute_reply":"2023-08-02T08:35:42.997609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CNN1D(nn.Module):\n    def __init__(self, input_dim, num_classes):\n        super(CNN1D, self).__init__()\n        # (batch_size, channels, embed_size)\n        self.conv1 = nn.Conv1d(in_channels=1, out_channels=3, kernel_size=3, dilation=1, padding=1, stride=1)\n        # (batch_size, 3, embed_size)\n        self.pool1 = nn.MaxPool1d(kernel_size=2, stride=2)\n        # (batch_size, 3, embed_size/2 = 512)\n        self.conv2 = nn.Conv1d(in_channels=3, out_channels=8, kernel_size=3, dilation=1, padding=1, stride=1)\n        # (batch_size, 8, embed_size/2 = 512)\n        self.pool2 = nn.MaxPool1d(kernel_size=2, stride=2)\n        # (batch_size, 8, embed_size/4 = 256)\n        self.fc1 = nn.Linear(in_features=int(8 * input_dim/4), out_features=128)\n        self.fc2 = nn.Linear(in_features=128, out_features=num_classes)\n\n    def forward(self, x):\n        x = x.reshape(x.shape[0], 1, x.shape[1])\n        x = self.pool1(nn.functional.relu(self.conv1(x)))\n        x = self.pool2(nn.functional.relu(self.conv2(x)))\n        x = torch.flatten(x, 1)\n        x = nn.functional.relu(self.fc1(x))\n        x = self.fc2(x)\n        x = torch.sigmoid(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:35:43.000607Z","iopub.execute_input":"2023-08-02T08:35:43.000979Z","iopub.status.idle":"2023-08-02T08:35:43.021014Z","shell.execute_reply.started":"2023-08-02T08:35:43.000945Z","shell.execute_reply":"2023-08-02T08:35:43.019796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train the Model","metadata":{}},{"cell_type":"code","source":"def train_model(embeddings_source, model_type=\"linear\", train_size=0.9):\n    \n    train_dataset = ProteinSequenceDataset(datatype=\"train\", embeddings_source = embeddings_source)\n    \n    train_set, val_set = random_split(train_dataset, lengths = [int(len(train_dataset)*train_size), len(train_dataset)-int(len(train_dataset)*train_size)])\n    train_dataloader = torch.utils.data.DataLoader(train_set, batch_size=config.batch_size, shuffle=True)\n    val_dataloader = torch.utils.data.DataLoader(val_set, batch_size=config.batch_size, shuffle=True)\n\n    if model_type == \"linear\":\n        model = MultiLayerPerceptron(input_dim=embeds_dim[embeddings_source], num_classes=config.num_labels).to(config.device)\n    if model_type == \"convolutional\":\n        model = CNN1D(input_dim=embeds_dim[embeddings_source], num_classes=config.num_labels).to(config.device)\n\n    optimizer = torch.optim.Adam(model.parameters(), lr = config.lr)\n    scheduler = ReduceLROnPlateau(optimizer, factor=0.1, patience=1)\n    CrossEntropy = torch.nn.CrossEntropyLoss()\n    f1_score = MultilabelF1Score(num_labels=config.num_labels).to(config.device)\n    n_epochs = config.n_epochs\n\n    print(\"BEGIN TRAINING...\")\n    train_loss_history=[]\n    val_loss_history=[]\n    \n    train_f1score_history=[]\n    val_f1score_history=[]\n    for epoch in range(n_epochs):\n        print(\"EPOCH \", epoch+1)\n        ## TRAIN PHASE :\n        losses = []\n        scores = []\n        for embed, targets in tqdm(train_dataloader):\n            embed, targets = embed.to(config.device), targets.to(config.device)\n            optimizer.zero_grad()\n            preds = model(embed)\n            loss= CrossEntropy(preds, targets)\n            score=f1_score(preds, targets)\n            losses.append(loss.item()) \n            scores.append(score.item())\n            loss.backward()\n            optimizer.step()\n        avg_loss = np.mean(losses)\n        avg_score = np.mean(scores)\n        print(\"Running Average TRAIN Loss : \", avg_loss)\n        print(\"Running Average TRAIN F1-Score : \", avg_score)\n        train_loss_history.append(avg_loss)\n        train_f1score_history.append(avg_score)\n        \n        ## VALIDATION PHASE : \n        losses = []\n        scores = []\n        for embed, targets in val_dataloader:\n            embed, targets = embed.to(config.device), targets.to(config.device)\n            preds = model(embed)\n            loss= CrossEntropy(preds, targets)\n            score=f1_score(preds, targets)\n            losses.append(loss.item())\n            scores.append(score.item())\n        avg_loss = np.mean(losses)\n        avg_score = np.mean(scores)\n        print(\"Running Average VAL Loss : \", avg_loss)\n        print(\"Running Average VAL F1-Score : \", avg_score)\n        val_loss_history.append(avg_loss)\n        val_f1score_history.append(avg_score)\n        \n        scheduler.step(avg_loss)\n        print(\"\\n\")\n        \n    print(\"TRAINING FINISHED\")\n    print(\"FINAL TRAINING SCORE : \", train_f1score_history[-1])\n    print(\"FINAL VALIDATION SCORE : \", val_f1score_history[-1])\n    \n    losses_history = {\"train\" : train_loss_history, \"val\" : val_loss_history}\n    scores_history = {\"train\" : train_f1score_history, \"val\" : val_f1score_history}\n    \n    return model, losses_history, scores_history","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:35:43.023589Z","iopub.execute_input":"2023-08-02T08:35:43.024526Z","iopub.status.idle":"2023-08-02T08:35:43.049306Z","shell.execute_reply.started":"2023-08-02T08:35:43.024481Z","shell.execute_reply":"2023-08-02T08:35:43.047544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Example","metadata":{}},{"cell_type":"code","source":"T5_model, T5_losses, T5_scores = train_model(embeddings_source=\"T5\",model_type=\"linear\")\n","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:35:43.051038Z","iopub.execute_input":"2023-08-02T08:35:43.052254Z","iopub.status.idle":"2023-08-02T08:41:13.088007Z","shell.execute_reply.started":"2023-08-02T08:35:43.052212Z","shell.execute_reply":"2023-08-02T08:41:13.086743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# todo - so it looks like he trains the model useing two different embeddings, protbert and ems2 - presumably to compare \n# from his results it seems T5 gets the best results (or worse if I'm reading the scores the wrong way round)\n# also seems like he never actually uses the CNN model - both training examples use MLP","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:41:13.09016Z","iopub.execute_input":"2023-08-02T08:41:13.090547Z","iopub.status.idle":"2023-08-02T08:41:13.096376Z","shell.execute_reply.started":"2023-08-02T08:41:13.090514Z","shell.execute_reply":"2023-08-02T08:41:13.095076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make prediction using model","metadata":{}},{"cell_type":"code","source":"def predict(embeddings_source):\n    \n    test_dataset = ProteinSequenceDataset(datatype=\"test\", embeddings_source = embeddings_source)\n    test_dataloader = torch.utils.data.DataLoader(test_dataset, batch_size=1, shuffle=False)\n    \n    if embeddings_source == \"T5\":\n        model = T5_model\n    if embeddings_source == \"ProtBERT\":\n        model = protbert_model\n    if embeddings_source == \"EMS2\":\n        model = ems2_model\n        \n    model.eval()\n    \n    labels = pd.read_csv(config.train_labels_path, sep = \"\\t\")\n    top_terms = labels.groupby(\"term\")[\"EntryID\"].count().sort_values(ascending=False)\n    labels_names = top_terms[:config.num_labels].index.values\n    print(\"GENERATE PREDICTION FOR TEST SET...\")\n\n    ids_ = np.empty(shape=(len(test_dataloader)*config.num_labels,), dtype=object)\n    go_terms_ = np.empty(shape=(len(test_dataloader)*config.num_labels,), dtype=object)\n    confs_ = np.empty(shape=(len(test_dataloader)*config.num_labels,), dtype=np.float32)\n    # todo - not quite clear on this for loop\n    for i, (embed, id) in tqdm(enumerate(test_dataloader)):\n        embed = embed.to(config.device)\n        confs_[i*config.num_labels:(i+1)*config.num_labels] = torch.nn.functional.sigmoid(model(embed)).squeeze().detach().cpu().numpy()\n        ids_[i*config.num_labels:(i+1)*config.num_labels] = id[0]\n        go_terms_[i*config.num_labels:(i+1)*config.num_labels] = labels_names\n\n    submission_df = pd.DataFrame(data={\"Id\" : ids_, \"GO term\" : go_terms_, \"Confidence\" : confs_})\n    print(\"PREDICTIONS DONE\")\n    return submission_df","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:41:13.098427Z","iopub.execute_input":"2023-08-02T08:41:13.100023Z","iopub.status.idle":"2023-08-02T08:41:13.116935Z","shell.execute_reply.started":"2023-08-02T08:41:13.09997Z","shell.execute_reply":"2023-08-02T08:41:13.115388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = predict(\"T5\")","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:41:13.11907Z","iopub.execute_input":"2023-08-02T08:41:13.119757Z","iopub.status.idle":"2023-08-02T08:43:27.189868Z","shell.execute_reply.started":"2023-08-02T08:41:13.119721Z","shell.execute_reply":"2023-08-02T08:43:27.1886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.head(50)","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:43:27.191804Z","iopub.execute_input":"2023-08-02T08:43:27.193519Z","iopub.status.idle":"2023-08-02T08:43:27.230251Z","shell.execute_reply.started":"2023-08-02T08:43:27.193462Z","shell.execute_reply":"2023-08-02T08:43:27.228741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(submission_df)\n","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:43:27.231796Z","iopub.execute_input":"2023-08-02T08:43:27.232168Z","iopub.status.idle":"2023-08-02T08:43:27.241732Z","shell.execute_reply.started":"2023-08-02T08:43:27.232139Z","shell.execute_reply":"2023-08-02T08:43:27.240329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.to_csv('submission.tsv', sep='\\t', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-08-02T08:43:27.243663Z","iopub.execute_input":"2023-08-02T08:43:27.244044Z","iopub.status.idle":"2023-08-02T08:49:00.107318Z","shell.execute_reply.started":"2023-08-02T08:43:27.244014Z","shell.execute_reply":"2023-08-02T08:49:00.105018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}