{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<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\" pydicom pylibjpeg \"opencv-python-headless\"","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-01-08T08:01:25.55509Z","iopub.execute_input":"2023-01-08T08:01:25.555748Z","iopub.status.idle":"2023-01-08T08:01:42.617807Z","shell.execute_reply.started":"2023-01-08T08:01:25.555651Z","shell.execute_reply":"2023-01-08T08:01:42.617069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2023-01-08T08:01:42.61946Z","iopub.execute_input":"2023-01-08T08:01:42.619747Z","iopub.status.idle":"2023-01-08T08:01:44.59951Z","shell.execute_reply.started":"2023-01-08T08:01:42.619722Z","shell.execute_reply":"2023-01-08T08:01:44.59858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 🐝 W&B Fork & Run\n\nIn order to run this notebook you will need to input your own **secret API key** within the `! wandb login $secret_value_0` line. \n\n🐝**How do you get your own API key?**\n\nSuper simple! Go to **https://wandb.ai/site** -> Login -> Click on your profile in the top right corner -> Settings -> Scroll down to API keys -> copy your very own key (for more info check [this amazing notebook for ML Experiment Tracking on Kaggle](https://www.kaggle.com/ayuraj/experiment-tracking-with-weights-and-biases)).\n\n<center><img src=\"https://i.imgur.com/fFccmoS.png\" width=500></center>","metadata":{}},{"cell_type":"code","source":"# 🐝 Secrets\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"wandb\")\n\n! wandb login $secret_value_0","metadata":{"execution":{"iopub.status.busy":"2023-01-08T08:01:44.603511Z","iopub.execute_input":"2023-01-08T08:01:44.605506Z","iopub.status.idle":"2023-01-08T08:01:46.488602Z","shell.execute_reply.started":"2023-01-08T08:01:44.60547Z","shell.execute_reply":"2023-01-08T08:01:46.487532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### ○ Helpers","metadata":{}},{"cell_type":"code","source":"# === 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        \n        \n# === 🐝 W&B ===\ndef save_dataset_artifact(run_name, artifact_name, path, data_type=\"dataset\"):\n    '''Saves dataset to W&B Artifactory.\n    run_name: name of the experiment\n    artifact_name: under what name should the dataset be stored\n    path: path to the dataset'''\n    \n    run = wandb.init(project='Otto', \n                     name=run_name, \n                     config=CONFIG)\n    artifact = wandb.Artifact(name=artifact_name, \n                              type=data_type)\n    artifact.add_file(path)\n\n    wandb.log_artifact(artifact)\n    wandb.finish()\n    print(\"Artifact has been saved successfully.\")\n    \n    \ndef create_wandb_plot(x_data=None, y_data=None, x_name=None, y_name=None, title=None, log=None, plot=\"line\"):\n    '''Create and save lineplot/barplot in W&B Environment.\n    x_data & y_data: Pandas Series containing x & y data\n    x_name & y_name: strings containing axis names\n    title: title of the graph\n    log: string containing name of log'''\n    \n    data = [[label, val] for (label, val) in zip(x_data, y_data)]\n    table = wandb.Table(data=data, columns = [x_name, y_name])\n    \n    if plot == \"line\":\n        wandb.log({log : wandb.plot.line(table, x_name, y_name, title=title)})\n    elif plot == \"bar\":\n        wandb.log({log : wandb.plot.bar(table, x_name, y_name, title=title)})\n    elif plot == \"scatter\":\n        wandb.log({log : wandb.plot.scatter(table, x_name, y_name, title=title)})\n        \n        \ndef create_wandb_hist(x_data=None, x_name=None, title=None, log=None):\n    '''Create and save histogram in W&B Environment.\n    x_data: Pandas Series containing x values\n    x_name: strings containing axis name\n    title: title of the graph\n    log: string containing name of log'''\n    \n    data = [[x] for x in x_data]\n    table = wandb.Table(data=data, columns=[x_name])\n    wandb.log({log : wandb.plot.histogram(table, x_name, title=title)})","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2023-01-08T08:01:46.490746Z","iopub.execute_input":"2023-01-08T08:01:46.491028Z","iopub.status.idle":"2023-01-08T08:01:46.642255Z","shell.execute_reply.started":"2023-01-08T08:01:46.490987Z","shell.execute_reply":"2023-01-08T08:01:46.641067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 🐝 Bonus: Cover Photo\nrun = wandb.init(project='RSNA_Breast_Cancer', name='CoverPhoto', config=CONFIG)\ncover = plt.imread(\"/kaggle/input/rsna-breast-cancer-helper-data/DKn4ofz.png\")\nwandb.log({\"cover\": wandb.Image(cover)})\nwandb.finish()","metadata":{"execution":{"iopub.status.busy":"2023-01-08T08:01:46.64312Z","iopub.execute_input":"2023-01-08T08:01:46.643381Z","iopub.status.idle":"2023-01-08T08:02:32.689223Z","shell.execute_reply.started":"2023-01-08T08:01:46.643357Z","shell.execute_reply":"2023-01-08T08:02:32.688233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-08T08:02:32.690519Z","iopub.execute_input":"2023-01-08T08:02:32.691933Z","iopub.status.idle":"2023-01-08T08:02:40.753155Z","shell.execute_reply.started":"2023-01-08T08:02:32.691872Z","shell.execute_reply":"2023-01-08T08:02:40.752191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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+\"Number of TOTAL images:\"+clr.E,\n      len(glob(\"/kaggle/input/rsna-breast-cancer-detection/train_images/*/*\")))\nprint(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":{"execution":{"iopub.status.busy":"2023-01-08T08:03:55.415129Z","iopub.execute_input":"2023-01-08T08:03:55.415473Z","iopub.status.idle":"2023-01-08T08:04:35.700305Z","shell.execute_reply.started":"2023-01-08T08:03:55.415445Z","shell.execute_reply":"2023-01-08T08:04:35.69921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"# 🐝 New Experiment\nrun = wandb.init(project='RSNA_Breast_Cancer', name='view_sample', config=CONFIG)","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:12:33.612326Z","iopub.execute_input":"2023-01-06T09:12:33.61271Z","iopub.status.idle":"2023-01-06T09:12:41.350219Z","shell.execute_reply.started":"2023-01-06T09:12:33.612664Z","shell.execute_reply":"2023-01-06T09:12:41.349043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_view(view_name, sample_size):\n    \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        wandb_images = []\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            wandb_images.append(wandb.Image(img))\n            axs[k].imshow(img, cmap=\"turbo\")\n            axs[k].axis(\"off\")\n\n        plt.tight_layout()\n        plt.show()\n\n        # 🐝 Log Image to W&B\n        wandb.log({f\"{view_name}\": wandb_images})\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        wandb_images = []\n        img = pydicom.dcmread(path).pixel_array\n        wandb_images.append(wandb.Image(img))\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        \n        wandb.log({f\"LMO\": wandb_images})","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2023-01-06T09:12:41.355342Z","iopub.execute_input":"2023-01-06T09:12:41.358213Z","iopub.status.idle":"2023-01-06T09:12:42.274885Z","shell.execute_reply.started":"2023-01-06T09:12:41.358155Z","shell.execute_reply":"2023-01-06T09:12:42.273844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for view_name in train[\"view\"].unique().tolist():\n    # Custom function to prin images & log into 🐝W&B\n    show_view(view_name, sample_size=5)","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:12:42.282872Z","iopub.execute_input":"2023-01-06T09:12:42.285659Z","iopub.status.idle":"2023-01-06T09:13:43.403826Z","shell.execute_reply.started":"2023-01-06T09:12:42.28562Z","shell.execute_reply":"2023-01-06T09:13:43.402792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:13:43.407225Z","iopub.execute_input":"2023-01-06T09:13:43.407631Z","iopub.status.idle":"2023-01-06T09:13:51.854134Z","shell.execute_reply.started":"2023-01-06T09:13:43.407595Z","shell.execute_reply":"2023-01-06T09:13:51.853079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"# 🐝 New Experiment\nrun = wandb.init(project='RSNA_Breast_Cancer', name='age_hist', config=CONFIG)","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:13:51.857075Z","iopub.execute_input":"2023-01-06T09:13:51.857721Z","iopub.status.idle":"2023-01-06T09:13:59.63976Z","shell.execute_reply.started":"2023-01-06T09:13:51.857669Z","shell.execute_reply":"2023-01-06T09:13:59.638652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:13:59.644907Z","iopub.execute_input":"2023-01-06T09:13:59.648606Z","iopub.status.idle":"2023-01-06T09:14:03.019996Z","shell.execute_reply.started":"2023-01-06T09:13:59.648556Z","shell.execute_reply":"2023-01-06T09:14:03.018985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 🐝 Log into dashboard\ncreate_wandb_hist(x_data=train[\"age\"], \n                  x_name=\"Age\",\n                  title=\"Age Distribution\",\n                  log=\"age_hist\")","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:14:03.024992Z","iopub.execute_input":"2023-01-06T09:14:03.027727Z","iopub.status.idle":"2023-01-06T09:14:07.288915Z","shell.execute_reply.started":"2023-01-06T09:14:03.027685Z","shell.execute_reply":"2023-01-06T09:14:07.288021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:14:07.29365Z","iopub.execute_input":"2023-01-06T09:14:07.295839Z","iopub.status.idle":"2023-01-06T09:14:10.977919Z","shell.execute_reply.started":"2023-01-06T09:14:07.295802Z","shell.execute_reply":"2023-01-06T09:14:10.97662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 🐝 Log into dashboard\ncreate_wandb_hist(x_data=train[train[\"cancer\"]==1][\"age\"], \n                  x_name=\"Age\",\n                  title=\"Age Distribution - patients with cancer\",\n                  log=\"age_hist_cancer\")","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:14:10.979392Z","iopub.execute_input":"2023-01-06T09:14:10.980434Z","iopub.status.idle":"2023-01-06T09:14:12.711703Z","shell.execute_reply.started":"2023-01-06T09:14:10.980395Z","shell.execute_reply":"2023-01-06T09:14:12.710749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:14:12.716542Z","iopub.execute_input":"2023-01-06T09:14:12.718837Z","iopub.status.idle":"2023-01-06T09:14:49.801075Z","shell.execute_reply.started":"2023-01-06T09:14:12.718799Z","shell.execute_reply":"2023-01-06T09:14:49.800084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"# 🐝 New Experiment\nrun = wandb.init(project='RSNA_Breast_Cancer', name='implant_sample', config=CONFIG)","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:14:49.802675Z","iopub.execute_input":"2023-01-06T09:14:49.803175Z","iopub.status.idle":"2023-01-06T09:14:57.594306Z","shell.execute_reply.started":"2023-01-06T09:14:49.803139Z","shell.execute_reply":"2023-01-06T09:14:57.593242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_images(col, col_flag, sample_size, cancer_flag=0):\n    \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    wandb_images = []\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        wandb_images.append(wandb.Image(img))\n        axs[k].imshow(img, cmap=\"turbo\")\n        axs[k].axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n\n    # 🐝 Log Image to W&B\n    wandb.log({f\"{col}_{col_flag}\": wandb_images})","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-01-06T09:14:57.599516Z","iopub.execute_input":"2023-01-06T09:14:57.602435Z","iopub.status.idle":"2023-01-06T09:14:59.120251Z","shell.execute_reply.started":"2023-01-06T09:14:57.602386Z","shell.execute_reply":"2023-01-06T09:14:59.119225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    # Custom function to prin images & log into 🐝W&B\n    show_images(col=\"implant\", col_flag=implant_flag, sample_size=5)","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:14:59.125226Z","iopub.execute_input":"2023-01-06T09:14:59.127996Z","iopub.status.idle":"2023-01-06T09:15:22.272869Z","shell.execute_reply.started":"2023-01-06T09:14:59.127938Z","shell.execute_reply":"2023-01-06T09:15:22.271939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:15:22.277702Z","iopub.execute_input":"2023-01-06T09:15:22.279984Z","iopub.status.idle":"2023-01-06T09:15:43.767222Z","shell.execute_reply.started":"2023-01-06T09:15:22.27993Z","shell.execute_reply":"2023-01-06T09:15:43.7662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"# 🐝 New Experiment\nrun = wandb.init(project='RSNA_Breast_Cancer', name='cancer_explore', config=CONFIG)","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:15:43.768651Z","iopub.execute_input":"2023-01-06T09:15:43.769037Z","iopub.status.idle":"2023-01-06T09:15:52.192192Z","shell.execute_reply.started":"2023-01-06T09:15:43.768996Z","shell.execute_reply":"2023-01-06T09:15:52.191113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    # Custom function to prin images & log into 🐝W&B\n    show_images(col=\"cancer\", col_flag=cancer_flag, sample_size=5)","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:15:52.197516Z","iopub.execute_input":"2023-01-06T09:15:52.200042Z","iopub.status.idle":"2023-01-06T09:16:14.760585Z","shell.execute_reply.started":"2023-01-06T09:15:52.199976Z","shell.execute_reply":"2023-01-06T09:16:14.759642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    # Custom function to prin images & log into 🐝W&B\n    show_images(col=\"invasive\", col_flag=invasive_flag, sample_size=5, cancer_flag=1)","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:16:14.76524Z","iopub.execute_input":"2023-01-06T09:16:14.767855Z","iopub.status.idle":"2023-01-06T09:16:35.6831Z","shell.execute_reply.started":"2023-01-06T09:16:14.767815Z","shell.execute_reply":"2023-01-06T09:16:35.682189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    # Custom function to prin images & log into 🐝W&B\n    show_images(col=\"difficult_negative_case\", col_flag=difficult_flag, sample_size=5, cancer_flag=0)","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:16:35.684607Z","iopub.execute_input":"2023-01-06T09:16:35.684982Z","iopub.status.idle":"2023-01-06T09:16:52.7713Z","shell.execute_reply.started":"2023-01-06T09:16:35.684926Z","shell.execute_reply":"2023-01-06T09:16:52.770393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:16:52.77278Z","iopub.execute_input":"2023-01-06T09:16:52.773146Z","iopub.status.idle":"2023-01-06T09:17:09.388Z","shell.execute_reply.started":"2023-01-06T09:16:52.773109Z","shell.execute_reply":"2023-01-06T09:17:09.387008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:17:09.389426Z","iopub.execute_input":"2023-01-06T09:17:09.390108Z","iopub.status.idle":"2023-01-06T09:17:09.437296Z","shell.execute_reply.started":"2023-01-06T09:17:09.390069Z","shell.execute_reply":"2023-01-06T09:17:09.436424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:17:09.443398Z","iopub.execute_input":"2023-01-06T09:17:09.444272Z","iopub.status.idle":"2023-01-06T09:17:09.452022Z","shell.execute_reply.started":"2023-01-06T09:17:09.444236Z","shell.execute_reply":"2023-01-06T09:17:09.45022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.3 Save Artifact","metadata":{}},{"cell_type":"code","source":"# Save new dataset\ntrain.to_csv(\"train_path.csv\", index=False)\n\n# 🐝 Save Artifacts\nsave_dataset_artifact(run_name=\"save_train_prep\", \n                      artifact_name=\"train_prep\",\n                      path=\"/kaggle/input/rsna-breast-cancer-helper-data/train_path.csv\",\n                      data_type=\"dataset\")","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:17:09.45358Z","iopub.execute_input":"2023-01-06T09:17:09.454217Z","iopub.status.idle":"2023-01-06T09:17:51.559303Z","shell.execute_reply.started":"2023-01-06T09:17:09.454171Z","shell.execute_reply":"2023-01-06T09:17:51.558155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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,"execution":{"iopub.status.busy":"2023-01-06T09:17:51.561055Z","iopub.execute_input":"2023-01-06T09:17:51.561789Z","iopub.status.idle":"2023-01-06T09:18:02.925411Z","shell.execute_reply.started":"2023-01-06T09:17:51.561741Z","shell.execute_reply":"2023-01-06T09:18:02.92419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:02.928487Z","iopub.execute_input":"2023-01-06T09:18:02.929685Z","iopub.status.idle":"2023-01-06T09:18:05.548305Z","shell.execute_reply.started":"2023-01-06T09:18:02.929634Z","shell.execute_reply":"2023-01-06T09:18:05.547255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:05.549709Z","iopub.execute_input":"2023-01-06T09:18:05.550782Z","iopub.status.idle":"2023-01-06T09:18:05.698076Z","shell.execute_reply.started":"2023-01-06T09:18:05.550742Z","shell.execute_reply":"2023-01-06T09:18:05.696797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Sample down for dev\n# train = train.sample(n=500, random_state=13).reset_index(drop=True)\n# train[\"cancer\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:18:05.699466Z","iopub.execute_input":"2023-01-06T09:18:05.700151Z","iopub.status.idle":"2023-01-06T09:18:05.706121Z","shell.execute_reply.started":"2023-01-06T09:18:05.70011Z","shell.execute_reply":"2023-01-06T09:18:05.705023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:05.707498Z","iopub.execute_input":"2023-01-06T09:18:05.708522Z","iopub.status.idle":"2023-01-06T09:18:05.716352Z","shell.execute_reply.started":"2023-01-06T09:18:05.708481Z","shell.execute_reply":"2023-01-06T09:18:05.715144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:05.718232Z","iopub.execute_input":"2023-01-06T09:18:05.718639Z","iopub.status.idle":"2023-01-06T09:18:05.730766Z","shell.execute_reply.started":"2023-01-06T09:18:05.718605Z","shell.execute_reply":"2023-01-06T09:18:05.7296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:05.732312Z","iopub.execute_input":"2023-01-06T09:18:05.732747Z","iopub.status.idle":"2023-01-06T09:18:05.74313Z","shell.execute_reply.started":"2023-01-06T09:18:05.732707Z","shell.execute_reply":"2023-01-06T09:18:05.741998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:05.745006Z","iopub.execute_input":"2023-01-06T09:18:05.745407Z","iopub.status.idle":"2023-01-06T09:18:14.833628Z","shell.execute_reply.started":"2023-01-06T09:18:05.745351Z","shell.execute_reply":"2023-01-06T09:18:14.832512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:14.835831Z","iopub.execute_input":"2023-01-06T09:18:14.836862Z","iopub.status.idle":"2023-01-06T09:18:14.847078Z","shell.execute_reply.started":"2023-01-06T09:18:14.836818Z","shell.execute_reply":"2023-01-06T09:18:14.846074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:14.850528Z","iopub.execute_input":"2023-01-06T09:18:14.850841Z","iopub.status.idle":"2023-01-06T09:18:27.05714Z","shell.execute_reply.started":"2023-01-06T09:18:14.850807Z","shell.execute_reply":"2023-01-06T09:18:27.056034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.2 EffNet\n\n* The issue with ResNet & other architectures:\n    * CNNs are usually developed at a **fixed resource cost** and then scaled up when more resources are available.\n    * **The scaling up** - increase arbitrary the depth or width of the CNN, then manually tune.\n    * This is **not optimal**.\n* **[EffNets](https://ai.googleblog.com/2019/05/efficientnet-improving-accuracy-and.html) try to solve this issue** through a more \"principled\" method of scaling up.\n* They use a *compound coefficient* to scale up CNNs in a more structured manner.\n    * uniformly scale each dimension (width, height, resolution) with a **fixed set of scaling coefficients**.\n    * this *balanced* scaling is much more efficient than scaling only 1 dimension at a time.\n* They also use **AutoML**:\n    * Automated Machine Learning\n    * uses Reinforcement Learning\n* The relationship between different scaling dimensions is found through a **Grid Search** (to determine the appropriate scaling coeff for each dimension). These coefficients are then applied to scale up the *baseline network* (from which we can then expand).\n<center><img src=\"https://1.bp.blogspot.com/-Cdtb97FtgdA/XO3BHsB7oEI/AAAAAAAAEKE/bmtkonwgs8cmWyI5esVo8wJPnhPLQ5bGQCLcBGAs/s640/image4.png\" width=600></center>\n* Many EffNets: B0, B1, B2, B3, B4, B5, B6, B7.\n    * smaller than other architectures - the number of parameters are smaller\n    * but higher efficiency\n<center><img src=\"https://1.bp.blogspot.com/-oNSfIOzO8ko/XO3BtHnUx0I/AAAAAAAAEKk/rJ2tHovGkzsyZnCbwVad-Q3ZBnwQmCFsgCEwYBhgL/s640/image3.png\" width=400></center>","metadata":{}},{"cell_type":"code","source":"class EffNetNetwork(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 = EfficientNet.from_pretrained('efficientnet-b2')\n        \n        # (CSV)\n        self.csv = nn.Sequential(nn.Linear(self.no_columns, 250),\n                                 nn.BatchNorm1d(250),\n                                 nn.ReLU(),\n                                 nn.Dropout(p=0.2),\n                                 \n                                 nn.Linear(250, 250),\n                                 nn.BatchNorm1d(250),\n                                 nn.ReLU(),\n                                 nn.Dropout(p=0.2))\n        \n        # Define Classification part\n        self.classification = nn.Sequential(nn.Linear(1408 + 250, self.output_size))\n        \n        \n    def forward(self, image, meta, prints=False):   \n        \n        if prints: print('Input Image shape:', image.shape, '\\n'+\n                         'Input metadata shape:', meta.shape)\n        \n        # Image CNN\n        image = self.features.extract_features(image)\n        image = F.avg_pool2d(image, image.size()[2:]).reshape(-1, 1408)\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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:27.058721Z","iopub.execute_input":"2023-01-06T09:18:27.059343Z","iopub.status.idle":"2023-01-06T09:18:27.07234Z","shell.execute_reply.started":"2023-01-06T09:18:27.059304Z","shell.execute_reply":"2023-01-06T09:18:27.071139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### ◾ Sanity Check\n\n> How does it work:\n<center><img src=\"https://i.imgur.com/vdonWT7.jpg\" width=700></center>","metadata":{}},{"cell_type":"code","source":"# Load Model\nmodel_example2 = EffNetNetwork(output_size=output_size, no_columns=no_columns).to(DEVICE)\n\n# Outputs\nout = model_example2(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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:27.075088Z","iopub.execute_input":"2023-01-06T09:18:27.075602Z","iopub.status.idle":"2023-01-06T09:18:29.081084Z","shell.execute_reply.started":"2023-01-06T09:18:27.075575Z","shell.execute_reply":"2023-01-06T09:18:29.07987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 6. Training ...\n\n### 🎗️ Small function to save in an external file","metadata":{}},{"cell_type":"code","source":"1408+250","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:18:29.082719Z","iopub.execute_input":"2023-01-06T09:18:29.083112Z","iopub.status.idle":"2023-01-06T09:18:29.089847Z","shell.execute_reply.started":"2023-01-06T09:18:29.083075Z","shell.execute_reply":"2023-01-06T09:18:29.088733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-01-06T09:18:29.091466Z","iopub.execute_input":"2023-01-06T09:18:29.092324Z","iopub.status.idle":"2023-01-06T09:18:29.100519Z","shell.execute_reply.started":"2023-01-06T09:18:29.092287Z","shell.execute_reply":"2023-01-06T09:18:29.099347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 🎗️ Main Training Function\n\n> How does it work:\n<center><img src=\"https://i.imgur.com/7MWwg9W.png\" width=800></center>","metadata":{}},{"cell_type":"code","source":"def train_folds(model, train_original):\n    # Creates a .txt file that will contain the logs\n    # logs == what we also print to console\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(X = np.zeros(len(train_original)), \n                               y = train_original['cancer'], \n                               groups = train_original['patient_id'].tolist())\n    \n    # For each fold\n    for i, (train_index, valid_index) in enumerate(k_folds):\n        \n        print(clr.S+f\"---------- Fold: {i+1} ----------\"+clr.E)\n        add_in_file(f\"---------- Fold: {i+1} ----------\", f)\n        \n        # 🐝 W&B Tracking\n        RUN_CONFIG = CONFIG.copy()\n        params = dict(model=MODEL, \n                      version=VERSION,\n                      fold=i,\n                      epochs=EPOCHS, \n                      batch=BATCH_SIZE1,\n                      lr=LR,\n                      weight_decay=WD)\n        RUN_CONFIG.update(params)\n        run = wandb.init(project='RSNA_Breast_Cancer', config=RUN_CONFIG)\n\n        wandb.watch(model, log_freq=100) # 🐝\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(model.parameters(), lr = LR, \n                                     weight_decay=WD)\n        scheduler = ReduceLROnPlateau(optimizer=optimizer, mode='max', \n                                      patience=LR_PATIENCE, verbose=True, factor=LR_FACTOR)\n        criterion = nn.BCEWithLogitsLoss()\n\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(train_data, vertical_flip, horizontal_flip, \n                            is_train=True)\n        valid = RSNADataset(valid_data, vertical_flip, horizontal_flip,\n                            is_train=True)\n\n        # Dataloaders\n        train_loader = DataLoader(train, batch_size=BATCH_SIZE1, \n                                  shuffle=True, num_workers=WORKERS)\n        valid_loader = DataLoader(valid, batch_size=BATCH_SIZE2, \n                                  shuffle=False, num_workers=WORKERS)\n\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                wandb.log({\"train_loss\": loss.item()}, step=epoch) # 🐝\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() == targets.cpu().unsqueeze(1)).sum().item()\n\n            # Compute Train Accuracy\n            train_acc = correct / len(train_index)\n            wandb.log({\"train_acc\": train_acc}) # 🐝\n\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(size = (len(valid_index), 1), \n                                      device=DEVICE, 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] + image.shape[0]] = pred\n\n                # Calculate accuracy\n                valid_acc = accuracy_score(valid_data['cancer'].values, \n                                           torch.round(valid_preds.cpu()))\n                wandb.log({\"valid_acc\": valid_acc}) # 🐝\n                # Calculate ROC\n                valid_roc = roc_auc_score(valid_data['cancer'].values, \n                                          valid_preds.cpu())\n                wandb.log({\"valid_roc\": valid_roc}) # 🐝\n\n                # Calculate time on Train + Eval\n                duration = str(dtime.timedelta(seconds=time() - start_time))[:7]\n\n\n                # PRINT INFO\n                final_logs = '{} | Epoch: {}/{} | Loss: {:.4} | Acc_tr: {:.3} | Acc_vd: {:.3} | ROC: {:.3}'.\\\n                                format(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\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: {}'.\\\n                                    format(best_roc)\n                        add_in_file(stop_logs, f)\n                        print(stop_logs)\n                        break\n\n\n        # === CLEANING ===\n        # Clear memory\n        del train, valid, train_loader, valid_loader, image, targets\n        gc.collect()\n        \n        # 🐝 Experiment End for this fold\n        wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:18:29.102202Z","iopub.execute_input":"2023-01-06T09:18:29.102622Z","iopub.status.idle":"2023-01-06T09:18:29.128389Z","shell.execute_reply.started":"2023-01-06T09:18:29.102587Z","shell.execute_reply":"2023-01-06T09:18:29.127477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 🎗️ Experimenting\n\n> 📌 **Note**: This cell below was ran locally on my ZBook Studio (as the Kaggle environment is too slow for all the data). I am sharing with you the logs that I've got :)","metadata":{}},{"cell_type":"code","source":"FOLDS = 3\nEPOCHS = 3\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 = 'v1'\nMODEL = 'resnet50'\n\nmodel1 = ResNet50Network(output_size=output_size, no_columns=no_columns).to(DEVICE)\n\n# ------------------\n\n# Run the cell below to train\n# Ran it locally on all data, see the results below\n# train_folds(model=model1, train_original=train)\n\n# Print the logs during training\nf = open('/kaggle/input/rsna-breast-cancer-helper-data/logs_v1.txt', \"r\")\ncontents = f.read()\nprint(contents)","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:18:29.131451Z","iopub.execute_input":"2023-01-06T09:18:29.131711Z","iopub.status.idle":"2023-01-06T09:18:29.660818Z","shell.execute_reply.started":"2023-01-06T09:18:29.131687Z","shell.execute_reply":"2023-01-06T09:18:29.659646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FOLDS = 3\nEPOCHS = 3\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 = 'v2'\nMODEL = 'effnet'\n\nmodel2 = EffNetNetwork(output_size=output_size, no_columns=no_columns).to(DEVICE)\n\n# ------------------\n\n# Run the cell below to train\n# Ran it locally on all data, see the results below\n# train_folds(model=model2, train_original=train)\n\n# Print the logs during training\nf = open('/kaggle/input/rsna-breast-cancer-helper-data/logs_v2.txt', \"r\")\ncontents = f.read()\nprint(contents)","metadata":{"execution":{"iopub.status.busy":"2023-01-08T08:08:26.103907Z","iopub.execute_input":"2023-01-08T08:08:26.104285Z","iopub.status.idle":"2023-01-08T08:08:26.116897Z","shell.execute_reply.started":"2023-01-08T08:08:26.104256Z","shell.execute_reply":"2023-01-08T08:08:26.116045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# WIP","metadata":{"execution":{"iopub.status.busy":"2023-01-06T09:18:29.671019Z","iopub.execute_input":"2023-01-06T09:18:29.671804Z","iopub.status.idle":"2023-01-06T09:18:29.678813Z","shell.execute_reply.started":"2023-01-06T09:18:29.671752Z","shell.execute_reply":"2023-01-06T09:18:29.677707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 🐝 [W&B Dashboard](https://wandb.ai/andrada/RSNA_Breast_Cancer)\n    \n<center><img src=\"https://i.imgur.com/HBPrqoJ.jpg\"></center>\n\n------\n\n<center><img src=\"https://i.imgur.com/knxTRkO.png\"></center>\n\n### My Specs\n\n* 🖥 Z8 G4 Workstation\n* 💾 2 CPUs & 96GB Memory\n* 🎮 2x NVIDIA A6000\n* 💻 Zbook Studio G7 on the go","metadata":{}}]}