{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"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","trusted":true,"execution":{"iopub.status.busy":"2025-03-12T09:30:40.531941Z","iopub.execute_input":"2025-03-12T09:30:40.532202Z","iopub.status.idle":"2025-03-12T09:30:43.484714Z","shell.execute_reply.started":"2025-03-12T09:30:40.532176Z","shell.execute_reply":"2025-03-12T09:30:43.483269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_StratifiedGroupKFold(X, y, groups, n_splits, random_state=42):\n    \"\"\"Apply StratifiedGroupKFold cross-validation to a dataframe\"\"\"\n    df_out = df.copy()\n\n    # Apply StratifiedGroupKFold splitting\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        df_out.loc[val_index, \"fold\"] = fold_index\n\n        # check\n        train_tomo_ids, val_tomo_ids = groups[train_index], groups[val_index]\n        assert len(set(train_tomo_ids) & set(val_tomo_ids)) == 0\n\n    df_out = df_out.astype({\"fold\": 'int64'})\n    return df_out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-12T09:34:39.105885Z","iopub.execute_input":"2025-03-12T09:34:39.106522Z","iopub.status.idle":"2025-03-12T09:34:39.113730Z","shell.execute_reply.started":"2025-03-12T09:34:39.106488Z","shell.execute_reply":"2025-03-12T09:34:39.112412Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_folds = 5\n\ndf = pd.read_csv(\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv\")\n\ndf_out = apply_StratifiedGroupKFold(\n    X=df,\n    y=df[\"Number of motors\"].values,\n    groups=df[\"tomo_id\"].values,\n    n_splits=num_folds, \n    random_state=42,\n)\n\ndf_out.to_csv(f\"train_{num_folds}folds.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-12T09:35:49.153577Z","iopub.execute_input":"2025-03-12T09:35:49.154015Z","iopub.status.idle":"2025-03-12T09:35:49.473900Z","shell.execute_reply.started":"2025-03-12T09:35:49.153983Z","shell.execute_reply":"2025-03-12T09:35:49.472549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n\ncols = [\n    #\"Motor axis 0\",\n    #\"Motor axis 1\",\n    #\"Motor axis 2\",\n    \"Array shape (axis 0)\",\n    \"Array shape (axis 1)\",\n    \"Array shape (axis 2)\",\n    \"Voxel spacing\",\n    \"Number of motors\",\n]\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-12T09:35:49.475319Z","iopub.execute_input":"2025-03-12T09:35:49.475705Z","iopub.status.idle":"2025-03-12T09:35:51.079738Z","shell.execute_reply.started":"2025-03-12T09:35:49.475624Z","shell.execute_reply":"2025-03-12T09:35:51.078256Z"}},"outputs":[],"execution_count":null}]}