{"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":"# Multitarget(binary) Classification with XGBoost","metadata":{}},{"cell_type":"markdown","source":"This notebook first converts the multilabel problem into multitarget binary classification problem. Then the scikit learn's implementatoin of the XGBoost model is used for this problem. The input to the model needs to be feature vectors extarcted from the protein sequences provided by the train_sequences.fasta files. The T5 embeddings generated by Grandmaster Sergei Fironov are used as the feature vectors.","metadata":{}},{"cell_type":"markdown","source":"Only the top 1499 most occouring labels are consdierd otherwise the model cannot train for 40k different classes","metadata":{}},{"cell_type":"code","source":"n_labels_to_consider = 1499 # We will choose only top frequent labels (in train) and predict only them. \nn_max_preds = 1499","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-07T21:05:34.915296Z","iopub.execute_input":"2023-07-07T21:05:34.915659Z","iopub.status.idle":"2023-07-07T21:05:34.919785Z","shell.execute_reply.started":"2023-07-07T21:05:34.915626Z","shell.execute_reply":"2023-07-07T21:05:34.91887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\nt0start = time.time() \n\nimport numpy as np\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\nfrom sklearn.model_selection import train_test_split, cross_val_score, GridSearchCV, KFold, RandomizedSearchCV\nfrom sklearn.linear_model import Ridge,RidgeCV\nfrom sklearn.neural_network import MLPClassifier\nfrom sklearn.multioutput import MultiOutputClassifier\n\nimport xgboost as xgb\n\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-07-07T21:05:36.090349Z","iopub.execute_input":"2023-07-07T21:05:36.09134Z","iopub.status.idle":"2023-07-07T21:05:37.045665Z","shell.execute_reply.started":"2023-07-07T21:05:36.091299Z","shell.execute_reply":"2023-07-07T21:05:37.044732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrainTerms = pd.read_csv(\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv\",sep=\"\\t\")\nprint(trainTerms.shape)\ndisplay(trainTerms.head(2))\nvec_freqCount = (trainTerms['term'].value_counts())\nprint(vec_freqCount )","metadata":{"execution":{"iopub.status.busy":"2023-07-07T21:05:37.716542Z","iopub.execute_input":"2023-07-07T21:05:37.717265Z","iopub.status.idle":"2023-07-07T21:05:41.803881Z","shell.execute_reply.started":"2023-07-07T21:05:37.717227Z","shell.execute_reply":"2023-07-07T21:05:41.802769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## drop very rarely occouring GO terms\nvec_freqCount = vec_freqCount[vec_freqCount>=30]\nprint(vec_freqCount.shape[0])\nvec_freqCount.describe().round()","metadata":{"execution":{"iopub.status.busy":"2023-07-07T21:05:41.805928Z","iopub.execute_input":"2023-07-07T21:05:41.806366Z","iopub.status.idle":"2023-07-07T21:05:41.825415Z","shell.execute_reply.started":"2023-07-07T21:05:41.806323Z","shell.execute_reply":"2023-07-07T21:05:41.824381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vec_freqCount[vec_freqCount>200].shape[0]","metadata":{"execution":{"iopub.status.busy":"2023-07-07T21:05:41.826948Z","iopub.execute_input":"2023-07-07T21:05:41.827405Z","iopub.status.idle":"2023-07-07T21:05:41.837444Z","shell.execute_reply.started":"2023-07-07T21:05:41.827371Z","shell.execute_reply":"2023-07-07T21:05:41.836375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# consider only top 1500 labels\nlabels_to_consider = list(vec_freqCount.index[:n_labels_to_consider] )\nprint('n_labels_to_consider:', len(labels_to_consider), 'First 10:', labels_to_consider[:10] ) ","metadata":{"execution":{"iopub.status.busy":"2023-07-07T21:05:41.840504Z","iopub.execute_input":"2023-07-07T21:05:41.84137Z","iopub.status.idle":"2023-07-07T21:05:41.847687Z","shell.execute_reply.started":"2023-07-07T21:05:41.841336Z","shell.execute_reply":"2023-07-07T21:05:41.846739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load the T5 emebeddings\nfn = '/kaggle/input/t5embeds/train_ids.npy'\nvec_train_protein_ids = np.load(fn)\nprint(vec_train_protein_ids.shape)\nvec_train_protein_ids","metadata":{"execution":{"iopub.status.busy":"2023-07-07T21:05:50.068105Z","iopub.execute_input":"2023-07-07T21:05:50.068457Z","iopub.status.idle":"2023-07-07T21:05:50.120598Z","shell.execute_reply.started":"2023-07-07T21:05:50.068431Z","shell.execute_reply":"2023-07-07T21:05:50.119732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Convert the multibinary classification into multi target and store it in form a 2D matrix where cell (i,j) is 1 if i-th protein is associated with j-th label, 0 otherwise","metadata":{}},{"cell_type":"code","source":"%%time \ntrain_size = 142246 # len(X)\nY = np.zeros( (train_size ,n_labels_to_consider) )\nprint(Y.shape)\n\nseries_train_protein_ids = pd.Series(vec_train_protein_ids ) # \n\ntrainTerms_smaller = trainTerms[ trainTerms['term'].isin( labels_to_consider ) ] # to speed-up the next step \nprint( trainTerms_smaller.shape)\n\nfor i in range(Y.shape[1]):\n    m = trainTerms_smaller['term'] ==  labels_to_consider[i]\n#     m.sum()\n    Y[:,i] =  series_train_protein_ids.isin(  set(trainTerms_smaller[m]['EntryID'] ) ).astype(float )\n    if (i % 10) == 0: \n        print(i, m.sum())\nY","metadata":{"execution":{"iopub.status.busy":"2023-07-07T21:05:51.845267Z","iopub.execute_input":"2023-07-07T21:05:51.845615Z","iopub.status.idle":"2023-07-07T21:23:20.927038Z","shell.execute_reply.started":"2023-07-07T21:05:51.845585Z","shell.execute_reply":"2023-07-07T21:23:20.925881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Generating the above matrix takes time, so it is saved in file to be loaded and used later","metadata":{}},{"cell_type":"code","source":"%%time \n# save for possible future reuse \nfn4saveY = 'Y_'+str(Y.shape[1])\nprint(fn4saveY)\nnp.save( fn4saveY , Y) ","metadata":{"execution":{"iopub.status.busy":"2023-06-19T07:22:32.849376Z","iopub.execute_input":"2023-06-19T07:22:32.849752Z","iopub.status.idle":"2023-06-19T07:22:34.289652Z","shell.execute_reply.started":"2023-06-19T07:22:32.849719Z","shell.execute_reply":"2023-06-19T07:22:34.288485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### saving the top 1500 labels for future reuse","metadata":{}},{"cell_type":"code","source":"%%time\nfn4save_labels = 'Y_'+str(Y.shape[1]) + '_labels'\nnp.save(fn4save_labels, labels_to_consider )","metadata":{"execution":{"iopub.status.busy":"2023-05-29T04:25:15.894603Z","iopub.execute_input":"2023-05-29T04:25:15.900754Z","iopub.status.idle":"2023-05-29T04:25:15.909897Z","shell.execute_reply.started":"2023-05-29T04:25:15.900708Z","shell.execute_reply":"2023-05-29T04:25:15.908535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load the train embeds","metadata":{}},{"cell_type":"code","source":"%%time\nX = np.load('/kaggle/input/t5embeds/train_embeds.npy')\nprint(X.shape)\nX","metadata":{"execution":{"iopub.status.busy":"2023-06-04T08:39:07.856548Z","iopub.execute_input":"2023-06-04T08:39:07.857018Z","iopub.status.idle":"2023-06-04T08:39:19.374213Z","shell.execute_reply.started":"2023-06-04T08:39:07.856974Z","shell.execute_reply":"2023-06-04T08:39:19.373226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### The corresponding protein ids are:","metadata":{}},{"cell_type":"code","source":"%%time\nvec_train_protein_ids = np.load('/kaggle/input/t5embeds/train_ids.npy')\nprint(vec_train_protein_ids.shape)\nvec_train_protein_ids","metadata":{"execution":{"iopub.status.busy":"2023-05-29T07:58:49.078473Z","iopub.execute_input":"2023-05-29T07:58:49.079616Z","iopub.status.idle":"2023-05-29T07:58:49.093572Z","shell.execute_reply.started":"2023-05-29T07:58:49.079578Z","shell.execute_reply":"2023-05-29T07:58:49.092322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create the train test split","metadata":{}},{"cell_type":"code","source":"IX = np.arange(len(X))\nprint(IX.shape)\nprint(IX)\nIX_train, IX_test, _,_ = train_test_split( IX, IX, train_size=0.1, random_state=42)\n# print(len(IX_train), len(IX_test),  IX_train[:10], IX_test[:10] )","metadata":{"execution":{"iopub.status.busy":"2023-05-29T07:58:49.105587Z","iopub.execute_input":"2023-05-29T07:58:49.106266Z","iopub.status.idle":"2023-05-29T07:58:49.12124Z","shell.execute_reply.started":"2023-05-29T07:58:49.106232Z","shell.execute_reply":"2023-05-29T07:58:49.120117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training the Model","metadata":{}},{"cell_type":"code","source":"clf_xgb = xgb.XGBClassifier(objective=\"binary:logistic\", random_state=42, tree_method=\"gpu_hist\", verbosity=2)","metadata":{"execution":{"iopub.status.busy":"2023-05-29T07:58:49.122961Z","iopub.execute_input":"2023-05-29T07:58:49.123305Z","iopub.status.idle":"2023-05-29T07:58:49.128597Z","shell.execute_reply.started":"2023-05-29T07:58:49.123274Z","shell.execute_reply":"2023-05-29T07:58:49.127703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"clf_xgb.fit(X[IX_train,:], Y[IX_train,:])","metadata":{"execution":{"iopub.status.busy":"2023-05-29T07:58:49.140159Z","iopub.execute_input":"2023-05-29T07:58:49.1406Z","iopub.status.idle":"2023-05-29T08:34:12.289874Z","shell.execute_reply.started":"2023-05-29T07:58:49.140566Z","shell.execute_reply":"2023-05-29T08:34:12.289029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Predictions","metadata":{}},{"cell_type":"code","source":"y_pred_test = clf_xgb.predict(X[IX_train[:10],:])","metadata":{"execution":{"iopub.status.busy":"2023-05-29T08:34:12.294011Z","iopub.execute_input":"2023-05-29T08:34:12.296098Z","iopub.status.idle":"2023-05-29T08:34:12.403572Z","shell.execute_reply.started":"2023-05-29T08:34:12.296047Z","shell.execute_reply":"2023-05-29T08:34:12.402646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(y_pred_test)","metadata":{"execution":{"iopub.status.busy":"2023-05-29T08:34:12.406277Z","iopub.execute_input":"2023-05-29T08:34:12.410753Z","iopub.status.idle":"2023-05-29T08:34:12.421021Z","shell.execute_reply.started":"2023-05-29T08:34:12.410719Z","shell.execute_reply":"2023-05-29T08:34:12.420073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Save the model for testing and parameter tuning","metadata":{}},{"cell_type":"code","source":"clf_xgb.save_model(\"model.json\")","metadata":{"execution":{"iopub.status.busy":"2023-05-29T16:18:29.295404Z","iopub.execute_input":"2023-05-29T16:18:29.296071Z","iopub.status.idle":"2023-05-29T16:18:29.70451Z","shell.execute_reply.started":"2023-05-29T16:18:29.296019Z","shell.execute_reply":"2023-05-29T16:18:29.703079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}