{"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":"code","source":"import os\nimport pandas as pd\n\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom typing import Counter","metadata":{"execution":{"iopub.status.busy":"2022-07-10T20:55:43.88644Z","iopub.execute_input":"2022-07-10T20:55:43.887309Z","iopub.status.idle":"2022-07-10T20:55:45.229511Z","shell.execute_reply.started":"2022-07-10T20:55:43.887175Z","shell.execute_reply":"2022-07-10T20:55:45.228153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def apply_StratifiedGroupKFold(X, y, groups, n_splits, random_state=42):\n\n    df_out = X.copy(deep=True)\n\n    # split\n    cv = StratifiedGroupKFold(n_splits=n_splits, random_state=random_state, shuffle=True)\n    for fold_index, (train_index, val_index) in enumerate(cv.split(X, y, groups)):\n\n        df_out.loc[val_index, \"Fold\"] = fold_index\n\n        # check\n        train_groups, val_groups = groups[train_index], groups[val_index]\n        assert len(set(train_groups) & set(val_groups)) == 0\n\n    df_out = df_out.astype({\"Fold\": 'int64'})\n\n    return df_out","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-10T20:55:45.231936Z","iopub.execute_input":"2022-07-10T20:55:45.233053Z","iopub.status.idle":"2022-07-10T20:55:45.246539Z","shell.execute_reply.started":"2022-07-10T20:55:45.233003Z","shell.execute_reply":"2022-07-10T20:55:45.245514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_dir = \"../input/mayo-clinic-strip-ai\"\noutput_dir = \"./\"\nnum_folds = 5\n\ndf = pd.read_csv(os.path.join(input_dir, \"train.csv\"))\n\nprint(f\"num_folds={num_folds}\")\ndf_out = apply_StratifiedGroupKFold(\n    X=df,\n    y=df[\"label\"].values,\n    groups=df[\"center_id\"].values,\n    n_splits=num_folds)\n\n# check\nfor fold_index in range(num_folds):\n    records = df_out[(df_out[\"Fold\"] == fold_index)]\n    label_counts = Counter(records[\"label\"].values)\n    center_counts = Counter(records[\"center_id\"].values)\n    print(f\"fold{fold_index}: {label_counts} {center_counts}\")\n\ndf_out.to_csv(os.path.join(output_dir, f\"mayo_clinic_strip_ai_{num_folds}folds.csv\"), index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-10T20:55:45.25033Z","iopub.execute_input":"2022-07-10T20:55:45.251084Z","iopub.status.idle":"2022-07-10T20:55:45.323651Z","shell.execute_reply.started":"2022-07-10T20:55:45.251007Z","shell.execute_reply":"2022-07-10T20:55:45.322792Z"},"trusted":true},"execution_count":null,"outputs":[]}]}