{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This code is based on \"[ProteiNet 🧬 PyTorch+EMS2/T5/ProtBERT Embeddings](https://www.kaggle.com/code/henriupton/proteinet-pytorch-ems2-t5-protbert-embeddings)\". Thank you for sharing [@Henri Upton](https://www.kaggle.com/henriupton)\n\nI change his code simlpy. I extract ESM-2 (3B) embeddings with three pooling methods ([CLS],sum, mean). And try to predict performance. ","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nsub = pd.read_csv(\"/kaggle/input/cafa-5-protein-function-prediction/sample_submission.tsv\", sep= \"\\t\", header = None)\nsub.columns = [\"The Protein ID\", \"The Gene Ontology term (GO) ID\", \"Predicted link probability that GO appear in Protein\"]\nsub.head(5)","metadata":{"execution":{"iopub.status.busy":"2023-05-21T06:27:38.389504Z","iopub.execute_input":"2023-05-21T06:27:38.390161Z","iopub.status.idle":"2023-05-21T06:27:38.651322Z","shell.execute_reply.started":"2023-05-21T06:27:38.390127Z","shell.execute_reply":"2023-05-21T06:27:38.650359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MAIN_DIR = \"/kaggle/input/cafa-5-protein-function-prediction\"\n\n# UTILITARIES\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":{"execution":{"iopub.status.busy":"2023-05-21T06:27:38.653499Z","iopub.execute_input":"2023-05-21T06:27:38.65389Z","iopub.status.idle":"2023-05-21T06:28:03.908479Z","shell.execute_reply.started":"2023-05-21T06:27:38.653855Z","shell.execute_reply":"2023-05-21T06:28:03.907341Z"},"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 = 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-05-21T06:28:03.909933Z","iopub.execute_input":"2023-05-21T06:28:03.910514Z","iopub.status.idle":"2023-05-21T06:28:03.931751Z","shell.execute_reply.started":"2023-05-21T06:28:03.91048Z","shell.execute_reply":"2023-05-21T06:28:03.930852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(config.device)","metadata":{"execution":{"iopub.status.busy":"2023-05-21T06:28:03.934502Z","iopub.execute_input":"2023-05-21T06:28:03.934839Z","iopub.status.idle":"2023-05-21T06:28:03.93971Z","shell.execute_reply.started":"2023-05-21T06:28:03.934809Z","shell.execute_reply":"2023-05-21T06:28:03.938745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"GENERATE TARGETS FOR ENTRY IDS (\"+str(config.num_labels)+\" MOST COMMON GO TERMS)\")\nids = np.load(\"/kaggle/input/cafa5-esm-2-3b-various-poolings/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\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\nnp.save(\"/kaggle/working/train_targets_top\"+str(config.num_labels)+\".npy\", np.array(labels_matrix))\nprint(\"GENERATION FINISHED!\")","metadata":{"execution":{"iopub.status.busy":"2023-05-21T06:28:03.941069Z","iopub.execute_input":"2023-05-21T06:28:03.942111Z","iopub.status.idle":"2023-05-21T06:29:19.452656Z","shell.execute_reply.started":"2023-05-21T06:28:03.942081Z","shell.execute_reply":"2023-05-21T06:29:19.451572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Directories for the different embedding vectors : \nembeds_map = {\n    'ids' : \"cafa5-esm-2-3b-various-poolings/\",\n    \"cls\" : \"cafa5-esm-2-3b-various-poolings/cls\",\n    \"sum\" : \"cafa5-esm-2-3b-various-poolings/sum\",\n    \"mean\" : \"cafa5-esm-2-3b-various-poolings/mean\"\n}\n\n# Length of the different embedding vectors :\nembeds_dim = {\n    \"ESM2\" : 2560\n}","metadata":{"execution":{"iopub.status.busy":"2023-05-21T06:29:19.454044Z","iopub.execute_input":"2023-05-21T06:29:19.454677Z","iopub.status.idle":"2023-05-21T06:29:19.460491Z","shell.execute_reply.started":"2023-05-21T06:29:19.454644Z","shell.execute_reply":"2023-05-21T06:29:19.459399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ProteinSequenceDataset(Dataset):\n    def __init__(self, datatype, pooling_method):\n        super(ProteinSequenceDataset).__init__()\n        self.datatype = datatype\n        \n        if pooling_method in ['cls', 'sum', 'mean']:\n            embeds = np.load(\"/kaggle/input/\"+embeds_map[pooling_method]+\"/\"+datatype+\"_embeddings.npy\")\n            ids = np.load(\"/kaggle/input/\"+embeds_map['ids']+\"/\"+datatype+\"_ids.npy\")\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        if datatype==\"train\":\n            labels_vect = np.load(\"/kaggle/working/train_targets_top\"+str(config.num_labels)+\".npy\")\n            df_labels = pd.DataFrame({\"EntryID\": ids, \"labels_vect\": labels_vect.tolist()})\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    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-05-21T06:29:19.461921Z","iopub.execute_input":"2023-05-21T06:29:19.462385Z","iopub.status.idle":"2023-05-21T06:29:19.476171Z","shell.execute_reply.started":"2023-05-21T06:29:19.46235Z","shell.execute_reply":"2023-05-21T06:29:19.475342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MultiLayerPerceptron(torch.nn.Module):\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-05-21T06:29:19.477727Z","iopub.execute_input":"2023-05-21T06:29:19.478323Z","iopub.status.idle":"2023-05-21T06:29:19.490841Z","shell.execute_reply.started":"2023-05-21T06:29:19.478291Z","shell.execute_reply":"2023-05-21T06:29:19.490034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(pooling_method, model_type=\"linear\", train_size=0.9):\n    train_dataset = ProteinSequenceDataset(datatype=\"train\", pooling_method = pooling_method)\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[\"ESM2\"], 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        \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    \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-05-21T06:29:19.492257Z","iopub.execute_input":"2023-05-21T06:29:19.493038Z","iopub.status.idle":"2023-05-21T06:29:19.510832Z","shell.execute_reply.started":"2023-05-21T06:29:19.492994Z","shell.execute_reply":"2023-05-21T06:29:19.509925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cls_model, cls_losses, cls_scores = train_model(pooling_method=\"cls\",model_type=\"linear\")","metadata":{"execution":{"iopub.status.busy":"2023-05-21T06:33:51.655341Z","iopub.execute_input":"2023-05-21T06:33:51.656041Z","iopub.status.idle":"2023-05-21T06:36:47.121319Z","shell.execute_reply.started":"2023-05-21T06:33:51.656008Z","shell.execute_reply":"2023-05-21T06:36:47.120155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sum_model, sum_losses, sum_scores = train_model(pooling_method=\"sum\",model_type=\"linear\")","metadata":{"execution":{"iopub.status.busy":"2023-05-21T06:36:47.123391Z","iopub.execute_input":"2023-05-21T06:36:47.123782Z","iopub.status.idle":"2023-05-21T06:39:53.863634Z","shell.execute_reply.started":"2023-05-21T06:36:47.123747Z","shell.execute_reply":"2023-05-21T06:39:53.862579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean_model, mean_losses, mean_scores = train_model(pooling_method=\"mean\",model_type=\"linear\")","metadata":{"execution":{"iopub.status.busy":"2023-05-21T06:39:53.865036Z","iopub.execute_input":"2023-05-21T06:39:53.865478Z","iopub.status.idle":"2023-05-21T06:42:59.67731Z","shell.execute_reply.started":"2023-05-21T06:39:53.865443Z","shell.execute_reply":"2023-05-21T06:42:59.676342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize = (10, 4))\nplt.plot(cls_losses[\"val\"], label = \"cls\")\nplt.plot(sum_losses[\"val\"], label = \"sum\")\nplt.plot(mean_losses[\"val\"], label = \"mean\")\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(cls_scores[\"val\"], label = \"cls\")\nplt.plot(sum_scores[\"val\"], label = \"sum\")\nplt.plot(mean_scores[\"val\"], label = \"mean\")\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-05-21T06:43:26.352211Z","iopub.execute_input":"2023-05-21T06:43:26.352594Z","iopub.status.idle":"2023-05-21T06:43:26.947943Z","shell.execute_reply.started":"2023-05-21T06:43:26.352562Z","shell.execute_reply":"2023-05-21T06:43:26.947093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(pooling_method):\n    \n    test_dataset = ProteinSequenceDataset(datatype=\"test\", pooling_method = pooling_method)\n    test_dataloader = torch.utils.data.DataLoader(test_dataset, batch_size=1, shuffle=False)\n    \n    model = mean_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-05-21T06:43:35.230901Z","iopub.execute_input":"2023-05-21T06:43:35.231262Z","iopub.status.idle":"2023-05-21T06:43:35.244238Z","shell.execute_reply.started":"2023-05-21T06:43:35.231228Z","shell.execute_reply":"2023-05-21T06:43:35.243266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = predict(\"mean\")","metadata":{"execution":{"iopub.status.busy":"2023-05-21T06:43:49.76867Z","iopub.execute_input":"2023-05-21T06:43:49.769298Z","iopub.status.idle":"2023-05-21T06:45:38.094631Z","shell.execute_reply.started":"2023-05-21T06:43:49.769259Z","shell.execute_reply":"2023-05-21T06:45:38.093457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.head(50)","metadata":{"execution":{"iopub.status.busy":"2023-05-21T06:45:38.096956Z","iopub.execute_input":"2023-05-21T06:45:38.097312Z","iopub.status.idle":"2023-05-21T06:45:38.115734Z","shell.execute_reply.started":"2023-05-21T06:45:38.09728Z","shell.execute_reply":"2023-05-21T06:45:38.114623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(submission_df)","metadata":{"execution":{"iopub.status.busy":"2023-05-21T06:45:38.119097Z","iopub.execute_input":"2023-05-21T06:45:38.119405Z","iopub.status.idle":"2023-05-21T06:45:38.125699Z","shell.execute_reply.started":"2023-05-21T06:45:38.119379Z","shell.execute_reply":"2023-05-21T06:45:38.12459Z"},"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-05-21T06:45:38.12792Z","iopub.execute_input":"2023-05-21T06:45:38.128578Z","iopub.status.idle":"2023-05-21T06:50:01.87603Z","shell.execute_reply.started":"2023-05-21T06:45:38.128542Z","shell.execute_reply":"2023-05-21T06:50:01.875059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}