{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":4696088,"sourceType":"datasetVersion","datasetId":2687741},{"sourceId":7469530,"sourceType":"datasetVersion","datasetId":4348179}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":8485.182933,"end_time":"2024-01-23T00:03:09.45319","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-01-22T21:41:44.270257","version":"2.4.0"},"widgets":{"application/vnd.jupyter.widget-state+json":{"state":{"21b011cf0c8f4dff8af3e1d24311faf0":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"23f7c88e059043ffaeb873280650a9dc":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_6e8c9a24c3a24336a54586867d59976d","placeholder":"​","style":"IPY_MODEL_e323e75c270f4d58bbcbea9b18f40000","value":" 193M/193M [00:01&lt;00:00, 195MB/s]"}},"3d6ea4e54a574cc8becbc86bd9d1fe9d":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"42487334e2504ea787417c08ae1bef03":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"5e407baec4cb4b5490de3473d8d08b5f":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"6e8c9a24c3a24336a54586867d59976d":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"8327ce97cc7440fb89b7f5dc58bb8930":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_e9d2511d74524343b3e1f175d5e210ff","IPY_MODEL_d8d4966b5f9741f8a14cc7e3e06ce210","IPY_MODEL_23f7c88e059043ffaeb873280650a9dc"],"layout":"IPY_MODEL_21b011cf0c8f4dff8af3e1d24311faf0"}},"8ccb182709b544b08fadfdad0366c6b7":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"d8d4966b5f9741f8a14cc7e3e06ce210":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"success","description":"","description_tooltip":null,"layout":"IPY_MODEL_8ccb182709b544b08fadfdad0366c6b7","max":193322802,"min":0,"orientation":"horizontal","style":"IPY_MODEL_42487334e2504ea787417c08ae1bef03","value":193322802}},"e323e75c270f4d58bbcbea9b18f40000":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"e9d2511d74524343b3e1f175d5e210ff":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_5e407baec4cb4b5490de3473d8d08b5f","placeholder":"​","style":"IPY_MODEL_3d6ea4e54a574cc8becbc86bd9d1fe9d","value":"model.safetensors: 100%"}}},"version_major":2,"version_minor":0}}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"33a263cb-2645-4ddc-9018-3bc4754caa95","cell_type":"markdown","source":"# I. Train","metadata":{"papermill":{"duration":0.013181,"end_time":"2024-01-22T21:41:47.957087","exception":false,"start_time":"2024-01-22T21:41:47.943906","status":"completed"},"tags":[]}},{"id":"9181f663-6994-4d41-9f6a-8c93fc75e834","cell_type":"markdown","source":"### Libraries","metadata":{"papermill":{"duration":0.012488,"end_time":"2024-01-22T21:41:47.982262","exception":false,"start_time":"2024-01-22T21:41:47.969774","status":"completed"},"tags":[]}},{"id":"5ccb7d48-a69d-44fc-837b-157588972ff7","cell_type":"code","source":"from collections import defaultdict\nimport random\n\nimport pandas as pd\nimport numpy as np\n\nimport torch\nimport fastai\nimport torchvision\nimport timm\n\nfrom timm.layers.adaptive_avgmax_pool import SelectAdaptivePool2d\nfrom torch.nn import Flatten\n\nfrom fastai.vision.learner import *\nfrom fastai.data.all import *\nfrom fastai.vision.all import *\nfrom fastai.metrics import ActivationType\n\nfrom torchvision.models import ResNet50_Weights\n\nfrom sklearn.model_selection import StratifiedKFold, train_test_split, GroupShuffleSplit\nfrom sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, roc_auc_score\n\n\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\n    \nfrom pathlib import Path\nimport multiprocessing as mp\nimport cv2","metadata":{"papermill":{"duration":9.739099,"end_time":"2024-01-22T21:41:57.734861","exception":false,"start_time":"2024-01-22T21:41:47.995762","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:43.579316Z","iopub.execute_input":"2024-01-24T09:50:43.57977Z","iopub.status.idle":"2024-01-24T09:50:43.59106Z","shell.execute_reply.started":"2024-01-24T09:50:43.579731Z","shell.execute_reply":"2024-01-24T09:50:43.589875Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"582f0f76-59fe-4174-9fb9-ea2c3ba9e282","cell_type":"markdown","source":"### Helpers","metadata":{"papermill":{"duration":0.013039,"end_time":"2024-01-22T21:41:57.760771","exception":false,"start_time":"2024-01-22T21:41:57.747732","status":"completed"},"tags":[]}},{"id":"c274adf6-973d-42b1-9cae-1bb31425dcfe","cell_type":"code","source":"def set_seed(use_cuda: bool, seed_value: int = 42) -> None:\n    \"\"\"\n    Set random number generator seeds for CPU and GPU (CUDA) processing to ensure reproducibility.\n\n    :param use_cuda: Boolean flag indicating if GPU (CUDA) processing is available.\n    :param seed_value: Seed value to be used for random number generation (default: 1234).\n    :return: None\n    \"\"\"\n    np.random.seed(seed_value) # CPU Numpy variable (NumPy CPU results)\n    torch.manual_seed(seed_value) # CPU PyTorch variable \n    random.seed(seed_value) # Python variable (random Python module)\n    if use_cuda: # if GPU (CUDA) processing is available\n        torch.cuda.manual_seed(seed_value) # GPU PyTorch variable\n        torch.cuda.manual_seed_all(seed_value) # GPU PyTorch variable\n        torch.backends.cudnn.deterministic = True # cuDNN (CUDA Deep Neural Network library) set as deterministic\n        torch.backends.cudnn.benchmark = False\n    print(f'Seed set to: {seed_value}, with use_cuda={use_cuda}')","metadata":{"papermill":{"duration":0.022158,"end_time":"2024-01-22T21:41:57.79598","exception":false,"start_time":"2024-01-22T21:41:57.773822","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:44.325823Z","iopub.execute_input":"2024-01-24T09:50:44.326854Z","iopub.status.idle":"2024-01-24T09:50:44.33558Z","shell.execute_reply.started":"2024-01-24T09:50:44.326807Z","shell.execute_reply":"2024-01-24T09:50:44.334403Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"b10d53cd-9312-46e2-a459-e228b1f18555","cell_type":"markdown","source":"### Config","metadata":{"papermill":{"duration":0.011784,"end_time":"2024-01-22T21:41:57.819813","exception":false,"start_time":"2024-01-22T21:41:57.808029","status":"completed"},"tags":[]}},{"id":"83f9c37b-3407-460e-bcc4-eb5ac75fb225","cell_type":"code","source":"SEED=42\nNUM_EPOCHS=20\nNUM_SPLITS=4\nBATCH_SIZE=32\nSIZE=512\nDEV=True\nEXTERNAL=False\nOVERSAMPLING=True\nUNDERSAMPLING=False\n\nassert not (OVERSAMPLING and UNDERSAMPLING), \"Either OVERSAMPLING or UNDERSAMPLING should be True, but not both.\"\n\nmodel_weights_path = '/kaggle/working/rsna-trained-model-weights'\n\ndata_base_path = \"/kaggle/input/rsna-breast-cancer-detection/\"\nimages_base_path = \"/kaggle/input/rsna-mammography-images-as-pngs/\"\ndf_images_path = os.path.join(images_base_path, f\"images_as_pngs_cv2_vl_{SIZE}/train_images_processed_cv2_vl_{SIZE}/\")\n\ndf = pd.read_csv(os.path.join(data_base_path, \"train.csv\"))\ndf[\"image_path\"] = df_images_path + df[\"patient_id\"].astype(str) + \"/\" + df[\"image_id\"].astype(str) + \".png\"\n\nuse_cuda = torch.cuda.is_available()\nset_seed(use_cuda, SEED)\nDEVICE = torch.device('cuda' if use_cuda else 'cpu')\nprint('Device available now:', DEVICE)","metadata":{"papermill":{"duration":0.361304,"end_time":"2024-01-22T21:41:58.193007","exception":false,"start_time":"2024-01-22T21:41:57.831703","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:45.249098Z","iopub.execute_input":"2024-01-24T09:50:45.250004Z","iopub.status.idle":"2024-01-24T09:50:45.45262Z","shell.execute_reply.started":"2024-01-24T09:50:45.249958Z","shell.execute_reply":"2024-01-24T09:50:45.451592Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"18458130-54ac-4b4e-9e66-d4ab79e48e36","cell_type":"markdown","source":"###  DEV","metadata":{"papermill":{"duration":0.012052,"end_time":"2024-01-22T21:41:58.218224","exception":false,"start_time":"2024-01-22T21:41:58.206172","status":"completed"},"tags":[]}},{"id":"5aa9dc5e-479c-492b-9c75-80976402c5dd","cell_type":"code","source":"if DEV:\n    print(\"DEV SPLIT\")\n    original_dist = df['cancer'].value_counts(normalize=True)\n    original_size = len(df)\n    sample_size = 0.2\n    unique_patients = df['patient_id'].unique()\n    external_patients, sampled_patients = train_test_split(unique_patients, test_size=sample_size, random_state=SEED, stratify=df.groupby('patient_id')['cancer'].max().values)\n    df = df[df['patient_id'].isin(sampled_patients)]\n    sampled_dist = df['cancer'].value_counts(normalize=True)\n    sampled_size = len(df)\n    print(\"Dev is true then the number of samples is reduced to: \" + str(sample_size*100) + \"%\\n\")\n    print(\"Original Size:\\n\", original_size)\n    print(\"Original Distribution:\\n\", original_dist)\n    print(\"----------\")\n    print(\"Sampled Size:\\n\", sampled_size)\n    print(\"Sampled Distribution:\\n\", sampled_dist)\n    print()","metadata":{"papermill":{"duration":0.062568,"end_time":"2024-01-22T21:41:58.293009","exception":false,"start_time":"2024-01-22T21:41:58.230441","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:45.531985Z","iopub.execute_input":"2024-01-24T09:50:45.532331Z","iopub.status.idle":"2024-01-24T09:50:45.56023Z","shell.execute_reply.started":"2024-01-24T09:50:45.532303Z","shell.execute_reply":"2024-01-24T09:50:45.559253Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"229603ce-bfb4-49dc-bd9c-646f61572902","cell_type":"code","source":"df = df.reset_index(drop=True)","metadata":{"papermill":{"duration":0.021265,"end_time":"2024-01-22T21:41:58.32808","exception":false,"start_time":"2024-01-22T21:41:58.306815","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:45.68404Z","iopub.execute_input":"2024-01-24T09:50:45.684459Z","iopub.status.idle":"2024-01-24T09:50:45.69229Z","shell.execute_reply.started":"2024-01-24T09:50:45.684428Z","shell.execute_reply":"2024-01-24T09:50:45.691054Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"8bbd79b4-df78-4efb-871a-6f7f2ee45e18","cell_type":"code","source":"print(\"TRAIN TEST SPLIT\")\nprint()\nunique_patients = df['patient_id'].unique()\ntrain_patients, test_patients = train_test_split(unique_patients, test_size=0.2, random_state=SEED, stratify=df.groupby('patient_id')['cancer'].max().values)\n\n# Create train and test dataframes based on the patient_id split\ntrain_df = df[df['patient_id'].isin(train_patients)]\ntest_df = df[df['patient_id'].isin(test_patients)]\n\ntrain_patients_set = set(train_df['patient_id'])\ntest_patients_set = set(test_df['patient_id'])\nassert train_patients_set.isdisjoint(test_patients_set), \"Overlap found between train and test sets\"\n\ntrain_dist = train_df['cancer'].value_counts(normalize=False)\ntest_dist = test_df['cancer'].value_counts(normalize=False)\ntrain_dist_normalized = train_df['cancer'].value_counts(normalize=True)\ntest_dist_normalized = test_df['cancer'].value_counts(normalize=True)\n\nprint(\"Train Distribution:\\n\", train_dist)\nprint()\nprint(\"Test Distribution:\\n\", test_dist)\nprint()\nprint(\"Train Distribution Normalized:\\n\", train_dist_normalized)\nprint()\nprint(\"Test Distribution Normalized:\\n\", test_dist_normalized)","metadata":{"papermill":{"duration":0.036275,"end_time":"2024-01-22T21:41:58.417362","exception":false,"start_time":"2024-01-22T21:41:58.381087","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:45.861523Z","iopub.execute_input":"2024-01-24T09:50:45.862099Z","iopub.status.idle":"2024-01-24T09:50:45.886706Z","shell.execute_reply.started":"2024-01-24T09:50:45.862047Z","shell.execute_reply":"2024-01-24T09:50:45.885863Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"3aa270e0-8e6f-4b61-9667-0be10d3fdc4f","cell_type":"code","source":"train_image_id_to_cancer_label = train_df.set_index('image_id')['cancer'].to_dict()\ntest_image_id_to_cancer_label = test_df.set_index('image_id')['cancer'].to_dict()\n\ndef labeling_function(path, labels_dict):\n    try:\n        image_id = int(path.stem) \n        return labels_dict[image_id] \n    except KeyError:\n        return -1","metadata":{"papermill":{"duration":0.026375,"end_time":"2024-01-22T21:41:58.45627","exception":false,"start_time":"2024-01-22T21:41:58.429895","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:46.143688Z","iopub.execute_input":"2024-01-24T09:50:46.144498Z","iopub.status.idle":"2024-01-24T09:50:46.15624Z","shell.execute_reply.started":"2024-01-24T09:50:46.14447Z","shell.execute_reply":"2024-01-24T09:50:46.155218Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"1057f842-da61-44be-8a9c-dc71c0948921","cell_type":"code","source":"patient_labels = train_df.groupby('patient_id')['cancer'].max().reset_index()\nskf = StratifiedKFold(n_splits=NUM_SPLITS, shuffle=True, random_state=SEED)\nsplits = list(skf.split(patient_labels['patient_id'], patient_labels['cancer']))","metadata":{"papermill":{"duration":0.023623,"end_time":"2024-01-22T21:41:58.493415","exception":false,"start_time":"2024-01-22T21:41:58.469792","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:46.497696Z","iopub.execute_input":"2024-01-24T09:50:46.498613Z","iopub.status.idle":"2024-01-24T09:50:46.507631Z","shell.execute_reply.started":"2024-01-24T09:50:46.498579Z","shell.execute_reply":"2024-01-24T09:50:46.506658Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"33ea192f-6a2d-46dc-ba1e-99c8e5e749de","cell_type":"code","source":"def filter_training_items(items, split_index):\n    \"\"\"Filter out items that belong to the training set.\"\"\"\n    train_patient_ids = set(patient_labels['patient_id'][splits[split_index][0]].values)\n    training_items = [item for item in items if int(Path(item).parent.name) in train_patient_ids]\n    return training_items","metadata":{"papermill":{"duration":0.020292,"end_time":"2024-01-22T21:41:58.527382","exception":false,"start_time":"2024-01-22T21:41:58.50709","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:46.699137Z","iopub.execute_input":"2024-01-24T09:50:46.69951Z","iopub.status.idle":"2024-01-24T09:50:46.705494Z","shell.execute_reply.started":"2024-01-24T09:50:46.699483Z","shell.execute_reply":"2024-01-24T09:50:46.704443Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"f4aa8968-12ca-4ee3-a8f8-9dc694ad92b7","cell_type":"code","source":"import random\n\ndef get_undersampling_factor(train_df):\n    class_counts = train_df['cancer'].value_counts()\n    factor = class_counts[1] // class_counts[0]  # Ratio of minority to majority class\n    return factor\n\ndef get_all_items_with_undersampling(image_dir_path, split_index):\n    train_patient_ids = set(patient_labels['patient_id'][splits[split_index][0]].values)\n    valid_patient_ids = set(patient_labels['patient_id'][splits[split_index][1]].values)\n\n    cancer_items = []\n    non_cancer_items = []\n    items = []\n\n    # Collect cancer and non-cancer items from training set\n    for p in get_image_files(image_dir_path):\n        patient_id = int(Path(p).parent.name)\n        label = labeling_function(p, train_image_id_to_cancer_label)\n        \n        if patient_id in train_patient_ids:\n            if label == 1:\n                cancer_items.append(p)\n            elif label == 0:\n                non_cancer_items.append(p)\n\n    # Handle undersampling if enabled\n    if UNDERSAMPLING:\n        undersampled_non_cancer_items = random.sample(non_cancer_items, len(cancer_items))\n        items.extend(cancer_items + undersampled_non_cancer_items)\n        random.shuffle(items)\n    else:\n        items.extend(cancer_items)\n        items.extend(non_cancer_items)\n\n    # Add valid items without any sampling\n    for p in get_image_files(image_dir_path):\n        patient_id = int(Path(p).parent.name)\n        if patient_id in valid_patient_ids:\n            items.append(p)\n\n    return items","metadata":{"papermill":{"duration":0.024312,"end_time":"2024-01-22T21:41:58.564291","exception":false,"start_time":"2024-01-22T21:41:58.539979","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:46.865544Z","iopub.execute_input":"2024-01-24T09:50:46.865825Z","iopub.status.idle":"2024-01-24T09:50:46.875975Z","shell.execute_reply.started":"2024-01-24T09:50:46.865801Z","shell.execute_reply":"2024-01-24T09:50:46.875087Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"a1ede032-74dc-4b6f-977d-01735bf061e6","cell_type":"code","source":"def get_oversampling_factor(train_df, desired_ratio=2):\n    total_parts = desired_ratio + 1 \n    majority_percentage = (desired_ratio / total_parts) * 100\n    minority_percentage = (1 / total_parts) * 100\n    majority_percentage, minority_percentage\n    print(f\"Percentage after oversampling\\nMajority Percentage: {round(majority_percentage, 2)}\\nMinority Percentage: {round(minority_percentage, 2)}\")\n    class_counts = train_df['cancer'].value_counts()\n    factor = (class_counts[0] / desired_ratio) / class_counts[1]\n    factor = math.floor(factor)\n    return factor\n\nif OVERSAMPLING:\n    oversampling_factor = get_oversampling_factor(train_df, 2)\n    print(f\"Oversampling factor, based on desired_ratio: {oversampling_factor}\" )\n\ndef get_all_items_with_oversampling(image_dir_path, split_index):\n    train_patient_ids = set(patient_labels['patient_id'][splits[split_index][0]].values)\n    valid_patient_ids = set(patient_labels['patient_id'][splits[split_index][1]].values)\n    items = []\n    for p in get_image_files(image_dir_path):\n        patient_id = int(Path(p).parent.name)\n        if patient_id in train_patient_ids or patient_id in valid_patient_ids:\n            label = labeling_function(p, train_image_id_to_cancer_label)\n            if OVERSAMPLING and patient_id in train_patient_ids:\n                items.extend([p] * (oversampling_factor if label == 1 else 1))\n            else:\n                items.append(p)\n    random.shuffle(items)\n    return items","metadata":{"papermill":{"duration":0.025272,"end_time":"2024-01-22T21:41:58.602249","exception":false,"start_time":"2024-01-22T21:41:58.576977","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:47.077595Z","iopub.execute_input":"2024-01-24T09:50:47.078452Z","iopub.status.idle":"2024-01-24T09:50:47.089951Z","shell.execute_reply.started":"2024-01-24T09:50:47.078412Z","shell.execute_reply":"2024-01-24T09:50:47.088978Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"f145e847-ec56-44fd-a0ce-d9c7457af8b0","cell_type":"code","source":"get_all_items = get_all_items_with_oversampling\nif OVERSAMPLING:\n    print(\"TRAIN OVERSAMPLING METHOD\")\n    get_all_items = get_all_items_with_oversampling\nelif UNDERSAMPLING:\n    print(\"TRAIN UNDERSAMPLING METHOD\")\n    get_all_items = get_all_items_with_undersampling","metadata":{"papermill":{"duration":0.020251,"end_time":"2024-01-22T21:41:58.634654","exception":false,"start_time":"2024-01-22T21:41:58.614403","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:47.307315Z","iopub.execute_input":"2024-01-24T09:50:47.308035Z","iopub.status.idle":"2024-01-24T09:50:47.313475Z","shell.execute_reply.started":"2024-01-24T09:50:47.308006Z","shell.execute_reply":"2024-01-24T09:50:47.312549Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"7eab38d1-42cb-4acb-95b4-3102232b2480","cell_type":"code","source":"def check_class_distribution(items, label_dict):\n    class_counts = {'cancer': 0, 'non-cancer': 0}\n    for item in items:\n        image_id = int(Path(item).stem)\n        label = label_dict.get(image_id, -1)\n        if label == 1:\n            class_counts['cancer'] += 1\n        elif label == 0:\n            class_counts['non-cancer'] += 1\n    return class_counts\n\ndef aggregate_class_distribution(image_dir_path, num_splits, label_dict):\n    total_class_counts = {'non-cancer': 0, 'cancer': 0}\n    for split_index in range(num_splits):\n        items = get_all_items(image_dir_path, split_index)\n        training_items = filter_training_items(items, split_index)\n        class_counts = check_class_distribution(training_items, label_dict)\n        total_class_counts['cancer'] += class_counts['cancer']\n        total_class_counts['non-cancer'] += class_counts['non-cancer']\n    return total_class_counts\n\ndef aggregate_and_normalize_class_distribution(image_dir_path, num_splits, label_dict):\n    total_class_counts = aggregate_class_distribution(image_dir_path, num_splits, label_dict)\n    total_samples = total_class_counts['cancer'] + total_class_counts['non-cancer']\n    normalized_class_counts = {k: v / total_samples for k, v in total_class_counts.items()}\n    return normalized_class_counts\n\nprint(\"TRAIN DISTRIBUTION:\")\ntotal_distribution = aggregate_class_distribution(df_images_path, NUM_SPLITS, train_image_id_to_cancer_label)\nnormalized_distribution = aggregate_and_normalize_class_distribution(df_images_path, NUM_SPLITS, train_image_id_to_cancer_label)\nprint(normalized_distribution)","metadata":{"papermill":{"duration":133.27256,"end_time":"2024-01-22T21:44:11.919491","exception":false,"start_time":"2024-01-22T21:41:58.646931","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:50:48.704657Z","iopub.execute_input":"2024-01-24T09:50:48.705049Z","iopub.status.idle":"2024-01-24T09:52:58.142032Z","shell.execute_reply.started":"2024-01-24T09:50:48.705019Z","shell.execute_reply":"2024-01-24T09:52:58.141028Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"8df6d2b0-402b-4bd5-b792-960807287922","cell_type":"code","source":"def calculate_label_weights(train_df, class_counts, smoothing_factor=1.0):\n    weights = (1 / class_counts) * (len(train_df) / len(class_counts))\n    weights = weights * smoothing_factor\n    return torch.tensor(weights.values).float()\n\nclass_counts = pd.Series({0: total_distribution['non-cancer'], 1: total_distribution['cancer']}, name='count')\n# label_smoothing_weights = calculate_label_weights(train_df, class_counts)\nlabel_smoothing_weights = torch.tensor([1,10]).float()\nif torch.cuda.is_available():\n  label_smoothing_weights = label_smoothing_weights.cuda()\nprint(label_smoothing_weights)","metadata":{"papermill":{"duration":0.508779,"end_time":"2024-01-22T21:44:12.441317","exception":false,"start_time":"2024-01-22T21:44:11.932538","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:52:58.143562Z","iopub.execute_input":"2024-01-24T09:52:58.143857Z","iopub.status.idle":"2024-01-24T09:52:58.153331Z","shell.execute_reply.started":"2024-01-24T09:52:58.143831Z","shell.execute_reply":"2024-01-24T09:52:58.152444Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"616ebaaf-43a0-42e6-8e4e-9d1e1b543952","cell_type":"code","source":"from fastai.vision.augment import aug_transforms, RandomResizedCrop, Brightness, Contrast, RandomErasing\n\ndef train_labeling_wrapper(path):\n    return labeling_function(path, train_image_id_to_cancer_label)\n\ndef custom_splitter(items, split_index):\n    train_patient_ids = set(patient_labels['patient_id'][splits[split_index][0]].values)\n    valid_patient_ids = set(patient_labels['patient_id'][splits[split_index][1]].values)\n    train_idxs, valid_idxs = [], []\n    for idx, item in enumerate(items):\n        patient_id = int(Path(item).parent.name)  \n        if patient_id in train_patient_ids:\n            train_idxs.append(idx)\n        elif patient_id in valid_patient_ids:\n            valid_idxs.append(idx)\n    return train_idxs, valid_idxs\n\n\naug_tfms = aug_transforms(\n    do_flip=True,\n    flip_vert=False,\n    max_rotate=10.0,\n    min_zoom=1.0,\n    max_zoom=1.1,\n    max_lighting=0.2,\n    max_warp=0.2,\n    p_affine=0.75,\n    p_lighting=0.75,\n)\n\n\ndef get_train_dataloaders(image_dir_path, split_index, batch_size=BATCH_SIZE):\n    dblock = DataBlock(\n        blocks=(ImageBlock, CategoryBlock),\n        get_items=partial(get_all_items, split_index=split_index),\n        splitter=partial(custom_splitter, split_index=split_index),\n        get_y=train_labeling_wrapper,\n        batch_tfms=aug_tfms\n    )\n    dls = dblock.dataloaders(df_images_path, batch_size=batch_size)\n    return dls\n\ndef get_learner(split_index, arch=resnet18):\n    learner = vision_learner(\n        get_train_dataloaders(df_images_path, split_index),\n        arch,\n        custom_head=nn.Sequential(SelectAdaptivePool2d(pool_type='avg', flatten=Flatten()), nn.Linear(1280, 2)),\n        metrics=[\n            RocAucBinary(),\n            BrierScore()\n        ],\n        loss_func=FocalLossFlat(weight=torch.tensor([1,10]).float().cuda(), gamma=2.0),\n        pretrained=True,\n        normalize=False\n    ).to_fp16()\n    return learner","metadata":{"papermill":{"duration":0.03096,"end_time":"2024-01-22T21:44:12.485583","exception":false,"start_time":"2024-01-22T21:44:12.454623","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:52:58.1543Z","iopub.execute_input":"2024-01-24T09:52:58.154614Z","iopub.status.idle":"2024-01-24T09:52:58.176169Z","shell.execute_reply.started":"2024-01-24T09:52:58.154591Z","shell.execute_reply":"2024-01-24T09:52:58.175312Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"11635637-bb2c-415a-b76a-7e4ae0514365","cell_type":"markdown","source":"### Sanity Check","metadata":{"papermill":{"duration":0.012597,"end_time":"2024-01-22T21:44:12.51106","exception":false,"start_time":"2024-01-22T21:44:12.498463","status":"completed"},"tags":[]}},{"id":"6ac02765-ddf5-4262-8300-ca97e8352a54","cell_type":"code","source":"def show_sample_images(dl, n=5):\n    for i, (images, labels) in enumerate(dl):\n        if i >= n: break\n        dls.show_batch((images, labels), max_n=4, nrows=1)\n        print(\"Labels:\", labels)","metadata":{"execution":{"iopub.execute_input":"2024-01-22T21:44:12.538287Z","iopub.status.busy":"2024-01-22T21:44:12.537626Z","iopub.status.idle":"2024-01-22T21:44:12.542636Z","shell.execute_reply":"2024-01-22T21:44:12.541725Z"},"papermill":{"duration":0.020631,"end_time":"2024-01-22T21:44:12.544501","exception":false,"start_time":"2024-01-22T21:44:12.52387","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"0e06bf7d-f89f-415d-bd54-5138b1003f43","cell_type":"code","source":"from collections import Counter\n\ndls = get_train_dataloaders(df_images_path, 0)\n\n# Validation\nprint(\"\\n-----VALIDATION------\\n\")\nvalid_dl = dls.valid\nprint(f\"Validation number of batches: {len(valid_dl)}\")\n\n# Check the shape of input tensors in a five validation batches\nfor i, batch in enumerate(valid_dl):\n    images, _ = batch\n    print(f\"Validation Batch {i} - Input shape:\", images.shape)\n    if i >= 5:  # Check first 5 batches \n        break\n        \n# Check class distribution by Counter\nlabel_counts = Counter()\nfor _, labels in valid_dl:\n    label_counts.update(labels.cpu().numpy())\nprint(f\"Validation set class distribution: {label_counts}\")\n\n# Check labels corresponding to given bathces\nfor i, batch in enumerate(valid_dl):\n    _, labels = batch\n    print(f\"Valid Batch {i} labels:\", labels)\n    if i >= 5:\n        break\n\n# Train\nprint(\"\\n-----TRAIN------\\n\")\ntrain_dl = dls.train\nprint(f\"Train number of batches: {len(train_dl)}\")\n\n# Check the shape of input tensors in a five train batches\nfor i, batch in enumerate(train_dl):\n    images, _ = batch\n    print(f\"Train Batch {i} - Input shape:\", images.shape)\n    if i >= 5:  # Check first 5 batches\n        break\n        \n# Check class distribution by Counter\nlabel_counts = Counter()\nfor _, labels in train_dl:\n    label_counts.update(labels.cpu().numpy())\nprint(f\"Validation set class distribution: {label_counts}\")\n\n# Check labels corresponding to given bathces\nfor i, batch in enumerate(train_dl):\n    _, labels = batch\n    print(f\"Train Batch {i} labels:\", labels)\n    if i >= 5:\n        break","metadata":{"execution":{"iopub.execute_input":"2024-01-22T21:44:12.570853Z","iopub.status.busy":"2024-01-22T21:44:12.57055Z","iopub.status.idle":"2024-01-22T21:45:05.481139Z","shell.execute_reply":"2024-01-22T21:45:05.480101Z"},"papermill":{"duration":52.926535,"end_time":"2024-01-22T21:45:05.483577","exception":false,"start_time":"2024-01-22T21:44:12.557042","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"47e3a7d3-12ac-41e4-9518-75683617f451","cell_type":"code","source":"#  show_sample_images(valid_dl, n=5)","metadata":{"execution":{"iopub.execute_input":"2024-01-22T21:45:05.514663Z","iopub.status.busy":"2024-01-22T21:45:05.514304Z","iopub.status.idle":"2024-01-22T21:45:05.518734Z","shell.execute_reply":"2024-01-22T21:45:05.517841Z"},"papermill":{"duration":0.022151,"end_time":"2024-01-22T21:45:05.520662","exception":false,"start_time":"2024-01-22T21:45:05.498511","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ace12df6-c254-4bc1-a163-10ab37014e7b","cell_type":"code","source":"#  show_sample_images(train_dl, n=5)","metadata":{"execution":{"iopub.execute_input":"2024-01-22T21:45:05.54958Z","iopub.status.busy":"2024-01-22T21:45:05.549294Z","iopub.status.idle":"2024-01-22T21:45:05.553847Z","shell.execute_reply":"2024-01-22T21:45:05.553181Z"},"papermill":{"duration":0.021163,"end_time":"2024-01-22T21:45:05.555755","exception":false,"start_time":"2024-01-22T21:45:05.534592","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"38b08fb9-9fbf-4b82-81b0-0dfe5395f251","cell_type":"markdown","source":"### Pipeline","metadata":{"papermill":{"duration":0.014199,"end_time":"2024-01-22T21:45:05.584862","exception":false,"start_time":"2024-01-22T21:45:05.570663","status":"completed"},"tags":[]}},{"id":"f680b4be-5b8e-47f0-b884-20818d3fa397","cell_type":"code","source":"from fastai.callback.mixup import MixUp\nimport shutil\nimport gc\n\ntorch.cuda.empty_cache()\ngc.collect()\n\n# ACTUAL\narch=\"tf_efficientnetv2_s.in21k\"\n# arch=\"resnet50.a1_in1k\"\n# arch=\"mobilenetv2_100.ra_in1k\"\n\n\n# GOOD - TEST FURHTER\n# arch=\"mobilenetv2_100.ra_in1k\"\n# arch=\"resnet34.a1_in1k\" ~60 AUC\n\n# OK - TO TEST FURTHER\n# arch=\"resnext50_32x4d.fb_swsl_ig1b_ft_in1k\"\n# arch=\"regnety_002.pycls_in1k\"\n\n# DECENT - MAYBE TEST\n# arch=\"resnext50_32x4d.fb_swsl_ig1b_ft_in1k\" OK BUT LONG\n# arch=\"convnextv2_nano.fcmae_ft_in22k_in1k\"\n# arch=\"convnextv2_atto.fcmae_ft_in1k\" TOO LONG\n# arch=\"mobilenetv3_large_100.ra_in1k\"\npreds, targets = [], []\n\narch_name = arch if isinstance(arch, str) else arch.__name__\ntrained_model_weights_path = Path(f\"{model_weights_path}/{arch_name}\")\n\nif trained_model_weights_path.exists() and trained_model_weights_path.is_dir():\n    shutil.rmtree(trained_model_weights_path)\ntrained_model_weights_path.mkdir(parents=True)\n\nSPLIT=0\nfor SPLIT in range(NUM_SPLITS):\n    torch.cuda.empty_cache()\n    gc.collect()\n    learn = get_learner(SPLIT, arch)\n    learn.unfreeze()\n    \n    learn.add_cb(EarlyStoppingCallback(monitor='valid_loss', min_delta=0.05, patience=5))\n\n    lr_valley = learn.lr_find(suggest_funcs=valley, num_it=200)\n    print(f\"Suggested Learning Rate for Split {SPLIT}: Valley point: {lr_valley}\")\n    chosen_lr = lr_valley \n\n    learn.fit_one_cycle(\n    NUM_EPOCHS,\n    chosen_lr,\n    pct_start=0.3 \n    )\n    model_path = trained_model_weights_path/str(SPLIT)\n    print(f\"Saving as: {model_path}\")\n    learn.save(model_path)\n    output = learn.get_preds()\n    preds.append(output[0])  \n    targets.append(output[1])\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"execution":{"iopub.execute_input":"2024-01-22T21:45:05.614648Z","iopub.status.busy":"2024-01-22T21:45:05.614325Z","iopub.status.idle":"2024-01-23T00:01:17.954986Z","shell.execute_reply":"2024-01-23T00:01:17.953839Z"},"papermill":{"duration":8172.3581,"end_time":"2024-01-23T00:01:17.957314","exception":false,"start_time":"2024-01-22T21:45:05.599214","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"578998ba-9bd2-411a-88cc-84c34c6cd299","cell_type":"markdown","source":"# II. Predicting On Test","metadata":{"papermill":{"duration":0.019718,"end_time":"2024-01-23T00:01:17.997989","exception":false,"start_time":"2024-01-23T00:01:17.978271","status":"completed"},"tags":[]}},{"id":"63f651aa-2c42-4710-abc2-333a55dcdb73","cell_type":"code","source":"test_image_id_to_cancer_label = test_df.set_index('image_id')['cancer'].to_dict()\n\ndef test_labeling_wrapper(path):\n    return labeling_function(path, test_image_id_to_cancer_label)","metadata":{"papermill":{"duration":0.033087,"end_time":"2024-01-23T00:01:18.050791","exception":false,"start_time":"2024-01-23T00:01:18.017704","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:54:47.852883Z","iopub.execute_input":"2024-01-24T09:54:47.853797Z","iopub.status.idle":"2024-01-24T09:54:47.865084Z","shell.execute_reply.started":"2024-01-24T09:54:47.853753Z","shell.execute_reply":"2024-01-24T09:54:47.86416Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"aa0a870a-6900-4af3-b497-72d12a1d61d6","cell_type":"code","source":"def get_test_items(test_df, image_dir_path, _=None):\n    return [Path(row[\"image_path\"]) for _, row in test_df.iterrows()]\n\ndef get_test_dataloader(test_df, image_dir_path, batch_size=BATCH_SIZE):\n    test_db = DataBlock(\n        blocks=(ImageBlock, CategoryBlock),\n        get_items=partial(get_test_items, test_df, image_dir_path),\n        get_y=test_labeling_wrapper,\n        batch_tfms=None\n    )\n    test_dls = test_db.dataloaders(Path(image_dir_path), batch_size=batch_size)\n    return test_dls\n\ntest_dls = get_test_dataloader(test_df, df_images_path)","metadata":{"papermill":{"duration":0.339816,"end_time":"2024-01-23T00:01:18.410835","exception":false,"start_time":"2024-01-23T00:01:18.071019","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:54:48.022466Z","iopub.execute_input":"2024-01-24T09:54:48.02278Z","iopub.status.idle":"2024-01-24T09:54:48.32738Z","shell.execute_reply.started":"2024-01-24T09:54:48.022755Z","shell.execute_reply":"2024-01-24T09:54:48.326391Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"dccb72ca-efc8-4447-ab81-0648735cffe8","cell_type":"markdown","source":"# PROBABILITIES","metadata":{"papermill":{"duration":0.01962,"end_time":"2024-01-23T00:01:18.451251","exception":false,"start_time":"2024-01-23T00:01:18.431631","status":"completed"},"tags":[]}},{"id":"0a0339a7-d2d6-4cfe-9c9d-190dd70bfe2f","cell_type":"code","source":"from sklearn.metrics import roc_auc_score, precision_recall_curve, auc\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import roc_curve, auc, precision_recall_curve, average_precision_score\nfrom sklearn.metrics import roc_auc_score\nfrom scipy.stats import sem\nimport numpy as np\nimport numpy as np\nfrom sklearn.metrics import precision_recall_curve, auc\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import precision_recall_curve\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.calibration import calibration_curve","metadata":{"papermill":{"duration":0.348926,"end_time":"2024-01-23T00:01:18.820158","exception":false,"start_time":"2024-01-23T00:01:18.471232","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:54:48.402583Z","iopub.execute_input":"2024-01-24T09:54:48.402915Z","iopub.status.idle":"2024-01-24T09:54:48.409822Z","shell.execute_reply.started":"2024-01-24T09:54:48.402887Z","shell.execute_reply":"2024-01-24T09:54:48.408725Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"ef3fc01d-9605-4613-97af-ea3688231b28","cell_type":"code","source":"def calculate_auc_confidence_interval(y_true, y_scores, alpha=0.95):\n    auc_scores = []\n    n_bootstraps = 1000\n    rng = np.random.RandomState(42)\n    for i in range(n_bootstraps):\n        indices = rng.randint(0, len(y_scores), len(y_scores))\n        if len(np.unique(y_true[indices])) < 2:\n            continue\n        score = roc_auc_score(y_true[indices], y_scores[indices])\n        auc_scores.append(score)\n    sorted_scores = np.array(auc_scores)\n    sorted_scores.sort()\n    confidence_lower = sorted_scores[int((1.0 - alpha) / 2.0 * n_bootstraps)]\n    confidence_upper = sorted_scores[int((alpha + (1.0 - alpha) / 2.0) * n_bootstraps)]\n    return confidence_lower, confidence_upper\n\n\n\ndef plot_roc_auc_curve(true_labels, predictions):\n    fpr, tpr, _ = roc_curve(true_labels, average_preds[:, 1])\n    roc_auc = roc_auc_score(true_labels, average_preds[:, 1])\n\n    plt.figure(figsize=(8, 6))\n    plt.plot(fpr, tpr, color='blue', lw=2, label='ROC curve (area = %0.2f)' % roc_auc)\n    plt.fill_between(fpr, tpr, alpha=0.2)\n    plt.plot([0, 1], [0, 1], color='gray', lw=2, linestyle='--')\n    plt.xlim([0.0, 1.0])\n    plt.ylim([0.0, 1.05])\n    plt.xlabel('False Positive Rate')\n    plt.ylabel('True Positive Rate')\n    plt.title('Receiver Operating Characteristic (ROC)')\n    plt.legend(loc=\"lower right\")\n    plt.show()\n    return roc_auc\n\ndef calculate_precision_recall_gain(true_labels, predictions):\n    pi = np.mean(true_labels)\n    precision, recall, _ = precision_recall_curve(true_labels, predictions)\n    precision_gain = (precision - pi) / ((1 - pi) * precision)\n    recall_gain = (recall - pi) / ((1 - pi) * recall)\n    precision_gain[precision_gain < 0] = 0\n    recall_gain[recall_gain < 0] = 0\n    return precision_gain, recall_gain\n\ndef plot_pr_gain_curve(true_labels, predictions):\n    precision_gain, recall_gain = calculate_precision_recall_gain(true_labels, predictions)\n    pr_gain_auc = auc(recall_gain, precision_gain)\n    plt.figure(figsize=(8, 6))\n    plt.plot(recall_gain, precision_gain, color='blue', lw=2, label=f'PR-Gain curve (AUPRG = {pr_gain_auc:.2f})')\n    plt.fill_between(recall_gain, precision_gain, alpha=0.2, color='blue')\n    plt.xlim([0.0, 1.0])\n    plt.ylim([0.0, max(precision_gain) + 0.05])\n    plt.xlabel('Recall Gain')\n    plt.ylabel('Precision Gain')\n    plt.title('Precision-Recall-Gain Curve')\n    plt.legend(loc=\"lower left\")\n    plt.show()\n    return pr_gain_auc\n\n","metadata":{"papermill":{"duration":0.039033,"end_time":"2024-01-23T00:01:18.880271","exception":false,"start_time":"2024-01-23T00:01:18.841238","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-24T09:54:48.593113Z","iopub.execute_input":"2024-01-24T09:54:48.594032Z","iopub.status.idle":"2024-01-24T09:54:48.616904Z","shell.execute_reply.started":"2024-01-24T09:54:48.593975Z","shell.execute_reply":"2024-01-24T09:54:48.615933Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"f5478d72-28e0-4c69-b0dc-4c50cd4f6e50","cell_type":"code","source":"def get_model(trained_model_dir_path, split):\n    model_path = Path(f\"{trained_model_dir_path}/{split}\")\n    print(f\"Loading model {model_path}\")\n    return learn.load(model_path)\n\npreds_all = []\ntargets_all = []\n\nfor SPLIT in range(NUM_SPLITS):\n    learn = get_model(trained_model_weights_path, SPLIT)\n    learn.dls = test_dls \n    preds, targets = learn.get_preds(dl=test_dls.train) \n    preds_all.append(preds)\n    targets_all.append(targets)\n\naverage_preds = torch.mean(torch.stack(preds_all), dim=0)\ntrue_labels = targets_all[0].numpy()","metadata":{"execution":{"iopub.execute_input":"2024-01-23T00:01:18.921365Z","iopub.status.busy":"2024-01-23T00:01:18.92109Z","iopub.status.idle":"2024-01-23T00:02:11.441325Z","shell.execute_reply":"2024-01-23T00:02:11.440419Z"},"papermill":{"duration":52.543409,"end_time":"2024-01-23T00:02:11.443839","exception":false,"start_time":"2024-01-23T00:01:18.90043","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ed7fe73a-bd14-4af8-a70c-6cfdecd45501","cell_type":"markdown","source":"### a. AUC ROC + Plot","metadata":{"execution":{"iopub.execute_input":"2024-01-15T16:57:40.865901Z","iopub.status.busy":"2024-01-15T16:57:40.865454Z","iopub.status.idle":"2024-01-15T16:57:40.870561Z","shell.execute_reply":"2024-01-15T16:57:40.869553Z","shell.execute_reply.started":"2024-01-15T16:57:40.865871Z"},"papermill":{"duration":0.021344,"end_time":"2024-01-23T00:02:11.48721","exception":false,"start_time":"2024-01-23T00:02:11.465866","status":"completed"},"tags":[]}},{"id":"8bb8a8b1-e313-4a2e-9515-7ef37d325cc3","cell_type":"code","source":"roc_auc = plot_roc_auc_curve(true_labels, average_preds[:, 1])\nprint(\"ROC AUC:\", roc_auc)\n\nauc_lower, auc_upper = calculate_auc_confidence_interval(true_labels, average_preds[:, 1])\nprint(f\"AUC: {roc_auc} with 95% confidence interval [{auc_lower:.3f}, {auc_upper:.3f}]\")","metadata":{"execution":{"iopub.execute_input":"2024-01-23T00:02:11.532641Z","iopub.status.busy":"2024-01-23T00:02:11.532328Z","iopub.status.idle":"2024-01-23T00:02:13.958133Z","shell.execute_reply":"2024-01-23T00:02:13.957254Z"},"papermill":{"duration":2.450299,"end_time":"2024-01-23T00:02:13.960341","exception":false,"start_time":"2024-01-23T00:02:11.510042","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8cc48bfb-03e2-4849-8fc3-9211ab021d33","cell_type":"markdown","source":"### b. PR-Gain-AUC + Plot","metadata":{"execution":{"iopub.execute_input":"2024-01-15T17:08:16.035371Z","iopub.status.busy":"2024-01-15T17:08:16.034988Z","iopub.status.idle":"2024-01-15T17:08:16.03961Z","shell.execute_reply":"2024-01-15T17:08:16.038637Z","shell.execute_reply.started":"2024-01-15T17:08:16.035342Z"},"papermill":{"duration":0.022796,"end_time":"2024-01-23T00:02:14.006435","exception":false,"start_time":"2024-01-23T00:02:13.983639","status":"completed"},"tags":[]}},{"id":"9a79e955-b94b-4783-8d53-a0d62c5409d3","cell_type":"code","source":"pr_gain_auc = plot_pr_gain_curve(true_labels, average_preds[:, 1])\nprint(\"PR-Gain AUC:\", pr_gain_auc)","metadata":{"execution":{"iopub.execute_input":"2024-01-23T00:02:14.053893Z","iopub.status.busy":"2024-01-23T00:02:14.053527Z","iopub.status.idle":"2024-01-23T00:02:14.349491Z","shell.execute_reply":"2024-01-23T00:02:14.348615Z"},"papermill":{"duration":0.321816,"end_time":"2024-01-23T00:02:14.352031","exception":false,"start_time":"2024-01-23T00:02:14.030215","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8aa17c5e-1ffc-4a02-83a7-e00d7807f580","cell_type":"markdown","source":"### c. Precision and Recall vs Threshold","metadata":{"papermill":{"duration":0.023402,"end_time":"2024-01-23T00:02:14.398425","exception":false,"start_time":"2024-01-23T00:02:14.375023","status":"completed"},"tags":[]}},{"id":"ec17a70a-464e-4c85-826a-27375450fa57","cell_type":"code","source":"precision, recall, thresholds = precision_recall_curve(true_labels, preds[:, 1])\nplt.plot(thresholds, precision[:-1], 'b--', label='Precision')\nplt.plot(thresholds, recall[:-1], 'g-', label='Recall')\nplt.xlabel('Threshold')\nplt.ylabel('Precision/Recall')\nplt.title('Precision and Recall vs. Threshold')\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.execute_input":"2024-01-23T00:02:14.448982Z","iopub.status.busy":"2024-01-23T00:02:14.448627Z","iopub.status.idle":"2024-01-23T00:02:14.753302Z","shell.execute_reply":"2024-01-23T00:02:14.752397Z"},"papermill":{"duration":0.33211,"end_time":"2024-01-23T00:02:14.755256","exception":false,"start_time":"2024-01-23T00:02:14.423146","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8842b088-286a-4c00-b353-4e05e9373d6c","cell_type":"markdown","source":"### e. Calibration Curve","metadata":{"papermill":{"duration":0.025962,"end_time":"2024-01-23T00:02:14.806635","exception":false,"start_time":"2024-01-23T00:02:14.780673","status":"completed"},"tags":[]}},{"id":"0e060e7f-54d0-4204-b393-b6058ed41610","cell_type":"code","source":"prob_true, prob_pred = calibration_curve(true_labels, average_preds[:, 1], n_bins=10)\nplt.figure(figsize=(8, 6))\nplt.plot(prob_pred, prob_true, 's-', label='Calibration curve')\nplt.plot([0, 1], [0, 1], '--', color='gray', label='Perfectly calibrated')\nplt.xlabel('Mean predicted probability')\nplt.ylabel('Fraction of positives')\nplt.legend(loc='lower right')\nplt.title('Calibration Curve')\nplt.show()","metadata":{"execution":{"iopub.execute_input":"2024-01-23T00:02:14.861158Z","iopub.status.busy":"2024-01-23T00:02:14.860774Z","iopub.status.idle":"2024-01-23T00:02:15.119673Z","shell.execute_reply":"2024-01-23T00:02:15.11873Z"},"papermill":{"duration":0.288412,"end_time":"2024-01-23T00:02:15.1216","exception":false,"start_time":"2024-01-23T00:02:14.833188","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7029d4ca-5ada-4928-b467-241575b198bf","cell_type":"markdown","source":"### f. Calculate LogLoss","metadata":{"papermill":{"duration":0.024914,"end_time":"2024-01-23T00:02:15.171083","exception":false,"start_time":"2024-01-23T00:02:15.146169","status":"completed"},"tags":[]}},{"id":"68a1bd36-a232-41ac-bebe-235171cca2a9","cell_type":"code","source":"from sklearn.metrics import log_loss\nlogloss = log_loss(true_labels, average_preds)\nprint(f\"Log Loss (Cross-Entropy Loss): {logloss}\")","metadata":{"execution":{"iopub.execute_input":"2024-01-23T00:02:15.220322Z","iopub.status.busy":"2024-01-23T00:02:15.219991Z","iopub.status.idle":"2024-01-23T00:02:15.227333Z","shell.execute_reply":"2024-01-23T00:02:15.226572Z"},"papermill":{"duration":0.034123,"end_time":"2024-01-23T00:02:15.229295","exception":false,"start_time":"2024-01-23T00:02:15.195172","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"61345e1a-e086-4f10-8a40-d47691238523","cell_type":"markdown","source":"### g. Brier Score","metadata":{"execution":{"iopub.execute_input":"2024-01-15T16:58:37.528359Z","iopub.status.busy":"2024-01-15T16:58:37.527531Z","iopub.status.idle":"2024-01-15T16:58:37.533179Z","shell.execute_reply":"2024-01-15T16:58:37.532127Z","shell.execute_reply.started":"2024-01-15T16:58:37.528323Z"},"papermill":{"duration":0.025359,"end_time":"2024-01-23T00:02:15.278646","exception":false,"start_time":"2024-01-23T00:02:15.253287","status":"completed"},"tags":[]}},{"id":"682056ba-f342-4b38-8b7e-3acf19bfc831","cell_type":"code","source":"from sklearn.metrics import brier_score_loss\nbrier_score = brier_score_loss(true_labels, average_preds[:, 1])\nprint(f\"Brier Score: {brier_score}\")","metadata":{"execution":{"iopub.execute_input":"2024-01-23T00:02:15.328575Z","iopub.status.busy":"2024-01-23T00:02:15.328256Z","iopub.status.idle":"2024-01-23T00:02:15.334918Z","shell.execute_reply":"2024-01-23T00:02:15.334068Z"},"papermill":{"duration":0.033789,"end_time":"2024-01-23T00:02:15.336918","exception":false,"start_time":"2024-01-23T00:02:15.303129","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d66e9d22-6f36-4eb5-ac7f-5eddd1d787dc","cell_type":"markdown","source":"### h. Probability Distribution Plots","metadata":{"papermill":{"duration":0.02376,"end_time":"2024-01-23T00:02:15.385955","exception":false,"start_time":"2024-01-23T00:02:15.362195","status":"completed"},"tags":[]}},{"id":"cc0e1a66-236d-4068-a0fe-d211ab3435d1","cell_type":"code","source":"preds_all = []\ntargets_all = []\n\nSPLIT = 0 \nlearn = get_learner(SPLIT, \"tf_efficientnetv2_s.in21k\")\n\nfor SPLIT in range(NUM_SPLITS):\n    learn.load(f\"/kaggle/input/effnet-weights-oversampled/{SPLIT}\")\n    learn.dls = test_dls \n    preds, targets = learn.get_preds(dl=test_dls.train) \n    preds_all.append(preds)\n    targets_all.append(targets)","metadata":{"execution":{"iopub.status.busy":"2024-01-24T09:52:58.178085Z","iopub.execute_input":"2024-01-24T09:52:58.178413Z","iopub.status.idle":"2024-01-24T09:54:17.177765Z","shell.execute_reply.started":"2024-01-24T09:52:58.178383Z","shell.execute_reply":"2024-01-24T09:54:17.176383Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"4e17f7df-49aa-44dc-a80d-0435a4c93880","cell_type":"code","source":"df = pd.DataFrame({'Predicted Probability': preds[:, 1], 'True Label': true_labels})\nplt.figure(figsize=(10, 6))\nsns.kdeplot(data=df, x='Predicted Probability', hue='True Label', common_norm=False)\nplt.title('Gaussian Probability Distribution Plot for Both Classes')\nplt.xlabel('Predicted Probability')\nplt.ylabel('Density')\nplt.xlim(0, 1)  \nplt.ylim(0, 7.1136639420207235)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-01-24T09:55:41.669891Z","iopub.execute_input":"2024-01-24T09:55:41.67087Z","iopub.status.idle":"2024-01-24T09:55:42.033581Z","shell.execute_reply.started":"2024-01-24T09:55:41.670827Z","shell.execute_reply":"2024-01-24T09:55:42.032426Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"ccf1121f-8cb5-4737-beab-02ed75ce4ed6","cell_type":"code","source":"df = pd.DataFrame({'Predicted Probability': preds[:, 1], 'True Label': true_labels})\nplt.figure(figsize=(10, 6))\nsns.kdeplot(data=df, x='Predicted Probability', hue='True Label', common_norm=False)\nplt.title('Gaussian Probability Distribution Plot for Both Classes')\nplt.xlabel('Predicted Probability')\nplt.ylabel('Density')\nplt.show()","metadata":{"execution":{"iopub.execute_input":"2024-01-23T00:02:15.435357Z","iopub.status.busy":"2024-01-23T00:02:15.435086Z","iopub.status.idle":"2024-01-23T00:02:15.898046Z","shell.execute_reply":"2024-01-23T00:02:15.89703Z"},"papermill":{"duration":0.491075,"end_time":"2024-01-23T00:02:15.90096","exception":false,"start_time":"2024-01-23T00:02:15.409885","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"1229e117-db80-4e23-bf51-c186b8159e27","cell_type":"markdown","source":"### Download","metadata":{"papermill":{"duration":0.034123,"end_time":"2024-01-23T00:02:15.967467","exception":false,"start_time":"2024-01-23T00:02:15.933344","status":"completed"},"tags":[]}},{"id":"8f4d78b0-181d-41cd-85a8-5775f01faf2b","cell_type":"code","source":"!zip -r file.zip /kaggle/working\nfrom IPython.display import FileLink\nFileLink(r'file.zip')","metadata":{"execution":{"iopub.execute_input":"2024-01-23T00:02:16.03178Z","iopub.status.busy":"2024-01-23T00:02:16.0311Z","iopub.status.idle":"2024-01-23T00:03:06.414653Z","shell.execute_reply":"2024-01-23T00:03:06.413508Z"},"papermill":{"duration":50.418154,"end_time":"2024-01-23T00:03:06.41672","exception":false,"start_time":"2024-01-23T00:02:15.998566","status":"completed"},"tags":[]},"outputs":[],"execution_count":null}]}