{"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":"<p style=\"font-family: monospace; \n          font-weight: bold; \n          letter-spacing: 1px; \n          color: black; \n          font-size: 200%; \n          text-align: left;\n          padding: 0px; \n          border-bottom: 5px solid #fec89a\" >RSNA Breast Cancer Detection with PyTorch</p>\n          \nMy primary motivation behind creating this notebook was to get more practice using PyTorch. The prediction task in this notebook is to predict the presence or absence of cancer in the mammography images. The model I trained leverages both imaging and tabular data to make predictions. ","metadata":{}},{"cell_type":"markdown","source":"<p style=\"font-family: monospace; \n          font-weight: bold; \n          letter-spacing: 1px; \n          color: black; \n          font-size: 200%; \n          text-align: left;\n          padding: 0px; \n          border-bottom: 5px solid #fec89a\" >Table of Contents</p>\n\n* [Data and Cohort Characteristics](#section-one)\n* [Exploratory Data Analysis (EDA)](#section-two)\n* [Data Set and Data Loader](#section-three)\n* [PyTorch Model](#section-four)\n* [Training Loop](#section-five)\n* [Evaluation](#section-six)\n* [Conclusion](#section-seven)","metadata":{}},{"cell_type":"code","source":"# Installs\n!pip install -qU python-gdcm pydicom pylibjpeg\n!pip install polars\n!pip install lets-plot","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-07-04T22:54:33.39299Z","iopub.execute_input":"2023-07-04T22:54:33.393351Z","iopub.status.idle":"2023-07-04T22:55:07.366873Z","shell.execute_reply.started":"2023-07-04T22:54:33.393328Z","shell.execute_reply":"2023-07-04T22:55:07.365828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Imports\nimport os\nimport cv2\nimport glob\nimport gdcm\nimport torch\nimport random\nimport pydicom\nimport torchvision\nimport numpy as np\nimport polars as pl\nimport torch.nn as nn\nimport statistics as stats\nimport torch.optim as optim\nimport matplotlib.pyplot as plt\nimport torch.nn.functional as F\n\nfrom PIL import Image\nfrom lets_plot import *\nfrom torchvision import transforms\nfrom sklearn.metrics import confusion_matrix\nfrom torchvision.models import resnet50, ResNet50_Weights\nfrom sklearn.metrics import roc_auc_score, accuracy_score, precision_score, recall_score\nfrom lets_plot.mapping import as_discrete\nfrom torch.utils.data import Dataset, DataLoader\n\n# So the plots look nice\nLetsPlot.setup_html()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-07-04T22:55:15.065102Z","iopub.execute_input":"2023-07-04T22:55:15.065487Z","iopub.status.idle":"2023-07-04T22:55:17.988674Z","shell.execute_reply.started":"2023-07-04T22:55:15.065457Z","shell.execute_reply":"2023-07-04T22:55:17.987437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-one\"></a>\n<p style=\"font-family: monospace; \n          font-weight: bold; \n          letter-spacing: 1px; \n          color: black; \n          font-size: 200%; \n          text-align: left;\n          padding: 0px; \n          border-bottom: 5px solid #fec89a\" >Data and Cohort Characteristics</p>","metadata":{}},{"cell_type":"code","source":"# Data\ndf_train = pl.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ndf_test = pl.read_csv('/kaggle/input/rsna-breast-cancer-detection/test.csv')","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2023-07-04T22:55:28.872484Z","iopub.execute_input":"2023-07-04T22:55:28.873079Z","iopub.status.idle":"2023-07-04T22:55:29.057434Z","shell.execute_reply.started":"2023-07-04T22:55:28.87305Z","shell.execute_reply":"2023-07-04T22:55:29.056406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.glimpse()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T22:55:30.281973Z","iopub.execute_input":"2023-07-04T22:55:30.282317Z","iopub.status.idle":"2023-07-04T22:55:30.299303Z","shell.execute_reply.started":"2023-07-04T22:55:30.282295Z","shell.execute_reply":"2023-07-04T22:55:30.297814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Missing values by column')\ndf_train.select(pl.all().is_null().sum())","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T22:55:32.368787Z","iopub.execute_input":"2023-07-04T22:55:32.369355Z","iopub.status.idle":"2023-07-04T22:55:32.42556Z","shell.execute_reply.started":"2023-07-04T22:55:32.369326Z","shell.execute_reply":"2023-07-04T22:55:32.424351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-two\"></a>\n<p style=\"font-family: monospace; \n          font-weight: bold; \n          letter-spacing: 1px; \n          color: black; \n          font-size: 200%; \n          text-align: left;\n          padding: 0px; \n          border-bottom: 5px solid #fec89a\" >Exploratory Data Analysis (EDA)</p>","metadata":{}},{"cell_type":"code","source":"# Initialize colors and stuff\ntarget='cancer'\ncolor1='#d8e2dc'\ncolor2='#f4acb7'\ncolor3='#ee4266'\n\ndf_plt = df_train.with_columns(\n    pl.when(pl.col('cancer') == 1).then('Cancer Present').otherwise('No Cancer').alias('cancer'),\n    pl.when((pl.col('cancer') == 1) & (pl.col('invasive') == 1))\\\n        .then('Invasive')\\\n        .when((pl.col('cancer') == 1) & (pl.col('invasive') == 0))\\\n        .then('Non-Invasive')\\\n        .when((pl.col('cancer') == 0) & (pl.col('invasive') == 0))\\\n        .then('No Cancer')\\\n        .otherwise('No Cancer')\\\n        .alias('invasive2')\n)\n\n# Target variable plot\nvar = 'cancer'\ntitle = 'Cancer Distribution'\nlegend_title = ''\n\nplt1 = \\\n    ggplot(df_plt)+\\\n    geom_bar(aes(x = as_discrete(var),\n                fill = as_discrete(var)),\n            color = None,\n            size = 0.5)+\\\n    scale_fill_manual(values = [color1, color2])+\\\n    theme_minimal()+\\\n    theme(\n        plot_title = element_text(hjust = 0.5, face = 'bold'),\n        legend_position = \"top\",\n        panel_grid_major = element_blank(),\n        panel_grid_minor = element_blank(),\n        legend_title = element_blank(),\n        axis_title_x = element_blank(),\n        axis_line_y = element_line(size = 1))+\\\n    coord_flip()+\\\n    labs(y = \"Count\", title = title)\n\n# Target variable by invasiveness\nvar = 'invasive2'\ntitle = 'Cancer Distribution by Invasiveness'\nlegend_title = ''\n\nplt2 = \\\n    ggplot(df_plt)+\\\n    geom_bar(aes(x = as_discrete(var),\n                fill = as_discrete(var)),\n            color = None,\n            size = 0.5)+\\\n    scale_fill_manual(values = [color1, color3, color2])+\\\n    theme_minimal()+\\\n    theme(\n        plot_title = element_text(hjust = 0.5, face = 'bold'),\n        legend_position = \"top\",\n        panel_grid_major = element_blank(),\n        panel_grid_minor = element_blank(),\n        legend_title = element_blank(),\n        axis_title_x = element_blank(),\n        axis_line_y = element_line(size = 1))+\\\n    coord_flip()+\\\n    labs(y = \"Count\", title = title)\n\n\n\nbunch = GGBunch()\nbunch.add_plot(plt1, 0, 0, 500, 250)\nbunch.add_plot(plt2, 520, 0, 500, 250)\nbunch.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T22:55:39.055819Z","iopub.execute_input":"2023-07-04T22:55:39.056216Z","iopub.status.idle":"2023-07-04T22:55:42.674848Z","shell.execute_reply.started":"2023-07-04T22:55:39.056186Z","shell.execute_reply":"2023-07-04T22:55:42.673786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# By cancer status\nvar = 'age'\ntitle = 'Age by Cancer Status'\nxlab = 'Age (years)'\n\nplt1 =\\\n    ggplot(df_plt)+\\\n    geom_density(aes(x = var, fill = target), \n                 color = 'gray', alpha = 0.5)+\\\n    scale_fill_manual(values = [color1, color2])+\\\n    theme(plot_title = element_text(hjust = 0.5, face = 'bold'),\n         legend_position = \"top\",\n         legend_title = element_blank())+\\\n    labs(x = xlab, y = 'Proportion', title = title)\n\nplt2 =\\\n    ggplot(df_plt)+\\\n    geom_boxplot(aes(x = target, y = var, \n                     fill = target), outlier_shape = 21)+\\\n    scale_fill_manual(values = [color1, color2])+\\\n    scale_x_continuous(labels=['     ','     '])+\\\n    coord_flip()+\\\n    theme_minimal()+\\\n    theme(plot_title = element_text(hjust = 0.5, face = 'bold'),\n          legend_position = \"top\",\n          legend_title = element_blank())+\\\n    labs(x = '', y = xlab, title = '')\n\n# By invasiveness\nvar = 'age'\ntitle = 'Age by Cancer Invasiveness'\nxlab = 'Age (years)'\n\nplt3 =\\\n    ggplot(df_plt)+\\\n    geom_density(aes(x = var, fill = 'invasive2'), \n                 color = 'gray', alpha = 0.5)+\\\n    scale_fill_manual(values = [color1, color3, color2])+\\\n    theme(plot_title = element_text(hjust = 0.5, face = 'bold'),\n         legend_position = \"top\",\n         legend_title = element_blank())+\\\n    labs(x = xlab, y = 'Proportion', title = title)\n\nplt4 =\\\n    ggplot(df_plt)+\\\n    geom_boxplot(aes(x = 'invasive2', y = var, \n                     fill = 'invasive2'), outlier_shape = 21)+\\\n    scale_fill_manual(values = [color1, color3, color2])+\\\n    scale_x_continuous(labels=['     ','     '])+\\\n    coord_flip()+\\\n    theme_minimal()+\\\n    theme(plot_title = element_text(hjust = 0.5, face = 'bold'),\n          legend_position = \"top\",\n          legend_title = element_blank())+\\\n    labs(x = '', y = xlab, title = '')\n\nbunch = GGBunch()\nbunch.add_plot(plt1, 0, 0, 500, 250)\nbunch.add_plot(plt2, 0, 250, 500, 250)\nbunch.add_plot(plt3, 520, 0, 500, 250)\nbunch.add_plot(plt4, 520, 250, 500, 250)\nbunch.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T22:55:47.898828Z","iopub.execute_input":"2023-07-04T22:55:47.899203Z","iopub.status.idle":"2023-07-04T22:55:58.300488Z","shell.execute_reply.started":"2023-07-04T22:55:47.899173Z","shell.execute_reply":"2023-07-04T22:55:58.29894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<p style=\"font-family: monospace; \n          font-weight: bold; \n          letter-spacing: 1px; \n          color: black; \n          font-size: 150%; \n          text-align: left;\n          padding: 0px; \n          border-bottom: 2px solid #fec89a\" >Images by Presence of Cancer</p>","metadata":{}},{"cell_type":"code","source":"# Function for getting a dicom plot from letsplot\ndef get_dicom_plt(dcm_path, title):\n    \n    dcm = pydicom.dcmread(dcm_path)\n    img = dcm.pixel_array\n    img = (img - img.min()) / (img.max() - img.min())\n    \n    if dcm.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n        \n    plt=\\\n        ggplot()+\\\n        geom_imshow(img)+\\\n        theme(\n            legend_position='none', \n            panel_grid=element_blank(), \n            axis=element_blank(),\n            plot_title=element_text(hjust=0.5, face='bold'))+\\\n        labs(title=title)\n    \n    return plt","metadata":{"execution":{"iopub.status.busy":"2023-07-04T22:56:15.611465Z","iopub.execute_input":"2023-07-04T22:56:15.611818Z","iopub.status.idle":"2023-07-04T22:56:15.619436Z","shell.execute_reply.started":"2023-07-04T22:56:15.611792Z","shell.execute_reply":"2023-07-04T22:56:15.618064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Getting info based on cancer status\ninvasive_patients = df_plt.filter(pl.col('invasive2') == 'Invasive').select(['patient_id', 'image_id'])\ninvasive_patient_ids = invasive_patients.get_column('patient_id')\ninvasive_img_ids = invasive_patients.get_column('image_id')\n\nnoninvasive_patients = df_plt.filter(pl.col('invasive2') == 'Non-Invasive').select(['patient_id', 'image_id'])\nnoninvasive_patient_ids = noninvasive_patients.get_column('patient_id')\nnoninvasive_img_ids = noninvasive_patients.get_column('image_id')\n\nno_cancer_patients = df_plt.filter(pl.col('invasive2') == 'No Cancer').select(['patient_id', 'image_id'])\nno_cancer_patient_ids = no_cancer_patients.get_column('patient_id')\nno_cancer_img_ids = no_cancer_patients.get_column('image_id')\n\n# Initializing stuff\nimg_dir = '/kaggle/input/rsna-breast-cancer-detection/train_images'\nnpatients = 3\nbunch = GGBunch()\n\n# Plotting\nfor i in range(npatients):\n    \n    # For invasive cancer patients\n    dcm_path = f'{img_dir}/{invasive_patient_ids[i]}/{invasive_img_ids[i]}.dcm'\n    bunch.add_plot(get_dicom_plt(dcm_path, title = 'Invasive Cancer'), 0 + i*300, 0, 300, 300)\n    \n    # For noninvasive cancer patients\n    dcm_path = f'{img_dir}/{noninvasive_patient_ids[i]}/{noninvasive_img_ids[i]}.dcm'\n    bunch.add_plot(get_dicom_plt(dcm_path, title = 'Non-Invasive Cancer'), 0 + i*300, 310, 300, 300)\n\n    # For those without cancer\n    dcm_path = f'{img_dir}/{no_cancer_patient_ids[i]}/{no_cancer_img_ids[i]}.dcm'\n    bunch.add_plot(get_dicom_plt(dcm_path, title = 'No Cancer'), 0 + i*300, 620, 300, 300)\n                   \nbunch.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T22:56:24.336007Z","iopub.execute_input":"2023-07-04T22:56:24.337292Z","iopub.status.idle":"2023-07-04T22:56:43.271255Z","shell.execute_reply.started":"2023-07-04T22:56:24.337233Z","shell.execute_reply":"2023-07-04T22:56:43.267275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-three\"></a>\n<p style=\"font-family: monospace; \n          font-weight: bold; \n          letter-spacing: 1px; \n          color: black; \n          font-size: 200%; \n          text-align: left;\n          padding: 0px; \n          border-bottom: 5px solid #fec89a\" >Data Set and Data Loader</p>\n          \n[Another notebook](https://www.kaggle.com/code/theoviel/dicom-resized-png-jpg) previously converted all the dicom files into 256x256 png files. Instead of re-creating the code I simply imported the data set (which can be found [here](https://www.kaggle.com/datasets/theoviel/rsna-breast-cancer-256-pngs)). Since my time is limited I decided to downsample the data set so the target is balanced. \n\nA number of additional pre-processing techniques were applied:\n* Missing values in the 'age' metadata column were imputed using the mean\n* The view, laterality, and implant features were dummy encoded\n* A column with the image filenames was added to the df","metadata":{}},{"cell_type":"code","source":"# Downsampling the non-cancerous targets and creating a final df\ndf_target1 = df_train.filter(pl.col('cancer') == 1)\ndf_target0 = df_train.filter(pl.col('cancer') == 0)\n\nn_idx0 = len(df_target0) # Number of noncancerous patients\n\ndf_target0 = df_target0\\\n    .with_row_count()\\\n    .filter(pl.col('row_nr').is_in(random.sample(range(n_idx0), len(df_target1))))\\\n    .drop('row_nr')\n\n# Final df\ndf_keep = \\\n    pl.concat([df_target1, df_target0], how='vertical')\\\n    .select(pl.all().shuffle(seed=19970507))\n\n# Imputing missing values for age, normalizing\nage_mean = round(df_keep.get_column('age').mean())\nage_min = df_keep.get_column('age').min()\nage_max = df_keep.get_column('age').max()\n\ndf_keep = \\\n    df_keep.with_columns(\n    \n        pl.when(pl.col('age') == None)\\\n        .then(age_mean)\\\n        .otherwise(pl.col('age'))\\\n        .alias('age')\n    \n    ).with_columns(\n    \n        ((pl.col('age') - age_min)/(age_max - age_min)).alias('age')\n    \n    ).to_dummies(\n    \n        ['view', 'laterality', 'implant']\n    \n    )\n\n# Adding a train/valid indicator column\nn = len(df_keep)\nntrain = round(n*.80)\n\ndf_keep =\\\n    df_keep.with_row_count()\\\n    .with_columns(\n        pl.when(pl.col('row_nr') <= ntrain)\\\n        .then('train')\\\n        .otherwise('valid')\\\n        .alias('trainvalid')\n    )\\\n    .drop('row_nr')\n    \n# Getting train/valid filenames and labels\n# Note the 256x256 pngs are names {patient_id}_{image_id}\ndf_keep = df_keep\\\n    .with_columns(pl.lit('_').alias('underscore'))\\\n    .with_columns(\n        pl.concat_str(\n            [\n                pl.col('patient_id'),\n                pl.col('underscore'),\n                pl.col('image_id')\n            ]\n        ).alias('fname')\n    ).drop('underscore')\n\ndf_train_meta = df_keep.filter(pl.col('trainvalid') == 'train')\ndf_valid_meta = df_keep.filter(pl.col('trainvalid') == 'valid')\n\ntrain_fnames = df_train_meta.get_column('fname')\nvalid_fnames = df_valid_meta.get_column('fname')\n\ntrain_labels = df_train_meta.get_column('cancer').to_numpy()\nvalid_labels = df_valid_meta.get_column('cancer').to_numpy()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T22:57:12.990699Z","iopub.execute_input":"2023-07-04T22:57:12.991085Z","iopub.status.idle":"2023-07-04T22:57:13.069415Z","shell.execute_reply.started":"2023-07-04T22:57:12.991053Z","shell.execute_reply":"2023-07-04T22:57:13.068682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Training Metadata Characteristics')\ndf_train_meta.glimpse()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T22:57:18.196186Z","iopub.execute_input":"2023-07-04T22:57:18.19683Z","iopub.status.idle":"2023-07-04T22:57:18.203349Z","shell.execute_reply.started":"2023-07-04T22:57:18.196798Z","shell.execute_reply":"2023-07-04T22:57:18.202182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Validation Metadata Characteristics')\ndf_valid_meta.glimpse()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T22:57:21.889667Z","iopub.execute_input":"2023-07-04T22:57:21.89074Z","iopub.status.idle":"2023-07-04T22:57:21.897984Z","shell.execute_reply.started":"2023-07-04T22:57:21.890701Z","shell.execute_reply":"2023-07-04T22:57:21.895487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get directories\nimg_dir='/kaggle/input/rsna-breast-cancer-256-pngs'\ntrain_dir='/kaggle/working/train'\nvalid_dir='/kaggle/working/valid'\n\n# Create directories for train/valid images\nos.makedirs(train_dir, exist_ok=True)\nos.makedirs(valid_dir, exist_ok=True)\n\n# Move the images into the correct directories\nfor file in train_fnames:\n    img = Image.open(f'{img_dir}/{file}.png')\n    img.save(f'{train_dir}/{file}.png')\n\nfor file in valid_fnames:\n    img = Image.open(f'{img_dir}/{file}.png')\n    img.save(f'{valid_dir}/{file}.png')","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T22:57:25.424417Z","iopub.execute_input":"2023-07-04T22:57:25.424794Z","iopub.status.idle":"2023-07-04T22:57:51.891798Z","shell.execute_reply.started":"2023-07-04T22:57:25.424765Z","shell.execute_reply":"2023-07-04T22:57:51.890462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Defining the data set\nclass MammographyDataset(Dataset):\n    def __init__(self, meta_df, img_dir, transform=None):\n        \n        self.df = meta_df\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        \n        # Get label from meta data\n        label = self.df.get_column('cancer')\n        label = label[idx]\n        \n        # Get image file paths from metadata\n        img_fname = self.df.get_column('fname')\n        img_fname = img_fname[idx]\n        \n        # Get image, transform\n        img_path = f'{self.img_dir}/{img_fname}.png'\n        img = Image.open(img_path)\n        \n        if self.transform:\n            img = self.transform(img)\n            \n        # Get metadata features\n        feature_names = [\n            'age', \n            'laterality_L', 'laterality_R', \n            'view_AT', 'view_CC', 'view_MLO',\n            'implant_0', 'implant_1'\n                        ]\n        \n        meta_features = self.df.select(feature_names)\n        meta_features = meta_features[idx,:].to_numpy()\n        \n        return img, meta_features, label\n\n# Defining the transformations\ntransform = transforms.Compose([\n    transforms.ToTensor(),\n])\n\n# Initializing the datasets\ntrain_dataset = MammographyDataset(\n    meta_df=df_train_meta,\n    img_dir='/kaggle/working/train',\n    transform=transform,\n)\n\nvalid_dataset = MammographyDataset(\n    meta_df=df_valid_meta,\n    img_dir='/kaggle/working/valid',\n    transform=transform,\n)\n\n# Initializing the dataloader\ntrain_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=64, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-07-04T22:58:30.094151Z","iopub.execute_input":"2023-07-04T22:58:30.094594Z","iopub.status.idle":"2023-07-04T22:58:30.104522Z","shell.execute_reply.started":"2023-07-04T22:58:30.094564Z","shell.execute_reply":"2023-07-04T22:58:30.103803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-four\"></a>\n<p style=\"font-family: monospace; \n          font-weight: bold; \n          letter-spacing: 1px; \n          color: black; \n          font-size: 200%; \n          text-align: left;\n          padding: 0px; \n          border-bottom: 5px solid #fec89a\" >PyTorch Model</p>\n          \nThe PyTorch model below leverages both the imaging data and the meta data to make predictions. More specifically, it uses the age, laterality, view, and implant features from the meta data. Previously age was min-max scaled and the remaining meta features were dummy encoded. \n\nA pre-trained ResNet50 model was used to output predicted probabilities which were then combined with the meta-data features and used as inputs to a fully connected network that returned the final predictions.","metadata":{}},{"cell_type":"code","source":"class MammographyModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n        # ResNet50 \n        self.rnet = torchvision.models.resnet50(weights=ResNet50_Weights.DEFAULT)\n        self.rnet.conv1 = torch.nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n        self.rnet.fc = torch.nn.Linear(in_features=2048, out_features=500)\n        \n        # Final classification network\n        self.fc1 = nn.Linear(508, 1)\n        self.sigmoid = nn.Sigmoid()\n        \n    def forward(self, img, meta_features):\n        \n        # ResNet50\n        resnet_out = self.rnet(img)\n        resnet_out = torch.sigmoid(resnet_out)\n        \n        # Reshape meta features\n        meta_features = meta_features.squeeze(1)\n        \n        # Get final predictions\n        x_final = torch.cat((resnet_out, meta_features), dim=1).to(torch.float32)        \n        x_final = self.fc1(x_final)\n        \n        out = self.sigmoid(x_final)\n        \n        return out \n    \nmodel = MammographyModel()","metadata":{"execution":{"iopub.status.busy":"2023-07-04T22:58:40.086452Z","iopub.execute_input":"2023-07-04T22:58:40.086824Z","iopub.status.idle":"2023-07-04T22:58:41.71305Z","shell.execute_reply.started":"2023-07-04T22:58:40.086797Z","shell.execute_reply":"2023-07-04T22:58:41.71133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Loss function and optimizer\nloss_fn = nn.BCELoss()\noptimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9)","metadata":{"execution":{"iopub.status.busy":"2023-07-04T22:58:46.250143Z","iopub.execute_input":"2023-07-04T22:58:46.250781Z","iopub.status.idle":"2023-07-04T22:58:46.25672Z","shell.execute_reply.started":"2023-07-04T22:58:46.250751Z","shell.execute_reply":"2023-07-04T22:58:46.255192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-five\"></a>\n<p style=\"font-family: monospace; \n          font-weight: bold; \n          letter-spacing: 1px; \n          color: black; \n          font-size: 200%; \n          text-align: left;\n          padding: 0px; \n          border-bottom: 5px solid #fec89a\" >Training Loop</p>\n          \n          \nBelow is the training loop fo the ResNet50 model. I did not try and adjust the hyperparameters or architecture to improve the performance as the main aim was to simply get more practice writing PyTorch code. ","metadata":{}},{"cell_type":"code","source":"# Handle the device stuff\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel.to(device)\nprint(f'Using device: {device}')","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T22:58:50.623931Z","iopub.execute_input":"2023-07-04T22:58:50.624586Z","iopub.status.idle":"2023-07-04T22:58:50.638891Z","shell.execute_reply.started":"2023-07-04T22:58:50.624556Z","shell.execute_reply":"2023-07-04T22:58:50.638153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Initialize number of epochs\nnpochs = 25\ntrain_loss, valid_loss = [], []\ntrain_accuracy, valid_accuracy = [], []\ntrain_sensitivity, valid_sensitivity = [], []\ntrain_specificity, valid_specificity = [], []\n\n# Helper function\ndef get_sens_spec(y_true, y_pred):\n    '''\n    Returns the sensitivity and specificity given\n    labels and class predictions\n    '''\n    \n    tn, fp, fn, tp = confusion_matrix(y_true, y_pred).ravel()\n    \n    sensitivity = tp / (tp + fn)\n    specificity = tn / (tn + fp)\n\n    return sensitivity, specificity\n\n# Training loop\nfor epoch in range(npochs):\n    \n    # Training section\n    model.train() \n    running_loss = 0.0   \n    train_preds, train_labels = [], []\n    \n    for batch, (img, meta_features, label) in enumerate(train_loader):\n        img = img.to(device)\n        meta_features = meta_features.to(device)\n        label = label.to(device)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(img, meta_features)\n        \n        loss = loss_fn(outputs.squeeze(), label.float())\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        \n        # Store preds and labels\n        preds = outputs.squeeze().detach().cpu().numpy().round()\n        train_preds.extend(preds)\n        train_labels.extend(label.cpu().numpy())\n        \n        if batch%5 == 0:\n            print(f'epoch {epoch + 1}  batch {batch + 1}  train loss: {loss.item():10.8f}')\n       \n    # Save performance metrics\n    train_sens, train_spec = get_sens_spec(train_labels, train_preds)\n    train_sensitivity.append(train_sens)\n    train_specificity.append(train_spec)\n    avg_train_loss = running_loss/len(train_loader)\n    train_loss.append(avg_train_loss)\n    train_accuracy.append(accuracy_score(train_labels, train_preds))\n    \n    # Eval section\n    model.eval()\n    running_loss = 0.0  \n    valid_preds, valid_labels = [], []\n    \n    with torch.no_grad():\n        for batch, (img, meta_features, label) in enumerate(valid_loader):\n            img = img.to(device)\n            meta_features = meta_features.to(device)\n            label = label.to(device)\n            \n            outputs = model(img, meta_features)\n            loss = loss_fn(outputs.squeeze(), label.float())\n            \n            running_loss += loss.item()\n            \n            # Store preds and labels\n            preds = outputs.squeeze().detach().cpu().numpy().round()\n            valid_preds.extend(preds)\n            valid_labels.extend(label.cpu().numpy())\n    \n    # Save performance metrics\n    valid_sens, valid_spec = get_sens_spec(valid_labels, valid_preds)\n    valid_sensitivity.append(valid_sens)\n    valid_specificity.append(valid_spec)\n    avg_valid_loss = running_loss/len(valid_loader)\n    valid_loss.append(avg_valid_loss)\n    valid_accuracy.append(accuracy_score(valid_labels, valid_preds))\n            \n    print(f'---------------------------------------------------------------------------------')\n    print(f'Metrics for epoch {epoch + 1}')\n    print(f'Accuracy     train: {train_accuracy[epoch]}  valid: {valid_accuracy[epoch]}')\n    print(f'Sensitivity  train: {train_sensitivity[epoch]}  valid: {valid_sensitivity[epoch]}')\n    print(f'Specificity  train: {train_specificity[epoch]}  valid: {valid_specificity[epoch]}')\n    print(f'---------------------------------------------------------------------------------')\n    ","metadata":{"execution":{"iopub.status.busy":"2023-07-04T22:58:58.914291Z","iopub.execute_input":"2023-07-04T22:58:58.914923Z","iopub.status.idle":"2023-07-04T23:33:48.678905Z","shell.execute_reply.started":"2023-07-04T22:58:58.914889Z","shell.execute_reply":"2023-07-04T23:33:48.677182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-six\"></a>\n<p style=\"font-family: monospace; \n          font-weight: bold; \n          letter-spacing: 1px; \n          color: black; \n          font-size: 200%; \n          text-align: left;\n          padding: 0px; \n          border-bottom: 5px solid #fec89a\" >Evaluation</p>\n          \n          \nBelow we can see that the model's performance is far from ideal. As my motivaiton behind this notebook was simply to practice writing PyTorch code I did not attempt to further adjust the architecture, hyperparameters, etc in order to improve the performance of the model.  ","metadata":{}},{"cell_type":"code","source":"# Initialize stuff\nepoch = 2*[x for x in range(1,npochs+1)]\nset_type = npochs*['Train'] + npochs*['Valid']\n\n# Loss plot\nlosses = train_loss + valid_loss # Lists\n\ndf_plt = pl.DataFrame({'epoch':epoch, 'set_type':set_type,'loss':losses})\n\nplt_loss=\\\n    ggplot(df_plt)+\\\n    geom_line(aes(x='epoch',y='loss',color='set_type'),size=2)+\\\n    labs(x='Epoch',y='Loss', title='Loss Tracking', color = '')+\\\n    scale_x_continuous(breaks=[x for x in range(1,npochs+1)])+\\\n    theme(\n        plot_title = element_text(hjust = 0.5, face = 'bold'),\n        legend_position = 'top',\n        axis_line_y = element_line(size = 1),\t\n        axis_line_x = element_line(size = 1),\n    )\n\n# Accuracy plot\naccuracies = train_accuracy + valid_accuracy # Lists\n\ndf_plt = pl.DataFrame({'epoch':epoch, 'set_type':set_type,'accuracy':accuracies})\n\nplt_acc=\\\n    ggplot(df_plt)+\\\n    geom_line(aes(x='epoch',y='accuracy',color='set_type'),size=2)+\\\n    labs(x='Epoch',y='Accuracy', title='Accuracy Tracking', color = '')+\\\n    scale_x_continuous(breaks=[x for x in range(1,npochs+1)])+\\\n    theme(\n        plot_title = element_text(hjust = 0.5, face = 'bold'),\n        legend_position = 'top',\n        axis_line_y = element_line(size = 1),\t\n        axis_line_x = element_line(size = 1),\n    )\n\n\n# Sensitivity plot\nsensitivities = train_sensitivity + valid_sensitivity # Lists\n\ndf_plt = pl.DataFrame({'epoch':epoch, 'set_type':set_type,'sensitivity':sensitivities})\n\nplt_sens=\\\n    ggplot(df_plt)+\\\n    geom_line(aes(x='epoch',y='sensitivity',color='set_type'),size=2)+\\\n    labs(x='Epoch',y='Sensitivity', title='Sensitivity Tracking', color = '')+\\\n    scale_x_continuous(breaks=[x for x in range(1,npochs+1)])+\\\n    theme(\n        plot_title = element_text(hjust = 0.5, face = 'bold'),\n        legend_position = 'top',\n        axis_line_y = element_line(size = 1),\t\n        axis_line_x = element_line(size = 1),\n    )\n\n# Specificity plot\nspecificities = train_specificity + valid_specificity # Lists\n\ndf_plt = pl.DataFrame({'epoch':epoch, 'set_type':set_type,'specificity':specificities})\n\nplt_spec=\\\n    ggplot(df_plt)+\\\n    geom_line(aes(x='epoch',y='specificity',color='set_type'),size=2)+\\\n    labs(x='Epoch',y='Specificity', title='Specificity Tracking', color = '')+\\\n    scale_x_continuous(breaks=[x for x in range(1,npochs+1)])+\\\n    theme(\n        plot_title = element_text(hjust = 0.5, face = 'bold'),\n        legend_position = 'top',\n        axis_line_y = element_line(size = 1),\t\n        axis_line_x = element_line(size = 1),\n    )\n\n# Bunch\nbunch = GGBunch()\nbunch.add_plot(plt_loss, 0, 0, 800, 400)\nbunch.add_plot(plt_acc, 0, 410, 800, 400)\nbunch.add_plot(plt_sens, 0, 820, 800, 400)\nbunch.add_plot(plt_spec, 0, 1230, 800, 400)\nbunch.show()","metadata":{"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-seven\"></a>\n<p style=\"font-family: monospace; \n          font-weight: bold; \n          letter-spacing: 1px; \n          color: black; \n          font-size: 200%; \n          text-align: left;\n          padding: 0px; \n          border-bottom: 5px solid #fec89a\" >Conclusion</p>\n          \nThis notebook provided an implementation of a PyTorch model for classifying mammography scans as cancerous or non-cancerous. If you made it this far thanks for reading!","metadata":{}}]}