{"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":"# Import dependencies","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom typing import Counter","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-04T06:43:18.499552Z","iopub.execute_input":"2023-01-04T06:43:18.500291Z","iopub.status.idle":"2023-01-04T06:43:19.599154Z","shell.execute_reply.started":"2023-01-04T06:43:18.500199Z","shell.execute_reply":"2023-01-04T06:43:19.598042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# StratifiedGroupKFold\nhttps://scikit-learn.org/stable/modules/generated/sklearn.model_selection.StratifiedGroupKFold.html\n\nThis cross-validation object is a variation of StratifiedKFold attempts to return stratified folds with non-overlapping groups. The folds are made by preserving the percentage of samples for each class.","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2023-01-04T06:43:19.601385Z","iopub.execute_input":"2023-01-04T06:43:19.601836Z","iopub.status.idle":"2023-01-04T06:43:19.611193Z","shell.execute_reply.started":"2023-01-04T06:43:19.601792Z","shell.execute_reply":"2023-01-04T06:43:19.609758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_folds = 4\n\ndf = pd.read_csv(\"/kaggle/input/rsna-breast-cancer-detection/train.csv\")\n\ndf_out = apply_StratifiedGroupKFold(\n    X=df,\n    y=df[\"cancer\"].values,\n    groups=df[\"patient_id\"].values,\n    n_splits=num_folds)\n\ndf_out.to_csv(f\"train_{num_folds}folds.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-01-04T06:43:19.613027Z","iopub.execute_input":"2023-01-04T06:43:19.613484Z","iopub.status.idle":"2023-01-04T06:43:24.009937Z","shell.execute_reply.started":"2023-01-04T06:43:19.61344Z","shell.execute_reply":"2023-01-04T06:43:24.008831Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n\ncols = [\"cancer\", \"machine_id\", \"site_id\", \"density\", \"difficult_negative_case\"]\nfig, axes = plt.subplots(ncols=1, nrows=len(cols), figsize=(12, 40))\n\nfor col, ax in zip(cols, axes):\n    sns.countplot(x=col, data=df_out, hue=\"fold\", ax=ax)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-04T06:47:05.326972Z","iopub.execute_input":"2023-01-04T06:47:05.32739Z","iopub.status.idle":"2023-01-04T06:47:06.377174Z","shell.execute_reply.started":"2023-01-04T06:47:05.327356Z","shell.execute_reply":"2023-01-04T06:47:06.376166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}