{"metadata":{"kaggle":{"accelerator":"none","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.13"},"papermill":{"default_parameters":{},"duration":2868.446311,"end_time":"2024-06-23T22:36:51.047483","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-06-23T21:49:02.601172","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"---\n\nThis is a second version, based on a different feature extraction strategy, of the original work as per the following notebook:\n\nhttps://www.kaggle.com/code/aiaiaidavid/belka-model-training-and-evaluation-descriptors\n\n---\n\n**THE PROBLEM**\n\nThree binary targets (protein binding affinity):\n- sEH\n- BRD4\n- HSA (ALB)\n\nFeatures will be generated as a sparse matrix of 1s and 0s from the molecules (given in SMILES format) through domain specific library computations (in this version with the so-called ECFP values).\n\n---\n\n**VERSIONS**\n\nv23  1 Million, ECFP 128, full hyper-param grid, reduced int8 features\n\nv21  800K, ECFP 300, full hyper-param grid, reduced int8 features - GAME OVER.\n\nv20  900K, ECFP 300, full hyper-param grid, reduced int8 features - TIMED OUT (12 hour+)\n\nv19  900K, ECFP 512, full hyper-param grid, reduced int8 features - FAILED (OOM)\n\nv18  900K, ECFP 512, full hyper-param grid, reduced int8 features - FAILED (OOM)\n\nv17  900K, ECFP 275 full hyper-param grid - CANCELLED\n\nv16  1 Million, ECFP 300 full hyper-param grid - FAILED (OOM)\n\nv15  1 Million (only 408K for HSA), ECFP 256 full hyper-param grid. Trying with radius = 3 - FAILED (OOM).\n\nV13  800K ECFP 256 full hyper-param grid.\n\nv12  700K 512 ECFP full hyper-param grid - FAILED (OOM)\n\nv11  700K 256 ECFP full hyper-param grid.\n\nv10  Going even bigger at 600K and ECFP 256 (with reduced selection grid)\n\nv9   Trying 500K 256 ECFP\n\nv8   Reducing Optuna grid - FAILED (OOM)\n\nv7   800K and 256 ECFP - FAILED (OOM)\n\nv6   Same as v4-5 (lost the count... something did not work last time)\n\nv4   Going big with 400K training per proteine and 256 ECFP.\n\nv3   Full run, saves models trained with 100K sample and 128 ECFP features.\n\nv2   Full run, saves models trained with 50K sample and 128 ECFP features.\n\nv1   First version with 5K sample size just to validate the pipeline all through submission.\n\n---","metadata":{"papermill":{"duration":0.008704,"end_time":"2024-06-23T21:49:05.271262","exception":false,"start_time":"2024-06-23T21:49:05.262558","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Environment setup, package installs, library imports\n---","metadata":{"papermill":{"duration":0.007949,"end_time":"2024-06-23T21:49:05.287783","exception":false,"start_time":"2024-06-23T21:49:05.279834","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install duckdb","metadata":{"papermill":{"duration":14.083403,"end_time":"2024-06-23T21:49:19.379497","exception":false,"start_time":"2024-06-23T21:49:05.296094","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-06T05:19:40.252095Z","iopub.execute_input":"2024-07-06T05:19:40.252698Z","iopub.status.idle":"2024-07-06T05:19:53.286166Z","shell.execute_reply.started":"2024-07-06T05:19:40.252653Z","shell.execute_reply":"2024-07-06T05:19:53.284639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install rdkit","metadata":{"papermill":{"duration":14.43603,"end_time":"2024-06-23T21:49:33.824832","exception":false,"start_time":"2024-06-23T21:49:19.388802","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-06T05:19:53.288788Z","iopub.execute_input":"2024-07-06T05:19:53.289245Z","iopub.status.idle":"2024-07-06T05:20:06.017706Z","shell.execute_reply.started":"2024-07-06T05:19:53.289207Z","shell.execute_reply":"2024-07-06T05:20:06.016297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install optuna-integration","metadata":{"papermill":{"duration":0.019279,"end_time":"2024-06-23T21:49:33.854688","exception":false,"start_time":"2024-06-23T21:49:33.835409","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-06T05:20:06.019562Z","iopub.execute_input":"2024-07-06T05:20:06.020004Z","iopub.status.idle":"2024-07-06T05:20:06.025626Z","shell.execute_reply.started":"2024-07-06T05:20:06.019959Z","shell.execute_reply":"2024-07-06T05:20:06.024424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport warnings\nimport duckdb\nfrom rdkit.Chem import AllChem, MolFromSmiles, rdFingerprintGenerator\nfrom rdkit.Chem import rdMolDescriptors, Descriptors, rdmolfiles, MolFromPDBFile\nimport rdkit\nimport numpy as np \nimport pandas as pd \nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.model_selection import cross_val_score, GridSearchCV\nfrom sklearn.metrics import average_precision_score\nfrom sklearn.metrics.pairwise import cosine_similarity\nfrom scipy.sparse import csr_matrix\nimport xgboost as xgb\nfrom xgboost import XGBClassifier\nimport shap\nimport optuna\n# from optuna.integration import XGBoostPruningCallback\n\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":8.346228,"end_time":"2024-06-23T21:49:42.212112","exception":false,"start_time":"2024-06-23T21:49:33.865884","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-06T05:20:06.028522Z","iopub.execute_input":"2024-07-06T05:20:06.028939Z","iopub.status.idle":"2024-07-06T05:20:06.040688Z","shell.execute_reply.started":"2024-07-06T05:20:06.028901Z","shell.execute_reply":"2024-07-06T05:20:06.039628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Functions definitions and general settings\n---","metadata":{"papermill":{"duration":0.009865,"end_time":"2024-06-23T21:49:42.240160","exception":false,"start_time":"2024-06-23T21:49:42.230295","status":"completed"},"tags":[]}},{"cell_type":"code","source":"proteins = ['sEH', 'BRD4', 'HSA']\n\nchemical_radius = 2\nchemical_nbits = 128\n\ntrain_file = \"/kaggle/input/leash-BELKA/train.parquet\"\ntest_file = \"/kaggle/input/leash-BELKA/test.csv\"\nsub_file = \"/kaggle/input/leash-BELKA/sample_submission.csv\"\n\n# pre_processed_train_data_path = '/kaggle/input/belka-900k-pre-processed-train-data-and-models/'","metadata":{"papermill":{"duration":0.019441,"end_time":"2024-06-23T21:49:42.269757","exception":false,"start_time":"2024-06-23T21:49:42.250316","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-06T05:20:06.042073Z","iopub.execute_input":"2024-07-06T05:20:06.042648Z","iopub.status.idle":"2024-07-06T05:20:06.051654Z","shell.execute_reply.started":"2024-07-06T05:20:06.042618Z","shell.execute_reply":"2024-07-06T05:20:06.050414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Generate ECFPs\ndef generate_ecfp(molecule, radius=chemical_radius, bits=chemical_nbits):\n    if molecule is None:\n        return None\n    else:\n        morgan_fp_generator = rdFingerprintGenerator.GetMorganGenerator(radius = chemical_radius, fpSize = chemical_nbits)\n        return morgan_fp_generator.GetFingerprint(molecule)","metadata":{"execution":{"iopub.status.busy":"2024-07-06T05:20:06.053307Z","iopub.execute_input":"2024-07-06T05:20:06.053867Z","iopub.status.idle":"2024-07-06T05:20:06.063711Z","shell.execute_reply.started":"2024-07-06T05:20:06.053825Z","shell.execute_reply":"2024-07-06T05:20:06.062465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data loading and processing\n---\n\nThe train dataset is way too big so it cannot be read in full neither as csv nor parquet. It can be loaded from parquet using a SQL query with the DuckDB, which is very convenient as it allows to specify the number of samples read.\n\nThe code block below allows to be run in two modes:\n - original raw data for full processing, or\n - pre-processed training data (with all features already generated) from an earlier execution.\n ","metadata":{"papermill":{"duration":0.00967,"end_time":"2024-06-23T21:49:42.330108","exception":false,"start_time":"2024-06-23T21:49:42.320438","status":"completed"},"tags":[]}},{"cell_type":"code","source":"LOAD_RAW_DATA = True","metadata":{"papermill":{"duration":0.018723,"end_time":"2024-06-23T21:49:42.358860","exception":false,"start_time":"2024-06-23T21:49:42.340137","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-06T05:20:06.065178Z","iopub.execute_input":"2024-07-06T05:20:06.066078Z","iopub.status.idle":"2024-07-06T05:20:06.073513Z","shell.execute_reply.started":"2024-07-06T05:20:06.066034Z","shell.execute_reply":"2024-07-06T05:20:06.072444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nif LOAD_RAW_DATA:\n    # Use DuckDB SQL query to load data from original train dataset\n    con = duckdb.connect()\n    sample_df = con.query(f\"\"\"(SELECT *\n                                FROM parquet_scan('{train_file}')\n                                WHERE binds = 0 AND protein_name = 'HSA'\n                                ORDER BY random()\n                                LIMIT 500000)\n                                UNION ALL\n                                (SELECT *\n                                FROM parquet_scan('{train_file}')\n                                WHERE binds = 1 AND protein_name = 'HSA'\n                                ORDER BY random()\n                                LIMIT 500000)\n\n                                UNION ALL\n                                (SELECT *\n                                FROM parquet_scan('{train_file}')\n                                WHERE binds = 0 AND protein_name = 'BRD4'\n                                ORDER BY random()\n                                LIMIT 500000)\n                                UNION ALL\n                                (SELECT *\n                                FROM parquet_scan('{train_file}')\n                                WHERE binds = 1 AND protein_name = 'BRD4'\n                                ORDER BY random()\n                                LIMIT 500000)\n\n                                UNION ALL\n                                (SELECT *\n                                FROM parquet_scan('{train_file}')\n                                WHERE binds = 0 AND protein_name = 'sEH'\n                                ORDER BY random()\n                                LIMIT 500000)\n                                UNION ALL\n                                (SELECT *\n                                FROM parquet_scan('{train_file}')\n                                WHERE binds = 1 AND protein_name = 'sEH'\n                                ORDER BY random()\n                                LIMIT 500000)\"\"\").df()\n    con.close()\n    print(sample_df[['protein_name', 'binds', 'molecule_smiles']].groupby(by=['protein_name', 'binds']).count())\nelse:\n    # Load previously saved (pre-processed training datasets)\n    train_df = {}\n    for p in proteins:\n#         file_name = p + '_900K_df.parquet'7\n        print(file_name)\n        train_df[p] = pd.read_parquet(pre_processed_train_data_path + file_name)\n        print(train_df[p][['protein_name', 'binds', 'molecule_smiles']].groupby(by=['protein_name', 'binds']).count())\n        print()","metadata":{"papermill":{"duration":4.591552,"end_time":"2024-06-23T21:49:46.960533","exception":false,"start_time":"2024-06-23T21:49:42.368981","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-06T05:20:06.075302Z","iopub.execute_input":"2024-07-06T05:20:06.076064Z","iopub.status.idle":"2024-07-06T05:22:23.966229Z","shell.execute_reply.started":"2024-07-06T05:20:06.076023Z","shell.execute_reply":"2024-07-06T05:22:23.964926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nif LOAD_RAW_DATA:\n    # Separate train data by protein\n    train_df = {}\n    for p in proteins:\n        train_df[p] = sample_df[sample_df['protein_name']==p].reset_index(drop=True)\n        # Create 'ecfp' description from the 'molecule_smiles'\n        train_df[p]['ecfp'] = train_df[p]['molecule_smiles'].apply(MolFromSmiles).apply(generate_ecfp)\n    #     # Save the datasets in parquet format\n    #     for p in proteins:\n    #     file_name = p + '_900K_df.parquet'\n    #     print('Saving dataset', file_name)\n    #     train_df[p].to_parquet(file_name)\nelse:\n    # Sample down pre-processed dataset\n    for p in proteins:\n#         train_df[p] = train_df[p].sample(400000).reset_index(drop=True)\n        print(train_df[p][['protein_name', 'binds', 'molecule_smiles']].groupby(by=['protein_name', 'binds']).count())","metadata":{"papermill":{"duration":1.384876,"end_time":"2024-06-23T21:49:48.355911","exception":false,"start_time":"2024-06-23T21:49:46.971035","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-06T05:22:23.967881Z","iopub.execute_input":"2024-07-06T05:22:23.968257Z","iopub.status.idle":"2024-07-06T05:45:22.244621Z","shell.execute_reply.started":"2024-07-06T05:22:23.968225Z","shell.execute_reply":"2024-07-06T05:45:22.243203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features = [('EC_' + str(i)) for i in range(chemical_nbits)]\ntarget = 'binds'","metadata":{"execution":{"iopub.status.busy":"2024-07-06T05:45:22.249087Z","iopub.execute_input":"2024-07-06T05:45:22.249492Z","iopub.status.idle":"2024-07-06T05:45:22.255835Z","shell.execute_reply.started":"2024-07-06T05:45:22.249460Z","shell.execute_reply":"2024-07-06T05:45:22.254461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define dictionaries (keys are the proteins)\nX = {} \ny = {}\n# spw = {}\nmodel = {}\noptimal_params = {}","metadata":{"papermill":{"duration":0.019091,"end_time":"2024-06-23T21:49:48.620961","exception":false,"start_time":"2024-06-23T21:49:48.601870","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-06T05:45:22.257334Z","iopub.execute_input":"2024-07-06T05:45:22.257731Z","iopub.status.idle":"2024-07-06T05:45:22.269993Z","shell.execute_reply.started":"2024-07-06T05:45:22.257699Z","shell.execute_reply":"2024-07-06T05:45:22.268770Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n# Separate X and y for each protein model\nfor p in proteins:\n    ecfp_array = np.array(train_df[p]['ecfp'].tolist())\n    X[p] = pd.DataFrame(data=csr_matrix(ecfp_array).todense(), columns=features)\n    # Reduce int64 to int8\n    X[p] = X[p].astype(np.int8)\n    y[p] = train_df[p][target]","metadata":{"papermill":{"duration":0.03637,"end_time":"2024-06-23T21:49:48.668027","exception":false,"start_time":"2024-06-23T21:49:48.631657","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-07-06T05:45:22.271300Z","iopub.execute_input":"2024-07-06T05:45:22.271733Z","iopub.status.idle":"2024-07-06T06:17:05.100098Z","shell.execute_reply.started":"2024-07-06T05:45:22.271698Z","shell.execute_reply":"2024-07-06T06:17:05.097182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X['sEH']","metadata":{"execution":{"iopub.status.busy":"2024-07-06T06:17:05.105157Z","iopub.execute_input":"2024-07-06T06:17:05.105697Z","iopub.status.idle":"2024-07-06T06:17:06.212597Z","shell.execute_reply.started":"2024-07-06T06:17:05.105653Z","shell.execute_reply":"2024-07-06T06:17:06.207553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"There are a few features with a costant value (1) which add no value, but there are not too many so I won´t worry about them.","metadata":{}},{"cell_type":"code","source":"[(col, X['sEH'][col].unique()) for col in X['sEH'].columns if X['sEH'][col].nunique()==1]","metadata":{"execution":{"iopub.status.busy":"2024-07-06T06:34:52.313596Z","iopub.execute_input":"2024-07-06T06:34:52.315259Z","iopub.status.idle":"2024-07-06T06:35:05.682106Z","shell.execute_reply.started":"2024-07-06T06:34:52.315202Z","shell.execute_reply":"2024-07-06T06:35:05.680810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y['sEH']","metadata":{"execution":{"iopub.status.busy":"2024-07-06T06:37:24.288770Z","iopub.execute_input":"2024-07-06T06:37:24.289326Z","iopub.status.idle":"2024-07-06T06:37:24.302156Z","shell.execute_reply.started":"2024-07-06T06:37:24.289293Z","shell.execute_reply":"2024-07-06T06:37:24.300469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Free up memory\ndel sample_df\nfor p in proteins:\n    del train_df[p]\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-07-06T06:37:27.113987Z","iopub.execute_input":"2024-07-06T06:37:27.114431Z","iopub.status.idle":"2024-07-06T06:37:29.110138Z","shell.execute_reply.started":"2024-07-06T06:37:27.114393Z","shell.execute_reply":"2024-07-06T06:37:29.108829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Models hyper-parameter tuning\n---\n\nEither Optuna or Grid-search cross-validation (GSCV).","metadata":{"papermill":{"duration":0.010484,"end_time":"2024-06-23T21:49:48.560843","exception":false,"start_time":"2024-06-23T21:49:48.550359","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"The following block implements the Optuna hyper-parameter tuning. Much faster than GSCV but similar results.","metadata":{"papermill":{"duration":0.010238,"end_time":"2024-06-23T21:49:48.689223","exception":false,"start_time":"2024-06-23T21:49:48.678985","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%%time\n\nwarnings.simplefilter(action='ignore')\n\nstudy = {}\noptimal_params = {}\n\noptuna.logging.set_verbosity(optuna.logging.ERROR)\n\nfor p in proteins:\n    \n    # Use XGB native DMatrix format\n    data = xgb.DMatrix(X[p], y[p])\n    \n    # Define Optuna objective function\n    def objective(trial: optuna.Trial):\n        # XGBoost parameters\n        params = {\n            \"objective\" : \"binary:logistic\",\n            \"eval_metric\" : \"map\",\n            \"n_estimators\" : trial.suggest_int(\"n_estimators\", 300, 3000),\n            \"learning_rate\" : trial.suggest_float(\"learning_rate\", 0.1, 0.5, log=False),\n            \"max_depth\" : trial.suggest_int(\"max_depth\", 10, 20),\n            \"min_child_weight\": trial.suggest_int('min_child_weight', 3, 30, log=False),\n            \"subsample\": trial.suggest_float(\"subsample\", 0.1, 1, log=False),\n            \"colsample_bytree\": trial.suggest_float(\"colsample_bytree\", 0.1, 1, log=False),\n            \"colsample_bynode\": trial.suggest_float(\"colsample_bynode\", 0.1, 1, log=False),\n            \"colsample_bylevel\": trial.suggest_float(\"colsample_bylevel\", 0.1, 1, log=False),\n        }\n\n        # Define pruning callback (speeds up the process and helps avoid overfitting)\n#         pruning_callback = XGBoostPruningCallback(trial, 'validation_0-map')   \n\n        xgb_cv_scores = xgb.cv(\n            params=params,\n            dtrain=data,\n            nfold=4,\n            stratified=True,\n            early_stopping_rounds=10,\n            verbose_eval=False,\n#             callbacks=[pruning_callback]\n        ) \n\n        avg_score = np.mean(xgb_cv_scores[\"test-map-mean\"].values)\n        return avg_score\n\n    # Create and run Optuna object\n    study[p] = optuna.create_study(direction=\"maximize\")\n    study[p].optimize(objective, n_trials=100, timeout=600)\n    optimal_params[p] = study[p].best_trial.params\n    \n    # Print results\n    print('Model', p)\n    print(\"Number of finished trials: {}\".format(len(study[p].trials)))\n    print('Best hyperparameters:\\n', optimal_params[p])\n    print('-'*140)\n    \n    gc.collect()","metadata":{"papermill":{"duration":1813.036539,"end_time":"2024-06-23T22:20:01.736382","exception":false,"start_time":"2024-06-23T21:49:48.699843","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Below are some Optuna visualizations on the training processes...","metadata":{"papermill":{"duration":0.010462,"end_time":"2024-06-23T22:20:01.757566","exception":false,"start_time":"2024-06-23T22:20:01.747104","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Visualize optimization history\noptuna.visualization.plot_optimization_history(study['sEH'])","metadata":{"papermill":{"duration":0.632239,"end_time":"2024-06-23T22:20:02.400423","exception":false,"start_time":"2024-06-23T22:20:01.768184","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize optimization history\noptuna.visualization.plot_optimization_history(study['BRD4'])","metadata":{"papermill":{"duration":0.041828,"end_time":"2024-06-23T22:20:02.517660","exception":false,"start_time":"2024-06-23T22:20:02.475832","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize optimization history\noptuna.visualization.plot_optimization_history(study['HSA'])","metadata":{"papermill":{"duration":0.042985,"end_time":"2024-06-23T22:20:02.572831","exception":false,"start_time":"2024-06-23T22:20:02.529846","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize parameter importances\noptuna.visualization.plot_param_importances(study['sEH'])","metadata":{"papermill":{"duration":1.490808,"end_time":"2024-06-23T22:20:04.075834","exception":false,"start_time":"2024-06-23T22:20:02.585026","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize parameter importances\noptuna.visualization.plot_param_importances(study['BRD4'])","metadata":{"papermill":{"duration":1.466462,"end_time":"2024-06-23T22:20:05.554947","exception":false,"start_time":"2024-06-23T22:20:04.088485","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize parameter importances\noptuna.visualization.plot_param_importances(study['HSA'])","metadata":{"papermill":{"duration":1.321724,"end_time":"2024-06-23T22:20:06.889323","exception":false,"start_time":"2024-06-23T22:20:05.567599","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize the hyperparameters and scores\n# optuna.visualization.plot_parallel_coordinate(study[p])","metadata":{"papermill":{"duration":0.021198,"end_time":"2024-06-23T22:20:06.924586","exception":false,"start_time":"2024-06-23T22:20:06.903388","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The following block implements the GSCV hyper-parameter tuning.","metadata":{"papermill":{"duration":0.012533,"end_time":"2024-06-23T22:20:06.949772","exception":false,"start_time":"2024-06-23T22:20:06.937239","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# # Parameters grid for the cross-validation exercise\n# hyperparam_grid = {\n#     'n_estimators': [1500, 2000, 2500],\n#     'learning_rate': [0.01, 0.03, 0.05, 0.07, 0.1],\n#     'max_depth': [3, 5, 7, 9],\n# #     'subsample': [0, 0.2, 0.4, 0.6, 0.8, 1],\n# #     'colsample_bynode': [0, 0.2, 0.4, 0.6, 0.8, 1],    \n# #     'colsample_bylevel': [0, 0.2, 0.4, 0.6, 0.8, 1],   \n# #     'colsample_bytree': [0, 0.2, 0.4, 0.6, 0.8, 1],\n# #     'min_child_weight': [3, 5, 9],\n# }","metadata":{"papermill":{"duration":0.020722,"end_time":"2024-06-23T22:20:06.983208","exception":false,"start_time":"2024-06-23T22:20:06.962486","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%time\n# # Create models and run grid search cross-validation\n# for p in proteins:\n#     # Instantiate XGBoost model\n#     model[p] = XGBClassifier( #scale_pos_weight=spw[p],\n#                              random_state=13,\n#                              tree_method = \"hist\",\n#                              device = \"cpu\")\n# #     dtrain = xgb.DMatrix(X[p], label=y[p])\n#     print('Model', p)\n#     print('Running grid search cross validation....')\n#     # Set up the gscv object with 4-fold\n#     gs_cv = GridSearchCV(estimator=model[p],\n#                          param_grid=hyperparam_grid,\n#                          scoring='average_precision',\n#                          cv=4,\n#                          return_train_score=False,\n#                          n_jobs=-1,\n#                          verbose=1)\n#     gs_cv.fit(X[p],y[p])\n#     optimal_params[p] = gs_cv.best_params_\n#     print('Optimal params for model', p)\n#     print(optimal_params[p])\n#     print('-'*80)","metadata":{"papermill":{"duration":0.021205,"end_time":"2024-06-23T22:20:07.017221","exception":false,"start_time":"2024-06-23T22:20:06.996016","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"papermill":{"duration":0.199955,"end_time":"2024-06-23T22:20:07.233212","exception":false,"start_time":"2024-06-23T22:20:07.033257","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Fit optimal models with full sample (optional model save)\n---\n","metadata":{"papermill":{"duration":0.012357,"end_time":"2024-06-23T22:20:07.258469","exception":false,"start_time":"2024-06-23T22:20:07.246112","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%%time\n# Fit and evaluate optimal models\nfor p in proteins:\n    model[p] = XGBClassifier(**optimal_params[p],\n                             random_state=13,\n                             tree_method = \"hist\",\n                             device = \"cpu\")\n    model[p].fit(X[p], y[p])\n    cv_scores = cross_val_score(model[p], X[p], y[p], cv=4, scoring='average_precision')\n    print('Cross validation score for model', p)\n    print('Average precision:', str(round(cv_scores.mean(), 3)))\n    print('-'*40)","metadata":{"papermill":{"duration":1001.001146,"end_time":"2024-06-23T22:36:48.272455","exception":false,"start_time":"2024-06-23T22:20:07.271309","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save the tuned models\nfor p in proteins:\n    model_name = 'XGB_' + p + '_1M_ECFP_128_int8_FHYP.model'\n    print('Saving model ', model_name)\n    model[p].save_model(model_name)","metadata":{"papermill":{"duration":0.247945,"end_time":"2024-06-23T22:36:48.533593","exception":false,"start_time":"2024-06-23T22:36:48.285648","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Features analysis and selection (only run with small samples, i.e. no more than 10K per protein)\n---\n\nHere I compute the correlation between features, and the models SHAP values.","metadata":{"papermill":{"duration":0.012879,"end_time":"2024-06-23T22:36:48.560616","exception":false,"start_time":"2024-06-23T22:36:48.547737","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%%time\n# Calculate and show Shap Values (this will take quite long with samples over a few thousand)\n# shap.initjs()\n# for p in proteins:\n#     shap_sample = X[p].sample(2000)\n#     explainer = shap.TreeExplainer(model[p])\n#     shap_values = explainer.shap_values(shap_sample)\n#     plt.suptitle('Model ' + p, fontsize=15) \n#     shap.summary_plot(shap_values, shap_sample, plot_size=[8,8], plot_type='dot', max_display=25, show=True)","metadata":{"papermill":{"duration":0.02323,"end_time":"2024-06-23T22:36:48.597738","exception":false,"start_time":"2024-06-23T22:36:48.574508","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Analysis of correlation between features and target\n# fig, ax = plt.subplots(3,1, figsize=(16,18), sharex=True)\n# i = 0 \n# for p in proteins:\n#     ax[i].set_title(p, fontsize=14)\n#     ax[i].tick_params(axis='both', labelsize=10)\n#     sns.heatmap(X[p].corr(method='pearson'), annot=True, annot_kws={\"size\":9}, cmap='YlGn', ax=ax[i], square=True)\n#     i = i + 1\n# plt.show()","metadata":{"papermill":{"duration":0.022708,"end_time":"2024-06-23T22:36:48.633782","exception":false,"start_time":"2024-06-23T22:36:48.611074","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]}]}