{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":170595844,"sourceType":"kernelVersion"}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from pathlib import Path\nimport pickle\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nfrom sklearn.model_selection import KFold","metadata":{"ExecuteTime":{"end_time":"2024-05-28T05:52:16.647305Z","start_time":"2024-05-28T05:52:14.307630Z"},"execution":{"iopub.status.busy":"2024-05-28T06:41:09.228419Z","iopub.execute_input":"2024-05-28T06:41:09.228846Z","iopub.status.idle":"2024-05-28T06:41:11.892913Z","shell.execute_reply.started":"2024-05-28T06:41:09.228812Z","shell.execute_reply":"2024-05-28T06:41:11.892038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_root = Path(\"/kaggle/input/belka-shrinking-the-dataset\")","metadata":{"ExecuteTime":{"end_time":"2024-05-28T05:16:30.390177Z","start_time":"2024-05-28T05:16:30.388257Z"},"execution":{"iopub.status.busy":"2024-05-28T06:41:19.986019Z","iopub.execute_input":"2024-05-28T06:41:19.986664Z","iopub.status.idle":"2024-05-28T06:41:19.992280Z","shell.execute_reply.started":"2024-05-28T06:41:19.986628Z","shell.execute_reply":"2024-05-28T06:41:19.991187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_parquet(data_root / \"train.parquet\", columns=[\n    \"buildingblock1_smiles\",\n    \"buildingblock2_smiles\",\n    \"buildingblock3_smiles\"\n])","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2024-05-28T05:16:30.659214Z","start_time":"2024-05-28T05:16:30.391173Z"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-28T06:41:20.344898Z","iopub.execute_input":"2024-05-28T06:41:20.345870Z","iopub.status.idle":"2024-05-28T06:41:22.637786Z","shell.execute_reply.started":"2024-05-28T06:41:20.345833Z","shell.execute_reply":"2024-05-28T06:41:22.636723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dicts = [\n    pickle.load(\n        open(data_root.joinpath(f\"train_dicts\", f\"BBs_dict_reverse_{i + 1}.p\"), \"br\"))\n    for i in range(3)]","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2024-05-28T05:16:38.308387Z","start_time":"2024-05-28T05:16:38.277195Z"},"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-05-28T06:41:33.462548Z","iopub.execute_input":"2024-05-28T06:41:33.462980Z","iopub.status.idle":"2024-05-28T06:41:33.498549Z","shell.execute_reply.started":"2024-05-28T06:41:33.462950Z","shell.execute_reply":"2024-05-28T06:41:33.497682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bb1 = np.array(list(dicts[0].values()))\nbb2 = np.array(list(dicts[1].values()))\nbb3 = np.array(list(dicts[2].values()))\nbb3 = np.array(sorted(set(bb3) - set(bb2)))","metadata":{"ExecuteTime":{"end_time":"2024-05-28T05:53:29.955349Z","start_time":"2024-05-28T05:53:29.952304Z"},"execution":{"iopub.status.busy":"2024-05-28T06:41:34.337529Z","iopub.execute_input":"2024-05-28T06:41:34.338410Z","iopub.status.idle":"2024-05-28T06:41:34.347043Z","shell.execute_reply.started":"2024-05-28T06:41:34.338367Z","shell.execute_reply":"2024-05-28T06:41:34.345808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_blocks = np.concatenate([bb1, bb2, bb3])\nlen(all_blocks)","metadata":{"ExecuteTime":{"end_time":"2024-05-28T05:53:55.429010Z","start_time":"2024-05-28T05:53:55.426528Z"},"execution":{"iopub.status.busy":"2024-05-28T06:41:36.634541Z","iopub.execute_input":"2024-05-28T06:41:36.634945Z","iopub.status.idle":"2024-05-28T06:41:36.644395Z","shell.execute_reply.started":"2024-05-28T06:41:36.634916Z","shell.execute_reply":"2024-05-28T06:41:36.642996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bb_str_to_global_id = {bb: i for i, bb in enumerate(all_blocks)}","metadata":{"ExecuteTime":{"end_time":"2024-05-28T05:53:57.945847Z","start_time":"2024-05-28T05:53:57.942578Z"},"execution":{"iopub.status.busy":"2024-05-28T06:41:37.827257Z","iopub.execute_input":"2024-05-28T06:41:37.827709Z","iopub.status.idle":"2024-05-28T06:41:37.836129Z","shell.execute_reply.started":"2024-05-28T06:41:37.827677Z","shell.execute_reply":"2024-05-28T06:41:37.834722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bb_indices = np.empty((len(train), 3), dtype=np.int16)\n\nfor i, row in tqdm(train.iterrows(), total=len(train)):\n    for bi in range(3):\n        bb_indices[i, bi] = bb_str_to_global_id[dicts[bi][row[f\"buildingblock{bi + 1}_smiles\"]]]","metadata":{"ExecuteTime":{"end_time":"2024-05-28T05:51:57.591040Z","start_time":"2024-05-28T05:31:08.536660Z"},"execution":{"iopub.status.busy":"2024-05-28T06:41:41.158100Z","iopub.execute_input":"2024-05-28T06:41:41.158560Z","iopub.status.idle":"2024-05-28T06:42:05.251732Z","shell.execute_reply.started":"2024-05-28T06:41:41.158526Z","shell.execute_reply":"2024-05-28T06:42:05.249842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kf1 = KFold(n_splits=5, shuffle=True, random_state=1)\nkf2 = KFold(n_splits=5, shuffle=True, random_state=2)\nkf3 = KFold(n_splits=5, shuffle=True, random_state=3)\n\ntrain_keeps = []\nval_keeps = []\n\nfor fold_i, ((train1, val1), (train2, val2), (train3, val3)) in enumerate(zip(kf1.split(bb1), kf2.split(bb2), kf3.split(bb3))):\n    train_keep = True\n    val_keep = False\n    val_ids = set(bb1[val1]) | set(bb2[val2]) | set(bb3[val3])\n    val_ids = set([bb_str_to_global_id[bb] for bb in val_ids])\n    \n    for val_id in tqdm(val_ids):\n        val_keep |= (bb_indices == val_id)\n        train_keep &= (bb_indices != val_id)\n        \n    train_keep = train_keep.all(axis=1)\n    val_keep = val_keep.all(axis=1)\n    \n    train_keeps.append(train_keep)\n    val_keeps.append(val_keep)","metadata":{"ExecuteTime":{"end_time":"2024-05-28T06:00:13.677684Z","start_time":"2024-05-28T05:58:30.936590Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_keeps = np.array(train_keeps)\nval_keeps = np.array(val_keeps)\n\n# save train_keeps and val_keeps to pickle\nwith open(\"split.pkl\", \"wb\") as f:\n    pickle.dump((train_keeps, val_keeps), f)","metadata":{"ExecuteTime":{"end_time":"2024-05-28T06:23:11.367020Z","start_time":"2024-05-28T06:23:09.965316Z"}},"execution_count":null,"outputs":[]}]}