{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.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":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":4821454,"sourceType":"datasetVersion","datasetId":2740008},{"sourceId":181667,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":154839,"modelId":177312}],"dockerImageVersionId":30302,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<img src=\"https://i.imgur.com/DKn4ofz.png\">\n\n<center><h1> RSNA Screening Mammography Breast Cancer Detection </h1></center>\n<center><h2> - Data Explore - </h2></center>\n\n> 📌 **Scope**: The goal of this competition is to identify cases of *breast cancer* in mammograms from screening exams. It is important to have True Positives (TP) as high as possible, but False Positives (FP) also have downsides for patients. Hence the **Precission should be as high as possible**.\n\n\n### ○ Libraries","metadata":{}},{"cell_type":"code","source":"! pip install -qU \"python-gdcm\" \n! pip install -qU pydicom \n! pip install -qU pylibjpeg\n#! pip install -qU \"opencv-python-headless\"","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-11-28T16:05:56.088746Z","iopub.execute_input":"2024-11-28T16:05:56.089044Z","iopub.status.idle":"2024-11-28T16:06:24.409924Z","shell.execute_reply.started":"2024-11-28T16:05:56.088976Z","shell.execute_reply":"2024-11-28T16:06:24.408719Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nprint(cv2.__version__)","metadata":{"execution":{"iopub.status.busy":"2024-11-28T16:06:24.412255Z","iopub.execute_input":"2024-11-28T16:06:24.413111Z","iopub.status.idle":"2024-11-28T16:06:24.579321Z","shell.execute_reply.started":"2024-11-28T16:06:24.413068Z","shell.execute_reply":"2024-11-28T16:06:24.578384Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# General Libraries\nimport os\nimport re\nimport gc\nimport cv2\nimport wandb\nimport random\nimport math\nfrom glob import glob\nfrom tqdm import tqdm\nfrom pprint import pprint\nfrom time import time\nimport datetime as dtime\nfrom datetime import datetime\nimport itertools\nimport warnings\nimport pandas as pd\nimport numpy as np\nimport pydicom # for DICOM images\nfrom skimage.transform import resize\nfrom sklearn.preprocessing import LabelEncoder, normalize\n\n# For the Visuals\nimport seaborn as sns\nimport matplotlib as mpl\nfrom matplotlib import cm\nimport matplotlib.patches as patches\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nfrom matplotlib.offsetbox import AnnotationBbox, OffsetImage\nfrom matplotlib.colors import ListedColormap, LinearSegmentedColormap\nfrom matplotlib.patches import Rectangle\nfrom IPython.display import display_html\nplt.rcParams.update({'font.size': 16})\n\n# Environment check\nwarnings.filterwarnings(\"ignore\")\nos.environ[\"WANDB_SILENT\"] = \"true\"\nCONFIG = {'competition': 'RSNA_Breast_Cancer', '_wandb_kernel': 'aot'}\n\n# Custom colors\nclass clr:\n    S = '\\033[1m' + '\\033[91m'\n    E = '\\033[0m'\n    \nmy_colors = [\"#517664\", \"#73AA90\", \"#94DDBC\", \"#DAB06C\", \n             \"#DF928E\", \"#C97973\", \"#B25F57\"]\nCMAP1 = ListedColormap(my_colors)\n\nprint(clr.S+\"Notebook Color Schemes:\"+clr.E)\nsns.palplot(sns.color_palette(my_colors))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-11-28T16:06:24.580595Z","iopub.execute_input":"2024-11-28T16:06:24.58086Z","iopub.status.idle":"2024-11-28T16:06:26.136026Z","shell.execute_reply.started":"2024-11-28T16:06:24.580834Z","shell.execute_reply":"2024-11-28T16:06:26.134648Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### ○ Helpers","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport random\nimport os\nimport torch\n\n# === General Functions ===\n\ndef set_seed(seed=1234):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\ndef show_values_on_bars(axs, h_v=\"v\", space=0.4):\n    '''Plots the value at the end of the a seaborn barplot.\n    axs: the ax of the plot\n    h_v: weather or not the barplot is vertical/ horizontal'''\n    \n    def _show_on_single_plot(ax):\n        if h_v == \"v\":\n            for p in ax.patches:\n                _x = p.get_x() + p.get_width() / 2\n                _y = p.get_y() + p.get_height()\n                value = int(p.get_height())\n                ax.text(_x, _y, format(value, ','), ha=\"center\") \n        elif h_v == \"h\":\n            for p in ax.patches:\n                _x = p.get_x() + p.get_width() + float(space)\n                _y = p.get_y() + p.get_height()\n                value = int(p.get_width())\n                ax.text(_x, _y, format(value, ','), ha=\"left\")\n\n    if isinstance(axs, np.ndarray):\n        for idx, ax in np.ndenumerate(axs):\n            _show_on_single_plot(ax)\n    else:\n        _show_on_single_plot(axs)\n","metadata":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:07:22.575909Z","iopub.execute_input":"2024-11-28T16:07:22.576284Z","iopub.status.idle":"2024-11-28T16:07:24.280601Z","shell.execute_reply.started":"2024-11-28T16:07:22.576251Z","shell.execute_reply":"2024-11-28T16:07:24.279828Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv(\"/kaggle/input/rsna-breast-cancer-detection/train.csv\")\n\n# Get image path\n# Example path: '/kaggle/input/rsna-breast-cancer-detection/train_images/10706/763186195.dcm'\nbase_path = \"/kaggle/input/rsna-breast-cancer-detection/train_images/\"\nall_paths = []\nfor k in tqdm(range(len(train))):\n    row = train.iloc[k, :]\n    all_paths.append(base_path + str(row.patient_id) + \"/\" + str(row.image_id) + \".dcm\")\n    \ntrain[\"path\"] = all_paths","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:11:41.934295Z","iopub.execute_input":"2024-11-28T16:11:41.934965Z","iopub.status.idle":"2024-11-28T16:11:48.756539Z","shell.execute_reply.started":"2024-11-28T16:11:41.934932Z","shell.execute_reply":"2024-11-28T16:11:48.755631Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 1. Brief overview\n\n**Q**: How does the data look like? What is the general feel of the numbers?\n\n**A**: We have the following info:\n* `site_id` - There are 2 total hospitals from where the records were gathered, split roughly 50-50\n* `patient_id` - The unique identifier of the patient. There are **11.913** total patients\n* `image_id` - The unique identidier of the image. There are **54.706** unique images in train. *Each patient has an average of 4.5 breast scans* (with 4 being the least number of scans and 14 being the maximum number of scans per patient).\n* `laterality` - L is for the left breast, R is for the right. There are slightly more images for the R (right) breast -> 27,439 than for the L (left) breast -> 27,267.","metadata":{}},{"cell_type":"code","source":"print(clr.S+\"Records gathered in Site 1:\"+clr.E, train[\"site_id\"].value_counts().values[0], \"\\n\"+\n      clr.S+\"Records gathered in Site 2:\"+clr.E, train[\"site_id\"].value_counts().values[1])\nprint(\"-------------------------------------------------\")\nprint(clr.S+\"Total unique patients:\"+clr.E, train[\"patient_id\"].nunique())\nprint(\"-------------------------------------------------\")\nprint(clr.S+\"Total unique images:\"+clr.E, train[\"image_id\"].nunique())\nprint(\"-------------------------------------------------\")\nprint(clr.S+\"Statistics: Images per Patient\"+clr.E)\nprint(train.groupby(\"patient_id\")[\"image_id\"].count().reset_index().describe()[\"image_id\"])\nprint(\"-------------------------------------------------\")\nprint(clr.S+\"Image records count per laterality (R):\"+clr.E, train[\"laterality\"].value_counts().values[0], \"\\n\"+\n      clr.S+\"Image records count per laterality (L):\"+clr.E, train[\"laterality\"].value_counts().values[1])\nprint(\"-------------------------------------------------\")\nprint(clr.S+\"Image records count per View:\"+clr.E)\nprint(train[\"view\"].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:11:50.507781Z","iopub.execute_input":"2024-11-28T16:11:50.508119Z","iopub.status.idle":"2024-11-28T16:11:50.548773Z","shell.execute_reply.started":"2024-11-28T16:11:50.50809Z","shell.execute_reply":"2024-11-28T16:11:50.547828Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Q**: What does `view` mean? How many images per view do we have? Is there a noticable difference between images and views?\n\n**A**: So here are the main takeaways:\n* `view` - The orientation of the image. The values for this variable are:\n    * MLO (mediolateral oblique) - 27903 images. \n    * CC (craniocaudal) - 26765 images.\n    * AT (axillary tail) - 19 images.\n    * LM (latero-medial) - 10 images.\n    * ML (medio-lateral) - 8 images.\n    * LMO (latero-medial oblique) - 1 images.\n* no, I do not see any noticeable difference between the images (I am not a doctor tho).","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport pydicom\n\ndef show_view(view_name, sample_size):\n    if view_name != \"LMO\":\n        # Get image info\n        data = train[train[\"view\"] == view_name].sample(sample_size, random_state=24)\n        image_path = data[\"path\"].to_list()\n\n        # Plot\n        fig, axs = plt.subplots(1, sample_size, figsize=(23, 4))\n        axs = axs.flatten()\n\n        for k, path in enumerate(image_path):\n            axs[k].set_title(f\"{k+1}. {view_name}\", \n                             fontsize=16, color=my_colors[0], weight='bold')\n\n            img = pydicom.dcmread(path).pixel_array\n            axs[k].imshow(img, cmap=\"turbo\")\n            axs[k].axis(\"off\")\n\n        plt.tight_layout()\n        plt.show()\n    else:\n        path = train[train[\"view\"] == \"LMO\"][\"path\"].item()\n        # Plot\n        fig, axs = plt.subplots(1, sample_size, figsize=(23, 4))\n        axs = axs.flatten()\n        img = pydicom.dcmread(path).pixel_array\n        axs[0].imshow(img, cmap=\"turbo\")\n        axs[0].set_title(f\"1. LMO\", \n                         fontsize=16, color=my_colors[0], weight='bold')\n        axs[0].axis(\"off\")\n        axs[1].axis(\"off\")\n        axs[2].axis(\"off\")\n        axs[3].axis(\"off\")\n        axs[4].axis(\"off\")\n\n        plt.tight_layout()\n        plt.show()\n","metadata":{"_kg_hide-input":false,"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:12:11.062364Z","iopub.execute_input":"2024-11-28T16:12:11.063056Z","iopub.status.idle":"2024-11-28T16:12:11.072362Z","shell.execute_reply.started":"2024-11-28T16:12:11.063024Z","shell.execute_reply":"2024-11-28T16:12:11.071515Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for view_name in train[\"view\"].unique().tolist():\n    show_view(view_name, sample_size=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:12:24.012548Z","iopub.execute_input":"2024-11-28T16:12:24.012922Z","iopub.status.idle":"2024-11-28T16:13:04.986505Z","shell.execute_reply.started":"2024-11-28T16:12:24.012895Z","shell.execute_reply":"2024-11-28T16:13:04.985616Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Q**: What is the age of the patients? Is the age spectrum wide (with patients from all categories) or is it narrowed on a specific age group?\n\n**A**: Looks like overall our patients are part of a more senior age group:\n* `age` - The patient's age in years. The **average age is 58 years old**, with the vast majority of the patients having between 50 and 65 years old. There are a few outliers with very young patients (26-30 years old), as well as a few more senior patients (89 years old).","metadata":{}},{"cell_type":"code","source":"# Plot\nf, (a0, a1) = plt.subplots(2, 1, gridspec_kw={'height_ratios': [3, 1]}, figsize=(24, 15))\nsns.distplot(a=train[\"age\"], rug=True, hist=False, \n             rug_kws={\"color\": my_colors[5]},\n             kde_kws={\"color\": my_colors[5], \"lw\": 5, \"alpha\": 0.7},\n             ax=a0)\n\na0.axvline(x=58, ls=\":\", lw=2, color=\"black\")\na0.text(x=58.5, y=0.018, s=\"mean: 58\", size=17, color=\"black\", weight=\"bold\")\na0.axvline(x=26, ls=\":\", lw=2, color=\"black\")\na0.text(x=26.5, y=0.008, s=\"min: 26\", size=17, color=\"black\", weight=\"bold\")\na0.axvline(x=89, ls=\":\", lw=2, color=\"black\")\na0.text(x=84, y=0.037, s=\"max: 89\", size=17, color=\"black\", weight=\"bold\")\n\nsns.boxenplot(x=train[\"age\"], ax=a1, color=my_colors[2])\n\nplt.suptitle(\"Age Distribution\", weight=\"bold\", size=25)\nsns.despine(right=True, top=True, left=True);","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:13:12.020863Z","iopub.execute_input":"2024-11-28T16:13:12.021204Z","iopub.status.idle":"2024-11-28T16:13:13.485098Z","shell.execute_reply.started":"2024-11-28T16:13:12.021177Z","shell.execute_reply":"2024-11-28T16:13:13.484207Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Q**: What about on cancer level? Are the distributions different?\n\n**A**: It looks like the minimum and average age have shifted. The *youngest patient to have cancer is 38 years old*, while the *mean of those patients is 63* years old.","metadata":{}},{"cell_type":"code","source":"# Plot\nf, (a0, a1) = plt.subplots(1, 2, figsize=(24, 12))\nsns.distplot(a=train[train[\"cancer\"]==0][\"age\"], rug=True, hist=False, \n             rug_kws={\"color\": my_colors[5]},\n             kde_kws={\"color\": my_colors[5], \"lw\": 5, \"alpha\": 0.7},\n             ax=a0)\na0.set_title(\"No Cancer Present\", weight=\"bold\", size=20)\na0.axvline(x=58, ls=\":\", lw=2, color=\"black\")\na0.text(x=58.5, y=0.018, s=\"mean: 58\", size=17, color=\"black\", weight=\"bold\")\na0.axvline(x=26, ls=\":\", lw=2, color=\"black\")\na0.text(x=26.5, y=0.008, s=\"min: 26\", size=17, color=\"black\", weight=\"bold\")\na0.axvline(x=89, ls=\":\", lw=2, color=\"black\")\na0.text(x=79, y=0.037, s=\"max: 89\", size=17, color=\"black\", weight=\"bold\")\n\n\nsns.distplot(a=train[train[\"cancer\"]==1][\"age\"], rug=True, hist=False, \n             rug_kws={\"color\": my_colors[2]},\n             kde_kws={\"color\": my_colors[2], \"lw\": 5, \"alpha\": 0.7},\n             ax=a1)\na1.set_title(\"Cancer Present\", weight=\"bold\", size=20)\na1.axvline(x=63, ls=\":\", lw=2, color=\"black\")\na1.text(x=63.5, y=0.018, s=\"mean: 63\", size=17, color=\"black\", weight=\"bold\")\na1.axvline(x=38, ls=\":\", lw=2, color=\"black\")\na1.text(x=38.5, y=0.008, s=\"min: 38\", size=17, color=\"black\", weight=\"bold\")\na1.axvline(x=89, ls=\":\", lw=2, color=\"black\")\na1.text(x=79, y=0.037, s=\"max: 89\", size=17, color=\"black\", weight=\"bold\")\n\n\nplt.suptitle(\"Age Distribution\", weight=\"bold\", size=25)\nsns.despine(right=True, top=True, left=True);","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:13:30.763907Z","iopub.execute_input":"2024-11-28T16:13:30.764287Z","iopub.status.idle":"2024-11-28T16:13:31.880373Z","shell.execute_reply.started":"2024-11-28T16:13:30.764256Z","shell.execute_reply":"2024-11-28T16:13:31.879461Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Q**: Are there many cases with implants? How do these implants show in the scans?\n\n**A**: Looks like **there are very few records with implants**. Some scans clearly show the implant, however in others it is not visible (at least to my unexpert eye):\n* `implant` - Flag whether or not the patient has breast implants. For `site`=1 we have this info on patient level, for `site`=2 we have this info at image level. There are **only 1477 images** that contain implants.","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport pydicom\n\ndef show_images(col, col_flag, sample_size, cancer_flag=0):\n    # Get image info\n    data = train[train[col] == col_flag].sample(sample_size, random_state=24)\n    if cancer_flag == 1:\n        data = train[train.cancer == 1]\n        data = data[data[col] == col_flag].sample(sample_size, random_state=24)\n    image_path = data[\"path\"].to_list()\n\n    # Plot\n    fig, axs = plt.subplots(1, sample_size, figsize=(23, 4))\n    axs = axs.flatten()\n\n    for k, path in enumerate(image_path):\n        axs[k].set_title(f\"{k+1}. {col_flag}\", \n                         fontsize=14, color=my_colors[0], weight='bold')\n\n        img = pydicom.dcmread(path).pixel_array\n        axs[k].imshow(img, cmap=\"turbo\")\n        axs[k].axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n","metadata":{"_kg_hide-input":true,"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:14:00.729127Z","iopub.execute_input":"2024-11-28T16:14:00.729871Z","iopub.status.idle":"2024-11-28T16:14:00.736991Z","shell.execute_reply.started":"2024-11-28T16:14:00.729841Z","shell.execute_reply":"2024-11-28T16:14:00.73602Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(clr.S+\"Records with no implants:\"+clr.E, train[\"implant\"].value_counts().values[0], \"\\n\"+\n      clr.S+\"Records with implants:\"+clr.E, train[\"implant\"].value_counts().values[1], \"\\n\")\n\nfor implant_flag in train[\"implant\"].unique().tolist():\n    show_images(col=\"implant\", col_flag=implant_flag, sample_size=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:14:02.467335Z","iopub.execute_input":"2024-11-28T16:14:02.468085Z","iopub.status.idle":"2024-11-28T16:14:17.026758Z","shell.execute_reply.started":"2024-11-28T16:14:02.46805Z","shell.execute_reply":"2024-11-28T16:14:17.025907Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Target Explore\n\n**Q**: How many cases with cancer are there? What can the variables given only for `train` dataset tell us?\n\n* `cancer` - our **target** variable. 1 if the image has cancer detected, 0 otherwise. The class imbalance is very big, with **only 2% of images being cancer positive**.\n* `invasive` - for cases when cancer was present, another flag that states if the cancer was invasive or not. Around 70% of cases that had cancer were invasive.","metadata":{}},{"cell_type":"code","source":"print(clr.S+\"Records with no cancer:\"+clr.E, train[\"cancer\"].value_counts().values[0], \"\\n\"+\n      clr.S+\"Records with cancer:\"+clr.E, train[\"cancer\"].value_counts().values[1], \"\\n\")\n\nfor cancer_flag in train[\"cancer\"].unique().tolist():\n    show_images(col=\"cancer\", col_flag=cancer_flag, sample_size=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:14:48.940141Z","iopub.execute_input":"2024-11-28T16:14:48.940791Z","iopub.status.idle":"2024-11-28T16:15:03.172239Z","shell.execute_reply.started":"2024-11-28T16:14:48.94076Z","shell.execute_reply":"2024-11-28T16:15:03.171355Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(clr.S+\"Records with invasion:\"+clr.E, train[train.cancer==1][\"invasive\"].value_counts().values[0], \"\\n\"+\n      clr.S+\"Records with no invasion:\"+clr.E, train[train.cancer==1][\"invasive\"].value_counts().values[1], \"\\n\")\n\nfor invasive_flag in train[\"invasive\"].unique().tolist():\n\n    show_images(col=\"invasive\", col_flag=invasive_flag, sample_size=5, cancer_flag=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:15:03.173944Z","iopub.execute_input":"2024-11-28T16:15:03.174283Z","iopub.status.idle":"2024-11-28T16:15:15.939226Z","shell.execute_reply.started":"2024-11-28T16:15:03.174249Z","shell.execute_reply":"2024-11-28T16:15:15.938302Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Q**: How many cases were difficult? Are the scans different between cases?\n\n**A**: Around 16% of cases proved to be difficult:\n* `difficult_negative_case` - True if the speciffic case was particularly difficult. There have been around 7705 cases that were exceptionally difficult in the train dataset. This variable is not provided for the test set.","metadata":{}},{"cell_type":"code","source":"print(clr.S+\"Cases not particularly difficult:\"+clr.E, train[\"difficult_negative_case\"].value_counts().values[0], \"\\n\"+\n      clr.S+\"Cases particularly difficult:\"+clr.E, train[\"difficult_negative_case\"].value_counts().values[1], \"\\n\")\n\nfor difficult_flag in train[\"difficult_negative_case\"].unique().tolist():\n    show_images(col=\"difficult_negative_case\", col_flag=difficult_flag, sample_size=5, cancer_flag=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:15:15.940537Z","iopub.execute_input":"2024-11-28T16:15:15.940831Z","iopub.status.idle":"2024-11-28T16:15:28.911656Z","shell.execute_reply.started":"2024-11-28T16:15:15.940805Z","shell.execute_reply":"2024-11-28T16:15:28.910838Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Train Preprocessing\n\nJust a little data preparation before going into the models.\n\n## 3.1 Label Encoding\n* `laterality`:\n    * 0: L (Left)\n    * 1: R (Right)\n* `view`:\n    * 0: AT (axillary tail)\n    * 1: CC (craniocaudal)\n    * 2: LM (latero-medial)\n    * 3: LMO (latero-medial oblique) \n    * 4: ML (medio-lateral)\n    * 5: MLO (mediolateral oblique)","metadata":{}},{"cell_type":"code","source":"# Keep only columns in test + target variable\ntrain = train[[\"patient_id\", \"image_id\", \"laterality\", \"view\", \"age\", \"implant\", \"path\", \"cancer\"]]\n\n# Encode categorical variables\nle_laterality = LabelEncoder()\nle_view = LabelEncoder()\n\ntrain['laterality'] = le_laterality.fit_transform(train['laterality'])\ntrain['view'] = le_view.fit_transform(train['view'])\n\ntrain.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:15:28.913164Z","iopub.execute_input":"2024-11-28T16:15:28.913454Z","iopub.status.idle":"2024-11-28T16:15:28.953123Z","shell.execute_reply.started":"2024-11-28T16:15:28.913427Z","shell.execute_reply":"2024-11-28T16:15:28.952313Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3.2 Impute\n\n**!!! Note**: `age` column has 37 missing values (duplicated for some patients). We will impute it with the mean and then *normalize*.","metadata":{}},{"cell_type":"code","source":"print(clr.S+\"Number of missing values in Age:\"+clr.E, train[\"age\"].isna().sum())\ntrain['age'] = train['age'].fillna(58)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:15:34.432075Z","iopub.execute_input":"2024-11-28T16:15:34.432433Z","iopub.status.idle":"2024-11-28T16:15:34.439922Z","shell.execute_reply.started":"2024-11-28T16:15:34.432404Z","shell.execute_reply":"2024-11-28T16:15:34.438923Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3.3 Save Artifact","metadata":{}},{"cell_type":"markdown","source":"---------------\n<img src=\"https://i.imgur.com/4MyF4Zh.jpg\">\n\n<center><h2> - PyTorch Baseline - </h2></center>\n\n### ○ (more) Libraries","metadata":{}},{"cell_type":"code","source":"!pip install -q efficientnet_pytorch","metadata":{"_kg_hide-output":true,"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:16:00.131498Z","iopub.execute_input":"2024-11-28T16:16:00.132383Z","iopub.status.idle":"2024-11-28T16:16:09.590151Z","shell.execute_reply.started":"2024-11-28T16:16:00.132348Z","shell.execute_reply":"2024-11-28T16:16:09.588997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# PyTorch\nimport torch\nimport torchvision\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch import FloatTensor, LongTensor\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\n# Data Augmentation for Image Preprocessing\nfrom albumentations import (ToFloat, Normalize, VerticalFlip, HorizontalFlip, Compose, Resize,\n                            RandomBrightnessContrast, HueSaturationValue, Blur, GaussNoise,\n                            Rotate, RandomResizedCrop, Cutout, ShiftScaleRotate, ToGray)\nfrom albumentations.pytorch import ToTensorV2\n\nfrom efficientnet_pytorch import EfficientNet\nfrom torchvision.models import resnet34, resnet50\n\n# SKlearn\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold\nfrom sklearn.metrics import accuracy_score, roc_auc_score, confusion_matrix","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:16:09.592276Z","iopub.execute_input":"2024-11-28T16:16:09.592612Z","iopub.status.idle":"2024-11-28T16:16:10.694236Z","shell.execute_reply.started":"2024-11-28T16:16:09.592562Z","shell.execute_reply":"2024-11-28T16:16:10.693478Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Seed\nset_seed()\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint('Device available now:', DEVICE)\n\n# Read in Data\ntrain = pd.read_csv(\"/kaggle/input/rsna-breast-cancer-helper-data/train_path.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:16:10.695291Z","iopub.execute_input":"2024-11-28T16:16:10.695652Z","iopub.status.idle":"2024-11-28T16:16:10.891981Z","shell.execute_reply.started":"2024-11-28T16:16:10.695618Z","shell.execute_reply":"2024-11-28T16:16:10.891007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(train))\ntrain1= train[train[\"cancer\"] !=0]\ntrain1 = train1.reset_index(drop=True)\nprint(len(train1))\ntrain0= train[train[\"cancer\"] ==0][:len(train1)]\ntrain0 = train0.reset_index(drop=True)\nframes = [train0, train1]\ntrain = pd.concat(frames)\ntrain = train.reset_index(drop=True)\nprint(len(train))\n# Sample down for dev\n#train = train.sample(n=500, random_state=13).reset_index(drop=True)\n#train[\"cancer\"].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:16:20.450028Z","iopub.execute_input":"2024-11-28T16:16:20.450771Z","iopub.status.idle":"2024-11-28T16:16:20.465671Z","shell.execute_reply.started":"2024-11-28T16:16:20.450738Z","shell.execute_reply":"2024-11-28T16:16:20.464626Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Sample down for dev\n#train = train.sample(n=500, random_state=13).reset_index(drop=True)\ntrain[\"cancer\"].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:16:27.089774Z","iopub.execute_input":"2024-11-28T16:16:27.090397Z","iopub.status.idle":"2024-11-28T16:16:27.097372Z","shell.execute_reply.started":"2024-11-28T16:16:27.090358Z","shell.execute_reply":"2024-11-28T16:16:27.09645Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### ○ Params","metadata":{}},{"cell_type":"code","source":"# ----- GLOBAL PARAMS -----\nvertical_flip = 0.5\nhorizontal_flip = 0.5\n\ncsv_columns = ['laterality', 'view', 'age', 'implant']\nno_columns = len(csv_columns)\noutput_size = 1\n# -------------------------","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:16:30.075293Z","iopub.execute_input":"2024-11-28T16:16:30.075684Z","iopub.status.idle":"2024-11-28T16:16:30.080675Z","shell.execute_reply.started":"2024-11-28T16:16:30.075652Z","shell.execute_reply":"2024-11-28T16:16:30.079538Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. PyTorch Dataset","metadata":{}},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    \n    def __init__(self, dataframe, vertical_flip, horizontal_flip,\n                 is_train=True):\n        self.dataframe, self.is_train = dataframe, is_train\n        self.vertical_flip, self.horizontal_flip = vertical_flip, horizontal_flip\n        \n        # Data Augmentation (custom for each dataset type)\n        if is_train:\n            self.transform = Compose([RandomResizedCrop(height=224, width=224),\n                                      ShiftScaleRotate(rotate_limit=90, scale_limit = [0.8, 1.2]),\n                                      HorizontalFlip(p = self.horizontal_flip),\n                                      VerticalFlip(p = self.vertical_flip),\n                                      ToTensorV2()])\n        else:\n            self.transform = Compose([ToTensorV2()])\n            \n    def __len__(self):\n        return len(self.dataframe)\n    \n    def __getitem__(self, index):\n        '''Take each row in batcj at a time.'''\n        \n        # Select path and read image\n        image_path = self.dataframe['path'][index]\n        image = pydicom.dcmread(image_path).pixel_array.astype(np.float32)\n        \n        # For this image also import .csv information\n        csv_data = np.array(self.dataframe.iloc[index][csv_columns].values, \n                            dtype=np.float32)\n        # Apply transforms\n        transf_image = self.transform(image=image)['image']\n        # Change image from 1 channel (B&W) to 3 channels\n        transf_image = np.concatenate([transf_image, transf_image, transf_image], axis=0)\n        \n        # Return info\n        if self.is_train:\n            return {\"image\": transf_image, \n                    \"meta\": csv_data, \n                    \"target\": self.dataframe['cancer'][index]}\n        else:\n            return {\"image\": transf_image, \n                    \"meta\": csv_data}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:16:32.468877Z","iopub.execute_input":"2024-11-28T16:16:32.469845Z","iopub.status.idle":"2024-11-28T16:16:32.482836Z","shell.execute_reply.started":"2024-11-28T16:16:32.469799Z","shell.execute_reply":"2024-11-28T16:16:32.481768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def data_to_device(data):\n    \n    image, metadata, targets = data.values()\n    return image.to(DEVICE), metadata.to(DEVICE), targets.to(DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:16:54.133198Z","iopub.execute_input":"2024-11-28T16:16:54.133962Z","iopub.status.idle":"2024-11-28T16:16:54.138462Z","shell.execute_reply.started":"2024-11-28T16:16:54.13393Z","shell.execute_reply":"2024-11-28T16:16:54.137539Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### ◾ Sanity Check","metadata":{}},{"cell_type":"code","source":"# Sample data\nsample_df = train.head(6)\n\n# Instantiate Dataset object\ndataset = RSNADataset(sample_df, vertical_flip, horizontal_flip,\n                      is_train=True)\n# The Dataloader\ndataloader = DataLoader(dataset, batch_size=3, shuffle=False)\n\n# Output of the Dataloader\nfor k, data in enumerate(dataloader):\n    image, meta, targets = data_to_device(data)\n    print(clr.S + f\"Batch: {k}\" + clr.E, \"\\n\" +\n          clr.S + \"Image:\" + clr.E, image.shape, \"\\n\" +\n          clr.S + \"Meta:\" + clr.E, meta, \"\\n\" +\n          clr.S + \"Targets:\" + clr.E, targets, \"\\n\" +\n          \"=\"*50)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:16:56.73644Z","iopub.execute_input":"2024-11-28T16:16:56.737123Z","iopub.status.idle":"2024-11-28T16:17:04.503479Z","shell.execute_reply.started":"2024-11-28T16:16:56.737089Z","shell.execute_reply":"2024-11-28T16:17:04.502475Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Network\n\n## 5.1 ResNet50\n\n* **ResNet** - stands for Residual Network.\n    * Residual Networks - form networks **by stacking residual blocks**.\n    * [Residual Block](https://towardsdatascience.com/resnets-residual-blocks-deep-residual-learning-a231a0ee73d2#:~:text=A%20residual%20block%20is%20a,layer%20in%20the%20main%20path.) - a stack of layers set in such a way that the output of a layer is taken and added to another layer deeper in the block. This method is also called **by-pass** or **skip-connection**.\n<center><img src=\"https://miro.medium.com/max/828/0*TyvNvZhiuzEbIiXw.webp\" width=700></center>\n* It's a *special* CNN introduced in 2015\n* Many ResNets - ResNet18, ResNet34, ResNet50, ResNet101, ResNet152 (numbers stand for how many layers)\n* **Deeper layers** (more convolutions), but without the [vanishing gradient problem](https://www.kdnuggets.com/2022/02/vanishing-gradient-problem.html).\n* The **residual blocks create an identity mapping** to activations *earlier* in the network to thwart the performance degradation problem associated with deep neural architectures.\n* The **by-pass connections help to address the problem of vanishing and exploding gradients**.\n\n<center><img src=\"https://i.imgur.com/YXA2OVM.jpg\" width=800></center>","metadata":{}},{"cell_type":"code","source":"class ResNet50Network(nn.Module):\n    def __init__(self, output_size, no_columns):\n        super().__init__()\n        self.no_columns, self.output_size = no_columns, output_size\n        \n        # Define Feature part (IMAGE)\n        self.features = resnet50(pretrained=True) # 1000 neurons out\n        # (metadata)\n        self.csv = nn.Sequential(nn.Linear(self.no_columns, 500),\n                                 nn.BatchNorm1d(500),\n                                 nn.ReLU(),\n                                 nn.Dropout(p=0.2))\n        \n        # Define Classification part\n        self.classification = nn.Linear(1000 + 500, output_size)\n        \n        \n    def forward(self, image, meta, prints=False):\n        if prints: print('Input Image shape:', image.shape, '\\n'+\n                         'Input metadata shape:', meta.shape)\n        \n        # Image CNN\n        image = self.features(image)\n        if prints: print('Features Image shape:', image.shape)\n        \n        # CSV FNN\n        meta = self.csv(meta)\n        if prints: print('Meta Data:', meta.shape)\n            \n        # Concatenate layers from image with layers from csv_data\n        image_meta_data = torch.cat((image, meta), dim=1)\n        if prints: print('Concatenated Data:', image_meta_data.shape)\n        \n        # CLASSIF\n        out = self.classification(image_meta_data)\n        if prints: print('Out shape:', out.shape)\n        \n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:39:15.125566Z","iopub.execute_input":"2024-11-28T16:39:15.126733Z","iopub.status.idle":"2024-11-28T16:39:15.135466Z","shell.execute_reply.started":"2024-11-28T16:39:15.126697Z","shell.execute_reply":"2024-11-28T16:39:15.134472Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### ◾ Sanity Check\n\n> How does it work:\n<center><img src=\"https://i.imgur.com/ChkZ53p.jpg\" width=700></center>","metadata":{}},{"cell_type":"code","source":"# Load Model\nmodel_example = ResNet50Network(output_size=output_size, no_columns=no_columns).to(DEVICE)\n\n# Outputs\nout = model_example(image, meta, prints=True)\n\n# Criterion example\ncriterion_example = nn.BCEWithLogitsLoss()\n# Unsqueeze(1) from shape=[3] to shape=[3, 1]\nloss = criterion_example(out, targets.unsqueeze(1).float()) \nprint(\"=\"*50)\nprint(clr.S+'Loss:'+clr.E, loss.item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:17:49.038069Z","iopub.execute_input":"2024-11-28T16:17:49.038401Z","iopub.status.idle":"2024-11-28T16:17:56.238078Z","shell.execute_reply.started":"2024-11-28T16:17:49.038375Z","shell.execute_reply":"2024-11-28T16:17:56.23699Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. Training ...\n","metadata":{}},{"cell_type":"code","source":"1408+250","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:18:30.686756Z","iopub.execute_input":"2024-11-28T16:18:30.687705Z","iopub.status.idle":"2024-11-28T16:18:30.69311Z","shell.execute_reply.started":"2024-11-28T16:18:30.68765Z","shell.execute_reply":"2024-11-28T16:18:30.692291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def add_in_file(text, f):\n    \n    with open(f'logs_{VERSION}.txt', 'a+') as f:\n        print(text, file=f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:18:32.119741Z","iopub.execute_input":"2024-11-28T16:18:32.120488Z","iopub.status.idle":"2024-11-28T16:18:32.124652Z","shell.execute_reply.started":"2024-11-28T16:18:32.120457Z","shell.execute_reply":"2024-11-28T16:18:32.123728Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 🎗️ Main Training Function\n\n> How does it work (W&B is removed from this notebook) :\n<center><img src=\"https://i.imgur.com/7MWwg9W.png\" width=800></center>","metadata":{}},{"cell_type":"code","source":"def train_folds(model, train_original, is_training=False):\n    if not is_training:\n        model_path = \"/kaggle/input/resnet50/pytorch/default/1/state_dict_model.pt\"\n        model.load_state_dict(torch.load(model_path))\n        return model\n\n    f = open(f\"logs_{VERSION}.txt\", \"w+\")\n\n    # Split in folds\n    group_fold = GroupKFold(n_splits=FOLDS)\n\n    # Generate indices to split data into training and test set.\n    k_folds = group_fold.split(\n        X=np.zeros(len(train_original)),\n        y=train_original['cancer'],\n        groups=train_original['patient_id'].tolist()\n    )\n\n    # For each fold\n    for i, (train_index, valid_index) in enumerate(k_folds):\n        print(clr.S + f\"---------- Fold: {i+1} ----------\" + clr.E)\n        add_in_file(f\"---------- Fold: {i+1} ----------\", f)\n\n        # --- Create Instances ---\n        # Best ROC score in this fold\n        best_roc = None\n        # Reset patience before every fold\n        patience_f = PATIENCE\n\n        # Optimizer/ Scheduler/ Criterion\n        optimizer = torch.optim.Adam(\n            model.parameters(),\n            lr=LR,\n            weight_decay=WD\n        )\n        scheduler = ReduceLROnPlateau(\n            optimizer=optimizer,\n            mode='max',\n            patience=LR_PATIENCE,\n            verbose=True,\n            factor=LR_FACTOR\n        )\n        criterion = nn.BCEWithLogitsLoss()\n\n        # --- Read in Data ---\n        train_data = train_original.iloc[train_index].reset_index(drop=True)\n        valid_data = train_original.iloc[valid_index].reset_index(drop=True)\n\n        # Create Data instances\n        train = RSNADataset(\n            train_data, vertical_flip, horizontal_flip, is_train=True)\n        valid = RSNADataset(\n            valid_data, vertical_flip, horizontal_flip, is_train=True)\n\n        # Dataloaders\n        train_loader = DataLoader(\n            train, batch_size=BATCH_SIZE1, shuffle=True, num_workers=WORKERS)\n        valid_loader = DataLoader(\n            valid, batch_size=BATCH_SIZE2, shuffle=False, num_workers=WORKERS)\n\n        # === EPOCHS ===\n        for epoch in range(EPOCHS):\n            start_time = time()\n            correct = 0\n            train_losses = 0\n\n            # === TRAIN ===\n            # Sets the module in training mode.\n            model.train()\n\n            # For each batch\n            for k, data in tqdm(enumerate(train_loader)):\n                # Save them to device\n                image, meta, targets = data_to_device(data)\n\n                # Clear gradients first; very important\n                # usually done BEFORE prediction\n                optimizer.zero_grad()\n\n                # Log Probabilities & Backpropagation\n                out = model(image, meta)\n                loss = criterion(out, targets.unsqueeze(1).float())\n                loss.backward()\n                optimizer.step()\n\n                # --- Save information after this batch ---\n                # Save loss\n                train_losses += loss.item()\n                # From log probabilities to actual probabilities\n                train_preds = torch.round(torch.sigmoid(out))  # 0 and 1\n                # Number of correct predictions\n                correct += (train_preds.cpu() ==\n                            targets.cpu().unsqueeze(1)).sum().item()\n\n            # Compute Train Accuracy\n            train_acc = correct / len(train_index)\n\n            # === EVAL ===\n            # Sets the model in evaluation mode.\n            model.eval()\n\n            # Create matrix to store evaluation predictions (for accuracy)\n            valid_preds = torch.zeros(\n                size=(len(valid_index), 1),\n                device=DEVICE,\n                dtype=torch.float32\n            )\n\n            # Disables gradients (we need to be sure no optimization happens)\n            with torch.no_grad():\n                for k, data in tqdm(enumerate(valid_loader)):\n                    # Save them to device\n                    image, meta, targets = data_to_device(data)\n\n                    out = model(image, meta)\n                    pred = torch.sigmoid(out)\n                    valid_preds[k*image.shape[0]:k*image.shape[0] +\n                                image.shape[0]] = pred\n\n                # Calculate accuracy\n                valid_acc = accuracy_score(\n                    valid_data['cancer'].values,\n                    torch.round(valid_preds.cpu())\n                )\n\n                # Calculate ROC\n                valid_roc = roc_auc_score(\n                    valid_data['cancer'].values,\n                    valid_preds.cpu()\n                )\n\n                # Calculate time on Train + Eval\n                duration = str(\n                    dtime.timedelta(seconds=time() - start_time))[:7]\n\n                # PRINT INFO\n                final_logs = '{} | Epoch: {}/{} | Loss: {:.4} | Acc_tr: {:.3} | Acc_vd: {:.3} | ROC: {:.3}'.format(\n                    duration, epoch+1, EPOCHS,\n                    train_losses, train_acc, valid_acc, valid_roc)\n                add_in_file(final_logs, f)\n                print(final_logs)\n\n                # === SAVE MODEL ===\n\n                # Update scheduler (for learning_rate)\n                scheduler.step(valid_roc)\n                # Name the model\n                model_name = f\"Fold{i+1}_Epoch{epoch+1}_ValidAcc{valid_acc:.3f}_ROC_{valid_roc:.3f}.pth\"\n\n                # Update best_roc\n                if not best_roc:  # If best_roc = None\n                    best_roc = valid_roc\n                    torch.save(model.state_dict(), model_name)\n                    continue\n\n                if valid_roc > best_roc:\n                    best_roc = valid_roc\n                    # Reset patience (because we have improvement)\n                    patience_f = PATIENCE\n                    torch.save(model.state_dict(), model_name)\n                else:\n                    # Decrease patience (no improvement in ROC)\n                    patience_f = patience_f - 1\n                    if patience_f == 0:\n                        stop_logs = 'Early stopping (no improvement since 3 models) | Best ROC: {}'.format(\n                            best_roc)\n                        add_in_file(stop_logs, f)\n                        print(stop_logs)\n                        break\n\n        # === CLEANING ===\n        # Clear memory\n        del train, valid, train_loader, valid_loader, image, targets\n        gc.collect()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T14:37:42.20651Z","iopub.execute_input":"2024-11-28T14:37:42.206893Z","iopub.status.idle":"2024-11-28T14:37:42.227793Z","shell.execute_reply.started":"2024-11-28T14:37:42.206862Z","shell.execute_reply":"2024-11-28T14:37:42.226573Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 🎗️ Experimenting\n\n*TODO: locally run and save logs here*\n\n> 📌 **Note**: This cell runs on only a small fraction of the data. I will add the logs for the full run in a later version.","metadata":{}},{"cell_type":"code","source":"FOLDS = 3\nEPOCHS = 5\nPATIENCE = 3\nWORKERS = 8\nLR = 0.0005\nWD = 0.0\nLR_PATIENCE = 1            # 1 model not improving until lr is decreasing\nLR_FACTOR = 0.4            # by how much the lr is decreasing\n\nBATCH_SIZE1 = 32           # for train\nBATCH_SIZE2 = 16           # for valid\n\nVERSION = 'v0'\nMODEL = 'resnet50'\nmodel1 = ResNet50Network(output_size=output_size, no_columns=no_columns).to(DEVICE)\n#MODEL = 'EffNetNetwork'\n#model1 = EffNetNetwork(output_size=output_size, no_columns=no_columns).to(DEVICE)\n\n\n# ------------------\n\ntrained_model = train_folds(model=model1, train_original=train)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T14:38:56.768453Z","iopub.execute_input":"2024-11-28T14:38:56.769346Z","iopub.status.idle":"2024-11-28T14:38:57.360276Z","shell.execute_reply.started":"2024-11-28T14:38:56.769311Z","shell.execute_reply":"2024-11-28T14:38:57.359521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# WIP\n# Save\ntorch.save(model1.state_dict(), \"state_dict_model.pt\")\n# Save\ntorch.save(model1, \"entire_model.pt\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 7. Inference","metadata":{}},{"cell_type":"code","source":"import torch\nimport pydicom\nimport numpy as np\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision.models import resnet50\nimport torch.nn as nn\n\n# Device configuration\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nMODEL_PATH = \"/kaggle/input/resnet50/pytorch/default/1/state_dict_model.pt\"\n\n\n\n# Dataset Class\nclass InferenceDataset(Dataset):\n    def __init__(self, dataframe, transform):\n        self.dataframe = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        # Read the DICOM file\n        image_path = self.dataframe.iloc[idx][\"path\"]\n        dicom = pydicom.dcmread(image_path)\n\n        # Extract the pixel array and metadata\n        image = dicom.pixel_array.astype(np.float32)\n\n        # Extract metadata\n        laterality = 0 if dicom.get(\"ImageLaterality\") == \"L\" else 1 \n        view = dicom.get(\"ViewPosition\", \"\") \n        age = int(dicom.get(\"PatientAge\", \"058\")[:3])\n        implant = int(dicom.get(\"ImplantPresent\", 0)) \n\n        # Encode view if necessary\n        view_mapping = {\"CC\": 0, \"MLO\": 1, \"AT\": 2, \"LM\": 3, \"LMO\": 4, \"ML\": 5}\n        view = view_mapping.get(view, -1)  # Default to -1 if the view is not recognized\n\n        # Duplicate channels to make the image 3-channel\n        image = np.stack([image, image, image], axis=-1)  # Shape: (H, W, 3)\n\n        # Apply transformations\n        transformed = self.transform(image=image)\n        transformed_image = transformed[\"image\"]\n\n        # Prepare metadata as a numpy array\n        metadata = np.array([laterality, view, age, implant], dtype=np.float32)\n\n        return {\"image\": transformed_image, \"meta\": metadata}\n\n\n# Load the Model\ndef load_model(model, checkpoint_path):\n    model.load_state_dict(torch.load(checkpoint_path, map_location=DEVICE))\n    print(\"Model loaded successfully.\")\n    return model\n\n# Inference Function\ndef infer(model, dataloader):\n    model.eval()\n    predictions = []\n\n    with torch.no_grad():\n        for batch in dataloader:\n            images = batch[\"image\"].to(DEVICE)\n            meta = batch[\"meta\"].to(DEVICE)\n            outputs = model(images, meta)\n            preds = torch.sigmoid(outputs).cpu().numpy()\n            predictions.extend(preds)\n\n    return predictions\n\n# Main Inference Script\nif __name__ == \"__main__\":\n\n    test_data = pd.DataFrame({\n        \"path\": [\n            \"/kaggle/input/rsna-breast-cancer-detection/test_images/10008/1591370361.dcm\",\n            \"/kaggle/input/rsna-breast-cancer-detection/test_images/10008/361203119.dcm\",\n            # Add more test image paths\n        ]\n    })\n\n    # Metadata columns\n    csv_columns = [\"laterality\", \"view\", \"age\", \"implant\"]\n\n    # Define image transformations\n    transform = Compose([\n        Resize(height=224, width=224),  # Resize images to match ResNet50 input size\n        ToTensorV2()  # Convert to tensor\n    ])\n\n    # Load the dataset and dataloader\n    dataset = InferenceDataset(test_data, transform)\n    dataloader = DataLoader(dataset, batch_size=4, shuffle=False)\n\n    # Load the model\n    model = ResNet50Network(output_size=1, no_columns=len(csv_columns)).to(DEVICE)\n    model = load_model(model, MODEL_PATH)\n\n    # Perform inference\n    predictions = infer(model, dataloader)\n\n    # Print predictions\n    for idx, pred in enumerate(predictions):\n        print(f\"Prediction for image {test_data.iloc[idx]['path']}: {pred[0]:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-28T16:47:18.803518Z","iopub.execute_input":"2024-11-28T16:47:18.803909Z","iopub.status.idle":"2024-11-28T16:47:20.081259Z","shell.execute_reply.started":"2024-11-28T16:47:18.80388Z","shell.execute_reply":"2024-11-28T16:47:20.080254Z"}},"outputs":[],"execution_count":null}]}