{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":41875,"databundleVersionId":5521661,"sourceType":"competition"},{"sourceId":5499219,"sourceType":"datasetVersion","datasetId":3167603}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Load the libraries.\nimport numpy as np \nimport pandas as pd \n%matplotlib inline\nimport matplotlib.pyplot as plt\nfrom sklearn.preprocessing import MultiLabelBinarizer\nimport locale\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import TensorDataset,random_split,DataLoader\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n\nprint(f\"Numpy.version:{np.__version__} has been loaded sucessfully. \")\nprint(f\"Pandas.version:{pd.__version__} has been loaded sucessfully. \")\nprint(f\"Pytorch.version:{torch.__version__} has been loaded sucessfully. \")\nprint(locale.getpreferredencoding())","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:10:44.774781Z","iopub.execute_input":"2025-04-06T09:10:44.775132Z","iopub.status.idle":"2025-04-06T09:10:48.783612Z","shell.execute_reply.started":"2025-04-06T09:10:44.775105Z","shell.execute_reply":"2025-04-06T09:10:48.782487Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load the files.\n\n# The t5embeds file path:\nt5_test_embeds_file_path = \"/kaggle/input/t5embeds/test_embeds.npy\"\nt5_test_ids_file_path = \"/kaggle/input/t5embeds/test_ids.npy\"\nt5_train_embeds_file_path = \"/kaggle/input/t5embeds/train_embeds.npy\"\nt5_train_ids_file_path = \"/kaggle/input/t5embeds/train_ids.npy\"\n\n# The train terms file path:\ntrain_terms_file_path = \"/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv\"\n\n# The IA weight file path:\nIA_file_path = \"/kaggle/input/cafa-5-protein-function-prediction/IA.txt\"\n\nprint(\"Files load successfully.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:10:48.784798Z","iopub.execute_input":"2025-04-06T09:10:48.785253Z","iopub.status.idle":"2025-04-06T09:10:48.789781Z","shell.execute_reply.started":"2025-04-06T09:10:48.785229Z","shell.execute_reply":"2025-04-06T09:10:48.788914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define a function to read the *.tsv file.\ndef read_tsv_to_df(file_path):\n    df = pd.read_csv(file_path,sep=\"\\t\")\n    return df\n\nprint(\"Functions loading successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:10:48.791374Z","iopub.execute_input":"2025-04-06T09:10:48.791644Z","iopub.status.idle":"2025-04-06T09:10:48.805876Z","shell.execute_reply.started":"2025-04-06T09:10:48.791623Z","shell.execute_reply":"2025-04-06T09:10:48.805221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define a function to read the *.txt file.\ndef read_txt_to_df(file_path):\n    with open(file_path,\"r\",encoding=\"UTF-8\") as f:\n        df_ia_weight = pd.DataFrame(columns=[\"terms\",\"ia_weight\"])\n        i = 0\n        for line in f:\n            terms,ia_weight = line.strip().split()\n            df_ia_weight.loc[i,\"terms\"] = terms\n            df_ia_weight.loc[i,\"ia_weight\"] = float(ia_weight)\n            i += 1\n        return df_ia_weight\n\nprint(\"Functions loading successfully.\")","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-04-06T09:10:48.806796Z","iopub.execute_input":"2025-04-06T09:10:48.807047Z","iopub.status.idle":"2025-04-06T09:10:48.820218Z","shell.execute_reply.started":"2025-04-06T09:10:48.807014Z","shell.execute_reply":"2025-04-06T09:10:48.819361Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Put the data into DataFrame.\n\n# Load the files.\nt5_train_embeds = np.load(t5_train_embeds_file_path)\nt5_test_embeds = np.load(t5_test_embeds_file_path)\nt5_train_ids = np.load(t5_train_ids_file_path)\nt5_test_ids = np.load(t5_test_ids_file_path)\ndf_train_terms = read_tsv_to_df(train_terms_file_path)\ndf_ia_weight = read_txt_to_df(IA_file_path)\n\n# Check the shape of the files.\nprint(t5_train_embeds.shape)\nprint(t5_train_ids.shape)\nprint(t5_test_embeds.shape)\nprint(t5_test_ids.shape)\nprint(df_train_terms.shape)\nprint(df_ia_weight.shape)\n\n# t5_train_embeds and t5_test_embeds share the same numeric value of the dimision for the embeddings.\n# embeds_dim = t5_train_embeds.shape[1]\n\n# Put t5_train_embeds, t5_test_embeds, t5_train_ids, t5_test_ids into DataFrame.\n# df_train_embeds = pd.DataFrame(t5_train_embeds,columns=[\"feature_\"+str(i) for i in range(1,embeds_dim+1)])\n# df_test_embeds = pd.DataFrame(t5_test_embeds,columns=[\"feature_\"+str(i) for i in range(1,embeds_dim+1)])\ndf_train_ids = pd.DataFrame(t5_train_ids,columns=[\"EntryID\"])\n# df_test_ids = pd.DataFrame(t5_test_ids,columns=[\"EntryID\"])\n\n# Check the dataframes.\n# print(df_train_embeds.head())\n# print(df_test_embeds.head())\nprint(df_train_ids.head())\n# print(df_test_ids.head())\nprint(df_train_terms.head())\nprint(df_ia_weight.head())","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:10:48.820976Z","iopub.execute_input":"2025-04-06T09:10:48.821287Z","iopub.status.idle":"2025-04-06T09:11:52.120503Z","shell.execute_reply.started":"2025-04-06T09:10:48.821257Z","shell.execute_reply":"2025-04-06T09:11:52.119699Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Prepare the label y of trainset.\n# Each index of the labels must be as the same as the features‘.\n\n# Check the total number of the GO terms\nprint(\"The total number of the GO terms:\")\nprint(df_train_terms.term.nunique())\n\n# Check the count number of the terms values. \nvalue_count = df_train_terms.term.value_counts()\nprint(value_count)\n\n# Confirm the number to be 1500, and use it to select the GO terms of the first 1500 in the value_count series.\nselect_value_count = value_count[0:1500]\nprint(\"The count number of the selected GO terms:\")\nprint(select_value_count)\n\n# Put the first 1600 high frequency terms into a list named \"select_terms\".\nselect_terms = select_value_count.index.tolist()\n# print(select_terms)\nprint(len(select_terms))\n\n# Drop the terms that haven't be chosen in the 1500, and create the selected terms dataframe.\ndf_select_terms = df_train_terms.loc[df_train_terms.term.isin(select_terms)]\nprint(len(df_select_terms))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:11:52.121341Z","iopub.execute_input":"2025-04-06T09:11:52.121557Z","iopub.status.idle":"2025-04-06T09:11:53.222298Z","shell.execute_reply.started":"2025-04-06T09:11:52.121539Z","shell.execute_reply":"2025-04-06T09:11:53.221382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get all the GO terms in the select_terms each protein has. \nseries_term_lst = df_select_terms.groupby(\"EntryID\").term.unique()\nprint(series_term_lst[0:5])\nprint(type(series_term_lst)) \n\n# Transform the series into a dataframe with the index reset.\ndf_label_obj = pd.DataFrame(series_term_lst).reset_index()\nprint(df_label_obj)\n\n# Merge the dataframes to sort by the order of df_train_ids ,corresponding to df_train_embeds.\ndf_train_ordered = df_train_ids.merge(df_label_obj,on=\"EntryID\",how=\"left\")\nprint(df_train_ordered)\nprint(df_train_ordered[\"term\"])\n\n# Use the MultiLAbelBinarizer to do multi binary classifacation.\nmyMLB = MultiLabelBinarizer()\ndf_train_encode_array = myMLB.fit_transform(df_train_ordered[\"term\"])\ndf_train_encode = pd.DataFrame(df_train_encode_array,columns=myMLB.classes_)\nlabel_order = myMLB.classes_\n\nprint(df_train_encode_array)\nprint(df_train_encode)\nprint(df_train_encode.loc[0,\"GO:0044249\"])\nprint(label_order)\n\n# Get the list of the IA weight in labels' order.\nlst_label_weight = []\nfor label in label_order:\n    label_weight = df_ia_weight.loc[df_ia_weight[\"terms\"] == label].ia_weight.values[0]\n    lst_label_weight.append(label_weight)\nprint(len(lst_label_weight))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:11:53.223206Z","iopub.execute_input":"2025-04-06T09:11:53.223468Z","iopub.status.idle":"2025-04-06T09:12:06.75579Z","shell.execute_reply.started":"2025-04-06T09:11:53.223446Z","shell.execute_reply":"2025-04-06T09:12:06.754995Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Prepare the dataset.\n\n# load the numpy ndarray.\nX_np = t5_train_embeds\ny_np = df_train_encode_array\ntest_np = t5_test_embeds\n\n# Transform them into tensor.\nX_tensor = torch.from_numpy(X_np).float()\ny_tensor = torch.from_numpy(y_np).float()\ntest_tensor = torch.from_numpy(test_np).float()\nweights_tensor = torch.tensor(lst_label_weight,dtype=torch.float32)\n\n# Create the dataset\ndataset_total = TensorDataset(X_tensor,y_tensor)\n\n# Split the dataset into train set and validation set.\ntrain_size = int(0.8 * len(dataset_total))\nvalid_size = len(dataset_total) - train_size\ntrain_set,valid_set = random_split(dataset_total,[train_size,valid_size])\n\n# Prepare the whole trainset and testset.\ntrain_set_whole = dataset_total\ntest_set = test_tensor\n\nprint(weights_tensor)\nprint(\"Loading successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:12:06.757969Z","iopub.execute_input":"2025-04-06T09:12:06.758222Z","iopub.status.idle":"2025-04-06T09:12:07.792647Z","shell.execute_reply.started":"2025-04-06T09:12:06.758201Z","shell.execute_reply":"2025-04-06T09:12:07.79185Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set up the models.\nnetwork = nn.Sequential(\n    nn.Linear(1024,512),nn.ReLU(),\n    nn.Linear(512,512),nn.ReLU(),\n    nn.Linear(512,512),nn.ReLU(),\n    nn.Linear(512,512),nn.ReLU(),\n    nn.Linear(512,1500)\n)\n\nnetwork_dropout = nn.Sequential(\n    nn.Linear(1024,1024),nn.ReLU(),\n    nn.Linear(1024,1024),nn.ReLU(),nn.Dropout(0.5),\n    nn.Linear(1024,1024),nn.ReLU(),nn.Dropout(0.5),\n    nn.Linear(1024,1024),nn.ReLU(),nn.Dropout(0.5),\n    nn.Linear(1024,1500)\n)\n\nnetwork_deep_drop = nn.Sequential(\n    nn.Linear(1024,512),nn.ReLU(),nn.Dropout(0.2),\n    nn.Linear(512,512),nn.ReLU(),nn.Dropout(0.2),\n    nn.Linear(512,512),nn.ReLU(),nn.Dropout(0.2),\n    nn.Linear(512,512),nn.ReLU(),nn.Dropout(0.3),\n    nn.Linear(512,512),nn.ReLU(),nn.Dropout(0.3),\n    nn.Linear(512,1500)\n)\n\nnetwork_deep_drop_bt_norm_0 = nn.Sequential(\n    nn.Linear(1024,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.2),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.2),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.2),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.3),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.3),\n    nn.Linear(512,1500)\n)\n\nnetwork_deep_drop_bt_norm_1 = nn.Sequential(\n    nn.Linear(1024,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.2),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.2),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.2),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.3),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.3),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.3),\n    nn.Linear(512,1500)\n)\n\nnetwork_deep_drop_bt_norm_2 = nn.Sequential(\n    nn.Linear(1024,512),nn.BatchNorm1d(512),nn.ReLU(),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.5),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.5),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.5),\n    nn.Linear(512,512),nn.BatchNorm1d(512),nn.ReLU(),nn.Dropout(0.5),\n    nn.Linear(512,1500)\n)\n\nprint(\"Setting successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:12:07.793734Z","iopub.execute_input":"2025-04-06T09:12:07.794029Z","iopub.status.idle":"2025-04-06T09:12:07.950342Z","shell.execute_reply.started":"2025-04-06T09:12:07.794007Z","shell.execute_reply":"2025-04-06T09:12:07.949524Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define a loss function to using the IA weights.\nclass ia_weighted_BCEWithLogitsLoss(nn.Module):\n    def __init__(self,weights):\n        super().__init__()\n        self.weights = weights\n\n    def forward(self,logits,targets):\n        ele_loss = nn.functional.binary_cross_entropy_with_logits(logits,targets,reduction=\"none\")\n        weighted_loss = ele_loss * self.weights + ele_loss\n        return weighted_loss.mean()\n\n# Define a function to get the accuracy of the batch.\ndef get_accuracy(y_hat,y,threshold=0.5):\n    y_probs = torch.sigmoid(y_hat)\n    y_pred_label = (y_probs >= threshold).float()\n    num_acc = (y_pred_label == y).sum().item()\n    return num_acc / y.numel()\n\n# Define a function to get the loss and accuracy.\ndef get_loss_and_acc(l,X,y,y_hat,running_loss,total_correct,total_labels):\n    running_loss += l.item() * X.size(0)\n    batch_accuracy = get_accuracy(y_hat,y,threshold=0.5)\n    total_correct += batch_accuracy * y.numel()\n    total_labels += y.numel()\n    return running_loss,total_correct,total_labels\n    \n# Print the results of loss and accuracy.\ndef get_loss_accuracy(running_loss,total_correct,total_labels,data_set,epoch):\n    epoch_loss = running_loss / len(data_set)\n    epoch_acc = 100 * total_correct / total_labels\n    return epoch_loss,epoch_acc\n\nprint(\"Functions loading successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:12:07.951103Z","iopub.execute_input":"2025-04-06T09:12:07.951339Z","iopub.status.idle":"2025-04-06T09:12:07.958666Z","shell.execute_reply.started":"2025-04-06T09:12:07.951319Z","shell.execute_reply":"2025-04-06T09:12:07.95776Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the function to train and score.\ndef train_and_score(net,train_set,valid_set,batch_size,weight_decay,num_epochs,learning_rate,):\n    # Create dataloader\n    train_loader = DataLoader(train_set,batch_size=batch_size,shuffle=True)\n    valid_loader = DataLoader(valid_set,batch_size=batch_size,shuffle=False)\n\n    # Set updater\n    updater = optim.Adam(net.parameters(),lr=learning_rate,weight_decay=weight_decay)\n\n    # Set Loss function\n    loss = ia_weighted_BCEWithLogitsLoss(weights_tensor)\n\n    # Set the loss and accuracy list.\n    train_losses,train_accs,valid_losses,valid_accs = [],[],[],[]\n            \n    # Train and valid.\n    for epoch in range(num_epochs):\n        \n        # Set the initial value.\n        running_loss = 0\n        total_correct = 0\n        total_labels = 0\n        \n        # Start to train.\n        net.train()\n        for X,y in train_loader:\n            updater.zero_grad()\n            y_hat = net(X)\n            l = loss(y_hat,y)\n            l.backward()\n            updater.step()\n            running_loss,total_correct,total_labels = get_loss_and_acc(l,X,y,y_hat,running_loss,total_correct,total_labels)\n        train_loss,train_acc = get_loss_accuracy(running_loss,total_correct,total_labels,train_set,epoch)\n        train_losses.append(train_loss)\n        train_accs.append(train_acc)\n\n        # Reset the initial value.\n        running_loss = 0\n        total_correct = 0\n        total_labels = 0\n        \n        # start to valid.\n        net.eval()\n        with torch.no_grad():\n            for X,y in valid_loader:\n                y_hat = net(X)\n                l = loss(y_hat,y)\n                running_loss,total_correct,total_labels = get_loss_and_acc(l,X,y,y_hat,running_loss,total_correct,total_labels)\n        valid_loss,valid_acc = get_loss_accuracy(running_loss,total_correct,total_labels,valid_set,epoch)\n        valid_losses.append(valid_loss)\n        valid_accs.append(valid_acc)\n\n        print(f\"Epoch: {epoch + 1} / {num_epochs}\")\n        print(f\"Train set: Loss: {train_loss:.4f} | Accuracy: {train_acc:.3f}%\")\n        print(f\"Valid set: Loss: {valid_loss:.4f} | Accuracy: {valid_acc:.3f}%\")\n        print(\"-\" * 50)\n\n    # Set the plot.\n    fig,(ax1,ax2) = plt.subplots(1,2,figsize=(12,5))\n    \n    x_data = list(range(1,num_epochs + 1))\n    \n    ax1.plot(x_data,train_losses,\"b-\",label=\"Training loss\")\n    ax1.plot(x_data,valid_losses,\"r--\",label=\"Validating loss\")\n    ax1.set_title(\"Loss Curve\")\n    ax1.set_xlabel(\"Epoch\")\n    ax1.set_ylabel(\"Loss\")\n    ax1.legend()\n\n    ax2.plot(x_data,train_accs,\"g-\",label=\"Training accuracy\")\n    ax2.plot(x_data,valid_accs,\"m--\",label=\"Validating accuracy\")\n    ax2.set_title(\"Accuracy Curve\")\n    ax2.set_xlabel(\"Epoch\")\n    ax2.set_ylabel(\"Accuracy\")\n    ax2.legend()\n        \n    plt.tight_layout()\n    \n    plt.draw()\n    plt.pause(0.1)\n\nprint(\"Function loading successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:12:07.959505Z","iopub.execute_input":"2025-04-06T09:12:07.95975Z","iopub.status.idle":"2025-04-06T09:12:07.977863Z","shell.execute_reply.started":"2025-04-06T09:12:07.959719Z","shell.execute_reply":"2025-04-06T09:12:07.977012Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set num_epochs.\nnum_epochs = 30\n# Set learning_rate.\nlearning_rate = 0.001\n# set weight decay range in (0.0001-0.01)\nweight_decay = 0.00001\n# Set batch size\nbatch_size = 5120\n\nprint(\"Loading successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:12:07.978608Z","iopub.execute_input":"2025-04-06T09:12:07.978891Z","iopub.status.idle":"2025-04-06T09:12:07.997534Z","shell.execute_reply.started":"2025-04-06T09:12:07.978862Z","shell.execute_reply":"2025-04-06T09:12:07.996651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test the model.\n# Comment the models that have not been selected.\n# train_and_score(\n#     network,train_set,valid_set,batch_size=batch_size,weight_decay=weight_decay,\\\n#     num_epochs=20,learning_rate=learning_rate,\n# )\n\n# train_and_score(\n#     network_dropout,train_set,valid_set,batch_size=batch_size,weight_decay=weight_decay,\\\n#     num_epochs=20,learning_rate=learning_rate,\n# )\n\n# train_and_score(\n#     network_deep_drop,train_set,valid_set,batch_size=batch_size,weight_decay=weight_decay,\\\n#     num_epochs=20,learning_rate=learning_rate,\n# )","metadata":{"execution":{"iopub.status.busy":"2025-04-06T09:12:07.998362Z","iopub.execute_input":"2025-04-06T09:12:07.998663Z","iopub.status.idle":"2025-04-06T09:12:08.012471Z","shell.execute_reply.started":"2025-04-06T09:12:07.998634Z","shell.execute_reply":"2025-04-06T09:12:08.011798Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_and_score(\n#     network_deep_drop_bt_norm_0,train_set,valid_set,batch_size=batch_size,weight_decay=weight_decay,\\\n#     num_epochs=30,learning_rate=learning_rate,\n# )\n\n# train_and_score(\n#     network_deep_drop_bt_norm_1,train_set,valid_set,batch_size=batch_size,weight_decay=weight_decay,\\\n#     num_epochs=30,learning_rate=learning_rate,\n# )\n\n# train_and_score(\n#     network_deep_drop_bt_norm_2,train_set,valid_set,batch_size=batch_size,weight_decay=weight_decay,\\\n#     num_epochs=30,learning_rate=learning_rate,\n# )","metadata":{"execution":{"iopub.status.busy":"2025-04-06T09:12:08.013296Z","iopub.execute_input":"2025-04-06T09:12:08.013542Z","iopub.status.idle":"2025-04-06T09:12:08.026951Z","shell.execute_reply.started":"2025-04-06T09:12:08.013523Z","shell.execute_reply":"2025-04-06T09:12:08.026376Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Use it to train the whole trainset.\ndef train_and_pred(net,train_set,test_set,batch_size,weight_decay,num_epochs,learning_rate,):\n    # Create dataloader\n    train_loader = DataLoader(train_set,batch_size=batch_size,shuffle=True)\n    \n    # Set updater\n    updater = optim.Adam(net.parameters(),lr=learning_rate,weight_decay=weight_decay)\n\n    # Set Loss function\n    loss = ia_weighted_BCEWithLogitsLoss(weights_tensor)\n\n    # Set the loss and accuracy list.\n    train_losses,train_accs = [],[]\n            \n    # Train and valid.\n    for epoch in range(num_epochs):\n        \n        # Set the initial value.\n        running_loss = 0\n        total_correct = 0\n        total_labels = 0\n        \n        # Start to train.\n        net.train()\n        for X,y in train_loader:\n            updater.zero_grad()\n            y_hat = net(X)\n            l = loss(y_hat,y)\n            l.backward()\n            updater.step()\n            running_loss,total_correct,total_labels = get_loss_and_acc(l,X,y,y_hat,running_loss,total_correct,total_labels)\n        train_loss,train_acc = get_loss_accuracy(running_loss,total_correct,total_labels,train_set,epoch)\n        train_losses.append(train_loss)\n        train_accs.append(train_acc)\n\n        print(f\"Epoch: {epoch + 1} / {num_epochs}\")\n        print(f\"Train set: Loss: {train_loss:.4f} | Accuracy: {train_acc:.3f}%\")\n        print(\"-\" * 50)\n\n    # Set the plot.\n    fig,(ax1,ax2) = plt.subplots(1,2,figsize=(12,5))\n    \n    x_data = list(range(1,num_epochs + 1))\n    \n    ax1.plot(x_data,train_losses,\"b-\",label=\"Training loss\")\n    ax1.set_title(\"Loss Curve\")\n    ax1.set_xlabel(\"Epoch\")\n    ax1.set_ylabel(\"Loss\")\n    ax1.legend()\n\n    ax2.plot(x_data,train_accs,\"g-\",label=\"Training accuracy\")\n    ax2.set_title(\"Accuracy Curve\")\n    ax2.set_xlabel(\"Epoch\")\n    ax2.set_ylabel(\"Accuracy\")\n    ax2.legend()\n        \n    plt.tight_layout()\n    \n    plt.draw()\n    plt.pause(0.1)\n    \n    # start to valid.\n    net.eval()\n    with torch.no_grad():\n        test_preds = net(test_set)\n        predictions = torch.sigmoid(test_preds)\n    return predictions\n\nprint(\"Function loading successfully.\") ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:12:08.027789Z","iopub.execute_input":"2025-04-06T09:12:08.02801Z","iopub.status.idle":"2025-04-06T09:12:08.047082Z","shell.execute_reply.started":"2025-04-06T09:12:08.027991Z","shell.execute_reply":"2025-04-06T09:12:08.046459Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# The best model is network_deep_drop_bt_norm_0, choose it.\nfinal_model = network_deep_drop_bt_norm_1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:12:08.047719Z","iopub.execute_input":"2025-04-06T09:12:08.047984Z","iopub.status.idle":"2025-04-06T09:12:08.068887Z","shell.execute_reply.started":"2025-04-06T09:12:08.047964Z","shell.execute_reply":"2025-04-06T09:12:08.068134Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train and get the final prediction.\nfinal_predictions = train_and_pred(final_model,train_set_whole,test_set,batch_size,weight_decay,num_epochs=50,learning_rate=0.001,)","metadata":{"execution":{"iopub.status.busy":"2025-04-06T09:12:08.06962Z","iopub.execute_input":"2025-04-06T09:12:08.069891Z","iopub.status.idle":"2025-04-06T09:24:07.204915Z","shell.execute_reply.started":"2025-04-06T09:12:08.06987Z","shell.execute_reply":"2025-04-06T09:24:07.203982Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Put the predictions into a dataframe by pandas.\nprint(final_predictions.shape)\ndf_final_pred = pd.DataFrame(final_predictions,index=t5_test_ids,columns=label_order)\nprint(df_final_pred.head())\n\n# Transform the dataframe into the shape that we want.\ndf_final_reset = df_final_pred.reset_index().rename(columns={\"index\":\"Protein Id\"})\ndf_submit = df_final_reset.melt(\n    id_vars=\"Protein Id\",\n    value_vars=label_order,\n    var_name=\"GO Terms\",\n    value_name=\"Prediction\"\n)\n\nprint(df_final_reset.head())\nprint(df_submit.head())\nprint(df_submit.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:24:07.205964Z","iopub.execute_input":"2025-04-06T09:24:07.206479Z","iopub.status.idle":"2025-04-06T09:24:24.817002Z","shell.execute_reply.started":"2025-04-06T09:24:07.206447Z","shell.execute_reply":"2025-04-06T09:24:24.81616Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# save the submission file to output.\ndf_submit.to_csv(\"submission.tsv\",header=False,index=False,sep=\"\\t\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-06T09:24:24.817947Z","iopub.execute_input":"2025-04-06T09:24:24.818294Z"}},"outputs":[],"execution_count":null}]}