{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"},{"sourceId":187052395,"sourceType":"kernelVersion"}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Package install, library imports and settings","metadata":{}},{"cell_type":"code","source":"!pip install duckdb","metadata":{"execution":{"iopub.status.busy":"2024-07-06T10:19:41.472798Z","iopub.execute_input":"2024-07-06T10:19:41.473290Z","iopub.status.idle":"2024-07-06T10:20:00.542830Z","shell.execute_reply.started":"2024-07-06T10:19:41.473252Z","shell.execute_reply":"2024-07-06T10:20:00.541071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install rdkit","metadata":{"execution":{"iopub.status.busy":"2024-07-06T10:20:00.545511Z","iopub.execute_input":"2024-07-06T10:20:00.545893Z","iopub.status.idle":"2024-07-06T10:20:19.002076Z","shell.execute_reply.started":"2024-07-06T10:20:00.545860Z","shell.execute_reply":"2024-07-06T10:20:19.000194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport numpy as np \nimport pandas as pd \nimport duckdb\nfrom rdkit.Chem import AllChem, MolFromSmiles, rdFingerprintGenerator\nfrom sklearn.model_selection import cross_val_score, GridSearchCV\nfrom sklearn.metrics import average_precision_score\nfrom scipy.sparse import csr_matrix\nimport xgboost as xgb\nfrom xgboost import XGBClassifier\n\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-06T10:20:19.005392Z","iopub.execute_input":"2024-07-06T10:20:19.005879Z","iopub.status.idle":"2024-07-06T10:20:20.808833Z","shell.execute_reply.started":"2024-07-06T10:20:19.005837Z","shell.execute_reply":"2024-07-06T10:20:20.807634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pre_trained_model_path = '/kaggle/input/belka-model-training-and-evaluation-ecfp-sparse/'\ntest_data_file = '/kaggle/input/leash-BELKA/test.parquet'\n\nproteins = ['sEH', 'BRD4', 'HSA']\n\nchemical_radius = 2\nchemical_nbits = 128\n\nfeatures = [('EC_' + str(i)) for i in range(chemical_nbits)]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load pre-trained models","metadata":{}},{"cell_type":"code","source":"%%time\n# Load previously saved models\nmodel = {}\nfor p in proteins:\n    model_name =  pre_trained_model_path + 'XGB_' + p + '_1M_ECFP_128_int8_FHYP.model'\n    model[p] = XGBClassifier()\n    model[p].load_model(model_name)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load and process test data","metadata":{}},{"cell_type":"code","source":"%%time\n# Read test data\n# Use DuckDB SQL query to load data from original train dataset\ncon = duckdb.connect()\ntest_data = con.query(f\"\"\"(SELECT * FROM parquet_scan('{test_data_file}'))\"\"\").df()\ncon.close()\n\ntest_data.info()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Sample down test data (DEVL ONLY)\n# test_data = test_data.sample(10000)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n# Generate ECFP function\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)\n    \n# Split test data by protein into three dataframes\ntest_df = {}\nfor p in proteins:\n    test_df[p] = test_data[test_data['protein_name']==p].reset_index(drop=True)\n    # Generate ECFPs\n    test_df[p]['ecfp'] = test_df[p]['molecule_smiles'].apply(MolFromSmiles).apply(generate_ecfp)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ecfp_array = np.array(test_df['sEH']['ecfp'].tolist())\npd.DataFrame(data=csr_matrix(ecfp_array).todense(), columns=features)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Make predictions and submit","metadata":{}},{"cell_type":"code","source":"%%time\n# Make predictions on test data\npredictions = {}\n\nfor p in proteins:\n    ecfp_array = np.array(test_df[p]['ecfp'].tolist())\n    X_test = pd.DataFrame(data=csr_matrix(ecfp_array).todense(), columns=features)\n    predictions[p] = model[p].predict_proba(X_test)   \n    predictions[p] = pd.DataFrame(predictions[p])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions[p][1].hist(bins=2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create submission DataFrame\nsubmission_df = pd.DataFrame()\nfor p in proteins:\n    sub_df = {\n        'id': test_df[p]['id'],\n        'binds': predictions[p][1]\n    }\n    sub_df = pd.DataFrame(sub_df)\n    submission_df = pd.concat([submission_df, sub_df], ignore_index=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save submission file to disk\nsubmission_df.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}