{"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":"code","source":"MAIN_DIR = \"/kaggle/input/cafa-5-protein-function-prediction\"\n\n# UTILITARIES\nimport pandas as pd\nimport numpy as np\nfrom tqdm import tqdm\nimport time\nimport matplotlib.pyplot as plt\nplt.style.use('ggplot')\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\nimport pytorch_lightning as pl\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning.loggers import WandbLogger\n\n# WANDB FOR LIGHTNING :\nimport wandb\n\n# FILES VISUALIZATION\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-08T08:21:33.462651Z","iopub.execute_input":"2023-08-08T08:21:33.463046Z","iopub.status.idle":"2023-08-08T08:21:33.480289Z","shell.execute_reply.started":"2023-08-08T08:21:33.463009Z","shell.execute_reply":"2023-08-08T08:21:33.479257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Using an imported dataset and dictionary mapping to reduce preprocessing time for Y\n\nI have created my own Y dataset for up to 1500 labels and uploaded it to Kaggle, and then imported it into the input folder for this notebook.\n\n* These labels can be created from scratch, but it is **very CPU intensive** to do so.\n\n* Unfortunately, Kaggle does not allow you to save files directly to your input folder, but rather into a working folder.\n\n* For the purpose of submission within this contest, Kaggle is unable to use anything within your working folder, meaning **you would need to recreate them each time you want to submit**.\n\nAs a workaround, I created the Y array for the top 1500 most common GO terms, downloaded them piece by piece, and reuploaded them. This step reduces the preprocessing time necessary to create the labels from around 30 minutes to around 15 seconds while allowing us to work with a number of n_labels_to_consider anywhere between 0 and 1500.\n\nThe code below shows how to import the labels using a dictionary map, this will work in increments of 50, although you could always subset the labels if you wanted to use a more specific number.","metadata":{}},{"cell_type":"code","source":"n_labels_to_consider = 500\n\nlabel_bins = np.arange(0,n_labels_to_consider,50)\n\nlabel_dict = dict()\nfor n in label_bins:\n    label_dict[n] = np.load(f'/kaggle/input/cafa-5-1500-labels/labels_{n}.npy')\n\nY = np.zeros((142246, 50))\n    \nfor key, value in label_dict.items():\n    if Y.max() == 0:\n        Y = label_dict.get(key)\n    else:\n        Y = np.concatenate((Y, label_dict.get(key)), axis = 1)\n\nY.shape","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:21:33.482201Z","iopub.execute_input":"2023-08-08T08:21:33.482707Z","iopub.status.idle":"2023-08-08T08:21:35.064354Z","shell.execute_reply.started":"2023-08-08T08:21:33.482674Z","shell.execute_reply":"2023-08-08T08:21:35.063249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config:\n    train_sequences_path = MAIN_DIR  + \"/Train/train_sequences.fasta\"\n    train_labels_path = MAIN_DIR + \"/Train/train_terms.tsv\"\n    test_sequences_path = MAIN_DIR + \"/Test (Targets)/testsuperset.fasta\"\n    \n    num_labels = n_labels_to_consider\n    \n    n_epochs = 100\n  \n    batch_size = 128\n\n    lr = 0.003\n \n    \n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:21:35.06604Z","iopub.execute_input":"2023-08-08T08:21:35.066415Z","iopub.status.idle":"2023-08-08T08:21:35.072565Z","shell.execute_reply.started":"2023-08-08T08:21:35.066382Z","shell.execute_reply":"2023-08-08T08:21:35.071437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Directories for the different embedding vectors : \nembeds_map = {\n    \"T5\" : \"t5embeds\",\n    \"ProtBERT\" : \"protbert-embeddings-for-cafa5\",\n}\n\n# Length of the different embedding vectors :\nembeds_dim = {\n    \"T5\" : 1024,\n    \"ProtBERT\" : 1024,\n}","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:21:35.075697Z","iopub.execute_input":"2023-08-08T08:21:35.076061Z","iopub.status.idle":"2023-08-08T08:21:35.083829Z","shell.execute_reply.started":"2023-08-08T08:21:35.07603Z","shell.execute_reply":"2023-08-08T08:21:35.082645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset loader","metadata":{}},{"cell_type":"code","source":"class ProteinSequenceDataset(Dataset):\n    \n    def __init__(self, datatype, embeddings_source):\n        super(ProteinSequenceDataset).__init__()\n        self.datatype = datatype\n        \n        if embeddings_source == \"ProtBERT\":\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            np_labels = Y\n            df_labels = pd.DataFrame(self.df['EntryID'])\n            df_labels['labels_vect']=[row for row in np_labels]\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","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:21:35.085476Z","iopub.execute_input":"2023-08-08T08:21:35.085858Z","iopub.status.idle":"2023-08-08T08:21:35.09946Z","shell.execute_reply.started":"2023-08-08T08:21:35.085822Z","shell.execute_reply":"2023-08-08T08:21:35.098462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Architeture for LNN","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.baseNorm = torch.nn.BatchNorm1d(input_dim)\n        self.linear1 = torch.nn.Linear(input_dim, 864)\n        self.norm1 = torch.nn.LayerNorm(864, elementwise_affine=True)\n        self.activation1 = torch.nn.ReLU()\n        self.norm2 = torch.nn.BatchNorm1d(864)\n        self.linear2 = torch.nn.Linear(864, 712)\n        self.norm3 = torch.nn.LayerNorm(712, elementwise_affine=True)\n        self.activation2 = torch.nn.ReLU()\n        self.norm4 = torch.nn.BatchNorm1d(712)\n        self.dropout = torch.nn.Dropout(.3)\n        self.linear3 = torch.nn.Linear(712, num_classes)\n        self.sigmoid = torch.nn.Sigmoid()\n \n      \n\n    def forward(self, x):\n        x = self.baseNorm(x)\n        x = self.linear1(x)\n        x = self.norm1(x)\n        x = self.activation1(x)\n        x = self.norm2(x)\n        x = self.linear2(x)\n        x = self.norm3(x)\n        x = self.activation2(x)\n        x = self.norm4(x)\n        x = self.dropout(x)\n        x = self.linear3(x)\n        x = self.sigmoid(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:21:35.100796Z","iopub.execute_input":"2023-08-08T08:21:35.101103Z","iopub.status.idle":"2023-08-08T08:21:35.112829Z","shell.execute_reply.started":"2023-08-08T08:21:35.10108Z","shell.execute_reply":"2023-08-08T08:21:35.11167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Architecture for CNN","metadata":{}},{"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=864)\n        self.fc2 = nn.Linear(in_features=864, 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.tanh(self.conv1(x)))\n        x = self.pool2(nn.functional.tanh(self.conv2(x)))\n        x = torch.flatten(x, 1)\n        x = nn.functional.tanh(self.fc1(x))\n        x = self.fc2(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:21:35.114167Z","iopub.execute_input":"2023-08-08T08:21:35.115216Z","iopub.status.idle":"2023-08-08T08:21:35.126226Z","shell.execute_reply.started":"2023-08-08T08:21:35.115184Z","shell.execute_reply":"2023-08-08T08:21:35.125255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Training Parameters\n\nThis includes training size, early stopping criteria, learning rate reduction, and a checkpoint function for capturing the weights of our best performing model.  We are also using a Binary F1 score with a threshold which provides a more accurate F1 score than Multilabel for the purpose of this challenge (because we are dropping labels from consideration).","metadata":{}},{"cell_type":"code","source":"def train_model(embeddings_source, model_type=\"linear\", train_size=0.8):\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    from torchmetrics.classification import BinaryF1Score\n    optimizer = torch.optim.Adam(model.parameters(), lr = config.lr)\n    scheduler = ReduceLROnPlateau(optimizer, mode= 'min', factor=0.1, patience=2)\n    CrossEntropy = torch.nn.CrossEntropyLoss()\n    f1_score = BinaryF1Score(threshold=0.25, multidim_average= 'samplewise').to(config.device)\n#     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    \n    early_stop_thresh = 5\n    best_score = -1\n    best_epoch = -1\n    \n    def checkpoint(model, filename):\n        torch.save(model.state_dict(), filename)\n        \n    def resume(model, filename):\n        model.load_state_dict(torch.load(filename))\n    \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.mean().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.mean().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        if avg_score > best_score:\n            best_score = avg_score\n            best_epoch = epoch\n            checkpoint(model, \"best_model.pth\")\n        elif epoch - best_epoch > early_stop_thresh:\n            print(\"Early stopped training at epoch %d\" % epoch)\n            break  # terminate the training loop\n        \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    resume(model, \"best_model.pth\")\n    \n    return model, losses_history, scores_history","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:21:35.127691Z","iopub.execute_input":"2023-08-08T08:21:35.128047Z","iopub.status.idle":"2023-08-08T08:21:35.149754Z","shell.execute_reply.started":"2023-08-08T08:21:35.128014Z","shell.execute_reply":"2023-08-08T08:21:35.148758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Choose embeds and model type that we want to train:","metadata":{}},{"cell_type":"code","source":" t5_model, t5_losses, t5_scores = train_model(embeddings_source=\"T5\",model_type=\"linear\")","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:21:35.151052Z","iopub.execute_input":"2023-08-08T08:21:35.1515Z","iopub.status.idle":"2023-08-08T08:37:37.295191Z","shell.execute_reply.started":"2023-08-08T08:21:35.151464Z","shell.execute_reply":"2023-08-08T08:37:37.293255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#protbert_model, protbert_losses, protbert_scores = train_model(embeddings_source=\"ProtBERT\",model_type=\"linear\")","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:37:37.296373Z","iopub.status.idle":"2023-08-08T08:37:37.297063Z","shell.execute_reply.started":"2023-08-08T08:37:37.296791Z","shell.execute_reply":"2023-08-08T08:37:37.296814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize = (10, 4))\nplt.plot(t5_losses[\"val\"], label = \"T5\")\n#plt.plot(protbert_losses[\"val\"], label = \"ProtBERT\") \nplt.title(\"Validation Losses for # Vector Embeddings\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Average Loss\")\nplt.legend()\nplt.show()\n\nplt.figure(figsize = (10, 4))\nplt.plot(t5_scores[\"val\"], label = \"T5\")\n#plt.plot(protbert_scores[\"val\"], label = \"ProtBERT\")\nplt.title(\"Validation F1-Scores for # Vector Embeddings\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Average F1-Score\")\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:37:37.2987Z","iopub.status.idle":"2023-08-08T08:37:37.299161Z","shell.execute_reply.started":"2023-08-08T08:37:37.298912Z","shell.execute_reply":"2023-08-08T08:37:37.298935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Prediction","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\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-08T08:37:37.300146Z","iopub.status.idle":"2023-08-08T08:37:37.300792Z","shell.execute_reply.started":"2023-08-08T08:37:37.300533Z","shell.execute_reply":"2023-08-08T08:37:37.300559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = predict(\"T5\")","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:37:37.304082Z","iopub.status.idle":"2023-08-08T08:37:37.304896Z","shell.execute_reply.started":"2023-08-08T08:37:37.304657Z","shell.execute_reply":"2023-08-08T08:37:37.304679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"submission_df.to_csv('submission.tsv', sep='\\t', header=False, index=False)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T08:37:37.306198Z","iopub.status.idle":"2023-08-08T08:37:37.306788Z","shell.execute_reply.started":"2023-08-08T08:37:37.306552Z","shell.execute_reply":"2023-08-08T08:37:37.306574Z"},"trusted":true},"execution_count":null,"outputs":[]}]}