{"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 -q torchmetrics","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\n!pip install -q torchsummary","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!git clone -q 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":{"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\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-02T10:13:16.108859Z","iopub.execute_input":"2023-08-02T10:13:16.109624Z","iopub.status.idle":"2023-08-02T10:13:26.55033Z","shell.execute_reply.started":"2023-08-02T10:13:16.109575Z","shell.execute_reply":"2023-08-02T10:13:26.549179Z"},"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-02T10:13:26.553628Z","iopub.execute_input":"2023-08-02T10:13:26.554042Z","iopub.status.idle":"2023-08-02T10:13:26.565405Z","shell.execute_reply.started":"2023-08-02T10:13:26.554005Z","shell.execute_reply":"2023-08-02T10:13:26.564104Z"},"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-02T10:13:26.566977Z","iopub.execute_input":"2023-08-02T10:13:26.567484Z","iopub.status.idle":"2023-08-02T10:13:26.574941Z","shell.execute_reply.started":"2023-08-02T10:13:26.567448Z","shell.execute_reply":"2023-08-02T10:13:26.573798Z"},"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-02T10:13:26.576443Z","iopub.execute_input":"2023-08-02T10:13:26.577595Z","iopub.status.idle":"2023-08-02T10:13:26.585678Z","shell.execute_reply.started":"2023-08-02T10:13:26.577539Z","shell.execute_reply":"2023-08-02T10:13:26.584627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# get_dataframe(train_terms).head()","metadata":{"execution":{"iopub.status.busy":"2023-08-02T10:13:26.587246Z","iopub.execute_input":"2023-08-02T10:13:26.588598Z","iopub.status.idle":"2023-08-02T10:13:26.596392Z","shell.execute_reply.started":"2023-08-02T10:13:26.588558Z","shell.execute_reply":"2023-08-02T10:13:26.59497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# get_dataframe(train_taxonomy).head()","metadata":{"execution":{"iopub.status.busy":"2023-08-02T10:13:26.598451Z","iopub.execute_input":"2023-08-02T10:13:26.599042Z","iopub.status.idle":"2023-08-02T10:13:26.609949Z","shell.execute_reply.started":"2023-08-02T10:13:26.599004Z","shell.execute_reply":"2023-08-02T10:13:26.608727Z"},"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-02T10:13:26.615437Z","iopub.execute_input":"2023-08-02T10:13:26.615842Z","iopub.status.idle":"2023-08-02T10:13:26.625983Z","shell.execute_reply.started":"2023-08-02T10:13:26.615808Z","shell.execute_reply":"2023-08-02T10:13:26.624793Z"},"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-02T10:13:26.628086Z","iopub.execute_input":"2023-08-02T10:13:26.628597Z","iopub.status.idle":"2023-08-02T10:13:26.644838Z","shell.execute_reply.started":"2023-08-02T10:13:26.628558Z","shell.execute_reply":"2023-08-02T10:13:26.643718Z"},"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-02T10:13:26.646966Z","iopub.execute_input":"2023-08-02T10:13:26.647355Z","iopub.status.idle":"2023-08-02T10:13:33.792783Z","shell.execute_reply.started":"2023-08-02T10:13:26.647321Z","shell.execute_reply":"2023-08-02T10:13:33.791537Z"},"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-02T10:13:33.795668Z","iopub.execute_input":"2023-08-02T10:13:33.796492Z","iopub.status.idle":"2023-08-02T10:13:34.277624Z","shell.execute_reply.started":"2023-08-02T10:13:33.796453Z","shell.execute_reply":"2023-08-02T10:13:34.276441Z"},"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-02T10:13:34.279691Z","iopub.execute_input":"2023-08-02T10:13:34.280574Z","iopub.status.idle":"2023-08-02T10:13:40.262445Z","shell.execute_reply.started":"2023-08-02T10:13:34.280534Z","shell.execute_reply":"2023-08-02T10:13:40.261432Z"},"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-02T10:13:40.264355Z","iopub.execute_input":"2023-08-02T10:13:40.265154Z","iopub.status.idle":"2023-08-02T10:13:46.118931Z","shell.execute_reply.started":"2023-08-02T10:13:40.265116Z","shell.execute_reply":"2023-08-02T10:13:46.117713Z"},"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-02T10:13:46.120613Z","iopub.execute_input":"2023-08-02T10:13:46.121328Z","iopub.status.idle":"2023-08-02T10:13:57.992899Z","shell.execute_reply.started":"2023-08-02T10:13:46.121289Z","shell.execute_reply":"2023-08-02T10:13:57.991709Z"},"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-02T10:13:57.994658Z","iopub.execute_input":"2023-08-02T10:13:57.995828Z","iopub.status.idle":"2023-08-02T10:13:58.255874Z","shell.execute_reply.started":"2023-08-02T10:13:57.99579Z","shell.execute_reply":"2023-08-02T10:13:58.253958Z"},"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-02T10:13:58.257323Z","iopub.execute_input":"2023-08-02T10:13:58.257766Z","iopub.status.idle":"2023-08-02T10:13:58.434259Z","shell.execute_reply.started":"2023-08-02T10:13:58.257723Z","shell.execute_reply":"2023-08-02T10:13:58.432888Z"},"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-02T10:13:58.440658Z","iopub.execute_input":"2023-08-02T10:13:58.441519Z","iopub.status.idle":"2023-08-02T10:14:49.617891Z","shell.execute_reply.started":"2023-08-02T10:13:58.441463Z","shell.execute_reply":"2023-08-02T10:14:49.61672Z"},"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-02T10:14:49.619615Z","iopub.execute_input":"2023-08-02T10:14:49.620302Z","iopub.status.idle":"2023-08-02T10:14:51.051851Z","shell.execute_reply.started":"2023-08-02T10:14:49.620262Z","shell.execute_reply":"2023-08-02T10:14:51.050769Z"},"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-02T10:14:51.053518Z","iopub.execute_input":"2023-08-02T10:14:51.053903Z","iopub.status.idle":"2023-08-02T10:14:51.128905Z","shell.execute_reply.started":"2023-08-02T10:14:51.053868Z","shell.execute_reply":"2023-08-02T10:14:51.127657Z"},"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-02T10:14:51.130885Z","iopub.execute_input":"2023-08-02T10:14:51.131966Z","iopub.status.idle":"2023-08-02T10:14:58.609094Z","shell.execute_reply.started":"2023-08-02T10:14:51.131926Z","shell.execute_reply":"2023-08-02T10:14:58.608033Z"},"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-02T10:14:58.610689Z","iopub.execute_input":"2023-08-02T10:14:58.61109Z","iopub.status.idle":"2023-08-02T10:14:58.6222Z","shell.execute_reply.started":"2023-08-02T10:14:58.611051Z","shell.execute_reply":"2023-08-02T10:14:58.620579Z"},"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, 1200)\n        self.ln1 = nn.LayerNorm(1200, elementwise_affine=True)\n        \n        self.bn2 = nn.BatchNorm1d(1200)\n        self.fc2 = nn.Linear(1200, 1200)\n        self.ln2 = nn.LayerNorm(1200, elementwise_affine=True)\n        \n        self.bn3 = nn.BatchNorm1d(1200)\n        self.fc3 = nn.Linear(1200, 1200)\n        self.ln3 = nn.LayerNorm(1200, elementwise_affine=True)\n        \n        self.bn4 = nn.BatchNorm1d(2400)\n        self.fc4 = nn.Linear(2400, 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.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-02T10:14:58.623939Z","iopub.execute_input":"2023-08-02T10:14:58.624655Z","iopub.status.idle":"2023-08-02T10:14:58.640862Z","shell.execute_reply.started":"2023-08-02T10:14:58.624619Z","shell.execute_reply":"2023-08-02T10:14:58.64007Z"},"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-02T10:45:53.34892Z","iopub.execute_input":"2023-08-02T10:45:53.349499Z","iopub.status.idle":"2023-08-02T10:45:53.435746Z","shell.execute_reply.started":"2023-08-02T10:45:53.349457Z","shell.execute_reply":"2023-08-02T10:45:53.43464Z"},"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-02T10:45:53.566851Z","iopub.execute_input":"2023-08-02T10:45:53.56763Z","iopub.status.idle":"2023-08-02T10:45:53.661263Z","shell.execute_reply.started":"2023-08-02T10:45:53.567582Z","shell.execute_reply":"2023-08-02T10:45:53.659907Z"},"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-02T10:45:53.900342Z","iopub.execute_input":"2023-08-02T10:45:53.900863Z","iopub.status.idle":"2023-08-02T10:45:53.906641Z","shell.execute_reply.started":"2023-08-02T10:45:53.90082Z","shell.execute_reply":"2023-08-02T10:45:53.905376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 32 #5120\nEPOCHS = 39\nLEARNING_RATE = 0.0001\nMOMENTUM = 0.9\nOPT_FUNC = SophiaG","metadata":{"execution":{"iopub.status.busy":"2023-08-02T10:45:54.058325Z","iopub.execute_input":"2023-08-02T10:45:54.058864Z","iopub.status.idle":"2023-08-02T10:45:54.065393Z","shell.execute_reply.started":"2023-08-02T10:45:54.058823Z","shell.execute_reply":"2023-08-02T10:45:54.064233Z"},"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-02T10:45:54.28768Z","iopub.execute_input":"2023-08-02T10:45:54.288069Z","iopub.status.idle":"2023-08-02T10:45:54.294443Z","shell.execute_reply.started":"2023-08-02T10:45:54.28803Z","shell.execute_reply":"2023-08-02T10:45:54.29341Z"},"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, 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\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            \n        result = evaluate(model, val_loader)\n        result['Train_loss'] = torch.stack(train_losses).mean().item()\n        model.epoch_end(epoch, result)\n        history.append(result)\n    \n    return history, optimizer","metadata":{"execution":{"iopub.status.busy":"2023-08-02T10:45:54.557497Z","iopub.execute_input":"2023-08-02T10:45:54.557941Z","iopub.status.idle":"2023-08-02T10:45:54.571139Z","shell.execute_reply.started":"2023-08-02T10:45:54.557908Z","shell.execute_reply":"2023-08-02T10:45:54.569792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history, optimizer = fit(EPOCHS, LEARNING_RATE, model, train_dl, val_dl,OPT_FUNC)","metadata":{"execution":{"iopub.status.busy":"2023-08-02T11:20:42.407826Z","iopub.execute_input":"2023-08-02T11:20:42.408268Z","iopub.status.idle":"2023-08-02T11:24:53.791062Z","shell.execute_reply.started":"2023-08-02T11:20:42.408221Z","shell.execute_reply":"2023-08-02T11:24:53.789973Z"},"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-02T11:24:53.799763Z","iopub.execute_input":"2023-08-02T11:24:53.800119Z","iopub.status.idle":"2023-08-02T11:24:54.121444Z","shell.execute_reply.started":"2023-08-02T11:24:53.800071Z","shell.execute_reply":"2023-08-02T11:24:54.120465Z"},"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(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 = model(X_batch_test)\n            prediction_target=predictions.detach().cpu().numpy()\n\n    return prediction_target\n\nprediction_target = eval_test_data(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 %1000000 == 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":[]}]}