{"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":"%%capture \n!pip install torchmetrics","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\n!pip install torchsummary","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!git clone https://github.com/kyegomez/Sophia.git\n! python Sophia/setup.py install","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm Sophia/Sophia/__init__.py","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from Sophia.Sophia.Sophia import SophiaG","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:13:42.152289Z","iopub.execute_input":"2023-08-01T15:13:42.153496Z","iopub.status.idle":"2023-08-01T15:13:44.641192Z","shell.execute_reply.started":"2023-08-01T15:13:42.153424Z","shell.execute_reply":"2023-08-01T15:13:44.64015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## https://www.kaggle.com/code/alexandervc/baseline-multilabel-to-multitarget-binary#Load-train-features---precalculated-embeddings-for-the-proteins\nimport os\nimport gc\nfrom sklearn.model_selection import train_test_split\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm\ntqdm.pandas()\nimport torch\nimport torch.nn as nn\nimport warnings\nwarnings.filterwarnings('ignore')\nfrom sklearn.model_selection import train_test_split\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchmetrics import AUROC,F1Score\nfrom torchmetrics.classification import BinaryF1Score\nfrom torchsummary import summary as torchsummary\nfrom torch.optim.swa_utils import AveragedModel, SWALR\nfrom torch.optim.lr_scheduler import CosineAnnealingLR","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-01T15:13:46.556562Z","iopub.execute_input":"2023-08-01T15:13:46.557148Z","iopub.status.idle":"2023-08-01T15:13:53.731696Z","shell.execute_reply.started":"2023-08-01T15:13:46.557115Z","shell.execute_reply":"2023-08-01T15:13:53.730667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nSEED = 42\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\nnp.random.seed(SEED)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:13:53.733591Z","iopub.execute_input":"2023-08-01T15:13:53.733962Z","iopub.status.idle":"2023-08-01T15:13:53.743974Z","shell.execute_reply.started":"2023-08-01T15:13:53.733929Z","shell.execute_reply":"2023-08-01T15:13:53.74295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_dataframe(path):\n    return pd.read_csv(path,sep='\\t')","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:13:53.745685Z","iopub.execute_input":"2023-08-01T15:13:53.74667Z","iopub.status.idle":"2023-08-01T15:13:53.75154Z","shell.execute_reply.started":"2023-08-01T15:13:53.746635Z","shell.execute_reply":"2023-08-01T15:13:53.750438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_terms = '/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv'\ntrain_taxonomy ='/kaggle/input/cafa-5-protein-function-prediction/Train/train_taxonomy.tsv'","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:13:53.754375Z","iopub.execute_input":"2023-08-01T15:13:53.75512Z","iopub.status.idle":"2023-08-01T15:13:53.759889Z","shell.execute_reply.started":"2023-08-01T15:13:53.755087Z","shell.execute_reply":"2023-08-01T15:13:53.758842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# get_dataframe(train_terms).head()","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:13:53.762212Z","iopub.execute_input":"2023-08-01T15:13:53.762944Z","iopub.status.idle":"2023-08-01T15:13:53.768317Z","shell.execute_reply.started":"2023-08-01T15:13:53.762912Z","shell.execute_reply":"2023-08-01T15:13:53.767232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# get_dataframe(train_taxonomy).head()","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:13:53.769712Z","iopub.execute_input":"2023-08-01T15:13:53.770943Z","iopub.status.idle":"2023-08-01T15:13:53.776517Z","shell.execute_reply.started":"2023-08-01T15:13:53.770911Z","shell.execute_reply":"2023-08-01T15:13:53.775528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def summary(text, df):\n    print(f'{text} shape: {df.shape}')\n    summ = pd.DataFrame(df.dtypes, columns=['dtypes'])\n    summ['null'] = df.isnull().sum()\n    summ['unique'] = df.nunique()\n    summ['min'] = df.min()\n    summ['median'] = df.median()\n    summ['max'] = df.max()\n    summ['mean'] = df.mean()\n    summ['std'] = df.std()\n    summ['duplicate'] = df.duplicated().sum()\n    return summ","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:13:53.778004Z","iopub.execute_input":"2023-08-01T15:13:53.778597Z","iopub.status.idle":"2023-08-01T15:13:53.788341Z","shell.execute_reply.started":"2023-08-01T15:13:53.778529Z","shell.execute_reply":"2023-08-01T15:13:53.787299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def reduce_mem_usage(df):\n    \"\"\" iterate through all the columns of a dataframe and modify the data type\n        to reduce memory usage.        \n    \"\"\"\n    start_mem = df.memory_usage().sum() / 1024**2\n    print('Memory usage of dataframe is {:.2f} MB'.format(start_mem))\n    \n    for col in df.columns:\n        col_type = df[col].dtype\n        \n        if col_type != object:\n            c_min = df[col].min()\n            c_max = df[col].max()\n            if str(col_type)[:3] == 'int':\n                if c_min > np.iinfo(np.int8).min and c_max < np.iinfo(np.int8).max:\n                    df[col] = df[col].astype(np.int8)\n                elif c_min > np.iinfo(np.int16).min and c_max < np.iinfo(np.int16).max:\n                    df[col] = df[col].astype(np.int16)\n                elif c_min > np.iinfo(np.int32).min and c_max < np.iinfo(np.int32).max:\n                    df[col] = df[col].astype(np.int32)\n                elif c_min > np.iinfo(np.int64).min and c_max < np.iinfo(np.int64).max:\n                    df[col] = df[col].astype(np.int64)  \n            else:\n                if c_min > np.finfo(np.float16).min and c_max < np.finfo(np.float16).max:\n                    df[col] = df[col].astype(np.float16)\n                elif c_min > np.finfo(np.float32).min and c_max < np.finfo(np.float32).max:\n                    df[col] = df[col].astype(np.float32)\n                else:\n                    df[col] = df[col].astype(np.float64)\n        else:\n            df[col] = df[col].astype('category')\n\n    end_mem = df.memory_usage().sum() / 1024**2\n    print('Memory usage after optimization is: {:.2f} MB'.format(end_mem))\n    print('Decreased by {:.1f}%'.format(100 * (start_mem - end_mem) / start_mem))\n    \n    return df","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:13:53.790122Z","iopub.execute_input":"2023-08-01T15:13:53.790402Z","iopub.status.idle":"2023-08-01T15:13:53.805967Z","shell.execute_reply.started":"2023-08-01T15:13:53.790381Z","shell.execute_reply":"2023-08-01T15:13:53.805334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary('train_terms',reduce_mem_usage(get_dataframe(train_terms)))","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:13:53.807362Z","iopub.execute_input":"2023-08-01T15:13:53.808045Z","iopub.status.idle":"2023-08-01T15:14:01.477446Z","shell.execute_reply.started":"2023-08-01T15:13:53.808012Z","shell.execute_reply":"2023-08-01T15:14:01.47631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary('train_terms',reduce_mem_usage(get_dataframe(train_taxonomy)))","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:14:01.48241Z","iopub.execute_input":"2023-08-01T15:14:01.482738Z","iopub.status.idle":"2023-08-01T15:14:01.920139Z","shell.execute_reply.started":"2023-08-01T15:14:01.482711Z","shell.execute_reply":"2023-08-01T15:14:01.919057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.countplot(data=reduce_mem_usage(get_dataframe(train_terms)),x='aspect',color='r')","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:14:01.921737Z","iopub.execute_input":"2023-08-01T15:14:01.923736Z","iopub.status.idle":"2023-08-01T15:14:07.114454Z","shell.execute_reply.started":"2023-08-01T15:14:01.923695Z","shell.execute_reply":"2023-08-01T15:14:07.113381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_terms=reduce_mem_usage(get_dataframe(train_terms))","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:14:07.116452Z","iopub.execute_input":"2023-08-01T15:14:07.11725Z","iopub.status.idle":"2023-08-01T15:14:11.882259Z","shell.execute_reply.started":"2023-08-01T15:14:07.11721Z","shell.execute_reply":"2023-08-01T15:14:11.881199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_dataset():\n#     train_protein_ids = np.load('/kaggle/input/4637427/train_ids_esm2_t36_3B_UR50D.npy')\n#     train_embeddings = np.load('/kaggle/input/4637427/train_embeds_esm2_t36_3B_UR50D.npy')\n#     train_protein_ids = np.load('/kaggle/input/t5embeds/train_ids.npy')\n#     train_embeddings = np.load('/kaggle/input/t5embeds/train_embeds.npy')\n\n    train_protein_ids = np.load('/kaggle/input/23468234/train_ids_esm2_t33_650M_UR50D.npy')\n    train_embeddings = np.load('/kaggle/input/23468234/train_embeds_esm2_t33_650M_UR50D.npy')\n    \n    column_num = train_embeddings.shape[1]\n    train = pd.DataFrame(train_embeddings, columns = [\"Column_\" + str(i) for i in range(1, column_num+1)])\n    return train,train_protein_ids\n\ntrain,train_protein_ids = get_train_dataset()\nprint(train.shape,train_protein_ids.shape)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:14:11.884603Z","iopub.execute_input":"2023-08-01T15:14:11.885379Z","iopub.status.idle":"2023-08-01T15:14:20.806266Z","shell.execute_reply.started":"2023-08-01T15:14:11.885341Z","shell.execute_reply":"2023-08-01T15:14:20.805206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_of_labels = 2000\ndef get_label_train_terms(df):\n    labels=df['term'].value_counts().index[:num_of_labels].tolist()\n    train_terms_updated=df.loc[df['term'].isin(labels)]\n    return labels,train_terms_updated\n\nlabels_count,train_terms_updated=get_label_train_terms(train_terms)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:14:20.807817Z","iopub.execute_input":"2023-08-01T15:14:20.808169Z","iopub.status.idle":"2023-08-01T15:14:21.02435Z","shell.execute_reply.started":"2023-08-01T15:14:20.808136Z","shell.execute_reply":"2023-08-01T15:14:21.023365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_pit_aspects():\n    pie_df = train_terms_updated['aspect'].value_counts()\n    palette_color = sns.color_palette('pastel')\n    plt.pie(pie_df.values, labels=np.array(pie_df.index), colors=palette_color, autopct='%.0f%%')\n    plt.show()\n    \nshow_pit_aspects()","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:14:21.026567Z","iopub.execute_input":"2023-08-01T15:14:21.027298Z","iopub.status.idle":"2023-08-01T15:14:21.176244Z","shell.execute_reply.started":"2023-08-01T15:14:21.027261Z","shell.execute_reply":"2023-08-01T15:14:21.175226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_labels(train_protein_ids):\n    train_size = train_protein_ids.shape[0] # len(X)\n    train_labels = np.zeros((train_size ,num_of_labels))\n    series_train_protein_ids = pd.Series(train_protein_ids)\n\n    for i in range(num_of_labels):\n        n_train_terms = train_terms_updated[train_terms_updated['term'] ==  labels_count[i]]\n        label_related_proteins = n_train_terms['EntryID'].unique()\n        train_labels[:,i] =  series_train_protein_ids.isin(label_related_proteins).astype(float)\n    return train_labels\n\n\ntrain_labels=get_labels(train_protein_ids)\n\nlabels = pd.DataFrame(data = train_labels, columns = labels_count)\nprint(labels.shape)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:14:21.178405Z","iopub.execute_input":"2023-08-01T15:14:21.179133Z","iopub.status.idle":"2023-08-01T15:15:09.846899Z","shell.execute_reply.started":"2023-08-01T15:14:21.179098Z","shell.execute_reply":"2023-08-01T15:15:09.845766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_test_dataset(features,labels):\n    return  train_test_split(features,labels,shuffle=True,random_state=42)\n\nX_train,X_val,y_train,y_val = train_test_dataset(train,labels)\nprint(X_train.shape,X_val.shape,y_train.shape,y_val.shape)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:09.848718Z","iopub.execute_input":"2023-08-01T15:15:09.849087Z","iopub.status.idle":"2023-08-01T15:15:11.015058Z","shell.execute_reply.started":"2023-08-01T15:15:09.849054Z","shell.execute_reply":"2023-08-01T15:15:11.013217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:11.016677Z","iopub.execute_input":"2023-08-01T15:15:11.017048Z","iopub.status.idle":"2023-08-01T15:15:11.050054Z","shell.execute_reply.started":"2023-08-01T15:15:11.017016Z","shell.execute_reply":"2023-08-01T15:15:11.049039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FluxData(Dataset):\n    def __init__(self, X_data, y_data):\n        self.X_data = X_data\n        self.y_data = y_data\n        \n    def __getitem__(self, index):\n            return self.X_data[index], self.y_data[index]\n        \n    def __len__ (self):\n        return len(self.X_data)\n\nX_data = torch.from_numpy(X_train.values).float().to(device)\ny_data = torch.from_numpy(y_train.values).float().to(device)\nX_val = torch.from_numpy(X_val.values).float().to(device)\ny_val = torch.from_numpy(y_val.values).float().to(device)\ntrain_data = FluxData(X_data,y_data)\ntest_data = FluxData(X_val,y_val)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:11.051794Z","iopub.execute_input":"2023-08-01T15:15:11.053152Z","iopub.status.idle":"2023-08-01T15:15:15.423369Z","shell.execute_reply.started":"2023-08-01T15:15:11.053116Z","shell.execute_reply":"2023-08-01T15:15:15.422299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CAFA5NNetBase(torch.nn.Module):\n    \n    def training_step(self,batch):\n        features,labels = batch\n        out = self(features)\n        loss = F.binary_cross_entropy(out,labels)\n        return loss\n    \n    def validation_step(self, batch):\n        features, labels = batch \n        out = self(features)                    # Generate predictions\n        loss = F.binary_cross_entropy(out, labels)   # Calculate loss\n        acc = auroc(out, labels)           # Calculate accuracy\n        return {'Validation_loss': loss.detach(), 'Validation_acc': acc}\n        \n    def validation_epoch_end(self, outputs):\n        batch_losses = [x['Validation_loss'] for x in outputs]\n        epoch_loss = torch.stack(batch_losses).mean()   # Combine losses\n        batch_accs = [x['Validation_acc'] for x in outputs]\n        epoch_acc = torch.stack(batch_accs).mean()      # Combine accuracies\n        return {'Validation_loss': epoch_loss.item(), 'Validation_acc': epoch_acc.item()}\n    \n    def epoch_end(self, epoch, result):\n        if epoch%5==0:\n            print(\"Epoch [{}], Train_loss: {:.4f}, Validation_loss: {:.4f}, Validation_acc: {:.4f}\".format(\n            epoch, result['Train_loss'], result['Validation_loss'], result['Validation_acc']))","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:15.424889Z","iopub.execute_input":"2023-08-01T15:15:15.425252Z","iopub.status.idle":"2023-08-01T15:15:15.437065Z","shell.execute_reply.started":"2023-08-01T15:15:15.425218Z","shell.execute_reply":"2023-08-01T15:15:15.435758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CAFA5NNet(CAFA5NNetBase):\n    def __init__(self,input_features,output_features):\n        super(CAFA5NNet,self).__init__()\n        \n        self.activation = nn.PReLU()\n        \n        self.bn1 = nn.BatchNorm1d(input_features)\n        self.fc1 = nn.Linear(input_features, 800)\n        self.ln1 = nn.LayerNorm(800, elementwise_affine=True)\n        \n        self.bn2 = nn.BatchNorm1d(800)\n        self.fc2 = nn.Linear(800, 600)\n        self.ln2 = nn.LayerNorm(600, elementwise_affine=True)\n        \n        self.bn3 = nn.BatchNorm1d(600)\n        self.fc3 = nn.Linear(600, 400)\n        self.ln3 = nn.LayerNorm(400, elementwise_affine=True)\n        \n        self.bn4 = nn.BatchNorm1d(1200)\n        self.fc4 = nn.Linear(1200, output_features)\n        self.ln4 = nn.LayerNorm(output_features, elementwise_affine=True)\n        \n        self.sigm = nn.Sigmoid()\n    def forward(self,inputs):\n#         print(inputs.shape)\n\n        fc1_out = self.bn1(inputs)\n        fc1_out = self.ln1(self.fc1(inputs))\n        fc1_out = self.activation(fc1_out)\n        \n        x = self.bn2(fc1_out)\n        \n        x = self.ln2(self.fc2(x))\n        x = self.activation(x)\n        \n        x = self.bn3(x)\n        \n        x = self.ln3(self.fc3(x))\n        x = self.activation(x)\n        \n        x = torch.cat([x, fc1_out], axis = -1)\n        \n        x = self.bn4(x)\n        \n        x = self.ln4(self.fc4(x))\n        out = self.sigm(x)\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:15.440244Z","iopub.execute_input":"2023-08-01T15:15:15.440597Z","iopub.status.idle":"2023-08-01T15:15:15.453434Z","shell.execute_reply.started":"2023-08-01T15:15:15.440554Z","shell.execute_reply":"2023-08-01T15:15:15.452376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CAFA5NNet(X_train.shape[1],y_train.shape[1])\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:15.454652Z","iopub.execute_input":"2023-08-01T15:15:15.454935Z","iopub.status.idle":"2023-08-01T15:15:15.516686Z","shell.execute_reply.started":"2023-08-01T15:15:15.454906Z","shell.execute_reply":"2023-08-01T15:15:15.515605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:15.51809Z","iopub.execute_input":"2023-08-01T15:15:15.518435Z","iopub.status.idle":"2023-08-01T15:15:15.523609Z","shell.execute_reply.started":"2023-08-01T15:15:15.518404Z","shell.execute_reply":"2023-08-01T15:15:15.522276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torchsummary(model, X_data.size(), batch_size=-1, device='cuda')","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:15.525572Z","iopub.execute_input":"2023-08-01T15:15:15.526199Z","iopub.status.idle":"2023-08-01T15:15:15.53167Z","shell.execute_reply.started":"2023-08-01T15:15:15.526163Z","shell.execute_reply":"2023-08-01T15:15:15.530655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 128 #5120\nEPOCHS = 39\nLEARNING_RATE = 0.001\nMOMENTUM = 0.9\nOPT_FUNC = SophiaG\nswa_model = AveragedModel(model)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:15.53307Z","iopub.execute_input":"2023-08-01T15:15:15.534277Z","iopub.status.idle":"2023-08-01T15:15:15.540939Z","shell.execute_reply.started":"2023-08-01T15:15:15.534241Z","shell.execute_reply":"2023-08-01T15:15:15.539862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_dataloaders(dataset_type,batch,shuffle):\n    if shuffle:\n         return DataLoader(dataset=dataset_type, batch_size=batch, shuffle=True)\n    else:\n        return DataLoader(dataset=dataset_type, batch_size=batch,shuffle=False)\n    \ntrain_dl = get_dataloaders(train_data,BATCH_SIZE,True)\nval_dl = get_dataloaders(test_data,BATCH_SIZE,False)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:15.542838Z","iopub.execute_input":"2023-08-01T15:15:15.543241Z","iopub.status.idle":"2023-08-01T15:15:15.551688Z","shell.execute_reply.started":"2023-08-01T15:15:15.54321Z","shell.execute_reply":"2023-08-01T15:15:15.550717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def auroc(outputs, labels):\n    auroc = AUROC(task=\"binary\")\n    return auroc(outputs, labels)\n\n  \n@torch.no_grad()\ndef evaluate(model, val_loader):\n    model.eval()\n    outputs = [model.validation_step(batch) for batch in val_loader]\n    return model.validation_epoch_end(outputs)\n\n  \ndef fit(epochs, lr, model,swa_model, train_loader, val_loader, opt_func = OPT_FUNC):\n    \n    history = []\n    optimizer = opt_func(model.parameters(),lr, betas=(0.965, 0.99), rho = 0.01, weight_decay=1e-1)\n#     scheduler = torch.optim.lr_scheduler.OneCycleLR(\n#                 optimizer, \n#                 max_lr=lr, \n#                 steps_per_epoch=len(train_loader), \n#                 epochs=epochs\n#                 )\n#     optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) #SGD(model.parameters(), lr=1e-3)\n    lr_sched = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2, eta_min=0.001, last_epoch=-1)\n    swa_model = AveragedModel(model)\n    scheduler = CosineAnnealingLR(optimizer, T_max=100)\n    swa_start = 5\n    swa_scheduler = SWALR(optimizer, swa_lr=0.05)\n    for epoch in tqdm(range(epochs)):\n        \n        model.train()\n        train_losses = []\n        for batch in train_loader:\n            loss = model.training_step(batch)\n            train_losses.append(loss)\n            loss.backward()\n            optimizer.step()\n            lr_sched.step()\n            optimizer.zero_grad()\n            if epoch > swa_start:\n                swa_model.update_parameters(model)\n                swa_scheduler.step()\n            else:\n                scheduler.step()\n            \n        result = evaluate(swa_model, val_loader)\n        result['Train_loss'] = torch.stack(train_losses).mean().item()\n        model.epoch_end(epoch, result)\n        history.append(result)\n    torch.optim.swa_utils.update_bn(val_loader, swa_model)\n        \n    return history","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:15.553269Z","iopub.execute_input":"2023-08-01T15:15:15.553643Z","iopub.status.idle":"2023-08-01T15:15:15.565745Z","shell.execute_reply.started":"2023-08-01T15:15:15.553611Z","shell.execute_reply":"2023-08-01T15:15:15.564841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = fit(EPOCHS, LEARNING_RATE, swa_model, train_dl, val_dl,OPT_FUNC)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:15:15.571295Z","iopub.execute_input":"2023-08-01T15:15:15.571582Z","iopub.status.idle":"2023-08-01T15:19:08.166682Z","shell.execute_reply.started":"2023-08-01T15:15:15.571558Z","shell.execute_reply":"2023-08-01T15:19:08.165564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_accuracies(history):\n    \"\"\" Plot the history of accuracies\"\"\"\n    accuracies = [x['Validation_acc'] for x in history]\n    plt.plot(accuracies, '-x')\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy')\n    plt.title('Accuracy vs. No. of epochs');\n    \n\nplot_accuracies(history)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:19:08.447771Z","iopub.execute_input":"2023-08-01T15:19:08.448126Z","iopub.status.idle":"2023-08-01T15:19:08.712583Z","shell.execute_reply.started":"2023-08-01T15:19:08.448095Z","shell.execute_reply":"2023-08-01T15:19:08.711436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_losses(history):\n    \"\"\" Plot the losses in each epoch\"\"\"\n    train_losses = [x.get('Train_loss') for x in history]\n    val_losses = [x['Validation_loss'] for x in history]\n    plt.plot(train_losses, '-bx')\n    plt.plot(val_losses, '-rx')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend(['Training', 'Validation'])\n    plt.title('Loss vs. No. of epochs');\n\nplot_losses(history)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T15:19:21.18553Z","iopub.execute_input":"2023-08-01T15:19:21.186344Z","iopub.status.idle":"2023-08-01T15:19:21.474325Z","shell.execute_reply.started":"2023-08-01T15:19:21.186304Z","shell.execute_reply":"2023-08-01T15:19:21.473368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del train, X_val, y_val, X_train, y_train, X_data,y_data, train_data, test_data, train_protein_ids, train_dl, val_dl\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_test_dataset():\n#     test_embeddings = np.load('/kaggle/input/4637427/test_embeds_esm2_t36_3B_UR50D.npy')\n    test_embeddings = np.load('/kaggle/input/23468234/test_embeds_esm2_t33_650M_UR50D.npy')\n#     test_embeddings = np.load('/kaggle/input/t5embeds/test_embeds.npy')\n    column_num = test_embeddings.shape[1]\n    test = pd.DataFrame(test_embeddings, columns = [\"Column_\" + str(i) for i in range(1, column_num+1)])\n    return test \n\ntest = get_test_dataset()\nprint(test.shape)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CAFA5TestData(Dataset):\n    \n    def __init__(self, X_test_data):\n        self.X_test_data = X_test_data\n        \n    def __getitem__(self, index):\n        return self.X_test_data[index]\n        \n    def __len__ (self):\n        return len(self.X_test_data)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = CAFA5TestData(torch.from_numpy(test.values).float().to(device))\ntest_data_loader = DataLoader(dataset=test_data, batch_size=test.shape[0])\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef eval_test_data(swa_model,testing_data_dl):\n    model.eval()\n    with torch.no_grad():\n        for X_batch_test in testing_data_dl:\n            X_batch_test = X_batch_test.to(device)\n            predictions = swa_model(X_batch_test)\n            prediction_target=predictions.detach().cpu().numpy()\n\n    return prediction_target\n\nprediction_target = eval_test_data(swa_model,test_data_loader)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"test predicted\")\ndel test_data_loader, model, test_data, test\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_predictions(prediction_target):\n#     test_protein_ids = np.load('/kaggle/input/4637427/test_ids_esm2_t36_3B_UR50D.npy')\n    test_protein_ids = np.load('/kaggle/input/23468234/test_ids_esm2_t33_650M_UR50D.npy')\n#     test_protein_ids = np.load('/kaggle/input/t5embeds/test_ids.npy')\n    protein_list = []\n    for k in list(test_protein_ids):\n        protein_list += [k] * prediction_target.shape[1] \n    return protein_list\n\nprotein_list=make_predictions(prediction_target)      ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# del labels\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\ndef submit(labels_count, protein_list, prediction_target):\n    labels_count = labels_count * prediction_target.shape[0]  # List of labels\n\n    with open(\"submission.tsv\", \"w\") as file:\n#         file.write(f\"Protein Id\\tGO Term Id\\tPrediction\\n\")\n        idx = 0\n        for row in prediction_target:\n            for element in row:\n                file.write(f\"{protein_list[idx]}\\t{labels_count[idx]}\\t{element}\\n\")\n                idx += 1\n                if idx %100000 == 0:\n                    print(f'{idx} passed')\n                \nsubmit(labels_count, protein_list, prediction_target)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def submit():\n#     df_submission = pd.DataFrame(columns = ['Protein Id', 'GO Term Id','Prediction'])\n#     df_submission['Protein Id'] = protein_list\n#     df_submission['GO Term Id'] = labels_count * prediction_target.shape[0]\n#     df_submission['Prediction'] = prediction_target.ravel()\n#     df_submission.to_csv(\"submission.tsv\",header=False, index=False,sep='\\t')\n    \n# submit()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# temp=pd.read_csv('/kaggle/working/submission.tsv',sep='\\t')\n# temp.count()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}