{"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":"code","source":"import os\nimport glob\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:41.263968Z","iopub.execute_input":"2022-12-09T11:12:41.264366Z","iopub.status.idle":"2022-12-09T11:12:41.271237Z","shell.execute_reply.started":"2022-12-09T11:12:41.26433Z","shell.execute_reply":"2022-12-09T11:12:41.27008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset preparation","metadata":{}},{"cell_type":"code","source":"data_path = \"/kaggle/input/rsna-breast-cancer-detection\"\ntrain_image_path = os.path.join(data_path,\"train_images\")\ntest_image_path = os.path.join(data_path,\"test_images\")\nos.path.exists(train_image_path)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:41.273806Z","iopub.execute_input":"2022-12-09T11:12:41.274341Z","iopub.status.idle":"2022-12-09T11:12:41.288242Z","shell.execute_reply.started":"2022-12-09T11:12:41.274303Z","shell.execute_reply":"2022-12-09T11:12:41.286944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# find all csv files in the dataset\ncsv_files = glob.glob(data_path+\"/*.csv\")\ncsv_files","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:41.290652Z","iopub.execute_input":"2022-12-09T11:12:41.290952Z","iopub.status.idle":"2022-12-09T11:12:41.299066Z","shell.execute_reply.started":"2022-12-09T11:12:41.290927Z","shell.execute_reply":"2022-12-09T11:12:41.297967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train csv to dataframe\ntrain_df = pd.read_csv(csv_files[1])","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:41.300558Z","iopub.execute_input":"2022-12-09T11:12:41.301192Z","iopub.status.idle":"2022-12-09T11:12:41.406426Z","shell.execute_reply.started":"2022-12-09T11:12:41.301146Z","shell.execute_reply":"2022-12-09T11:12:41.404774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.info()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:41.410028Z","iopub.execute_input":"2022-12-09T11:12:41.410716Z","iopub.status.idle":"2022-12-09T11:12:41.435909Z","shell.execute_reply.started":"2022-12-09T11:12:41.410676Z","shell.execute_reply":"2022-12-09T11:12:41.434718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.columns","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:41.436845Z","iopub.execute_input":"2022-12-09T11:12:41.437448Z","iopub.status.idle":"2022-12-09T11:12:41.445609Z","shell.execute_reply.started":"2022-12-09T11:12:41.437411Z","shell.execute_reply":"2022-12-09T11:12:41.444443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### What is the distribution of cancer cases","metadata":{}},{"cell_type":"code","source":"train_df[\"cancer\"].value_counts()\n# Highly imbalanced data","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:41.447105Z","iopub.execute_input":"2022-12-09T11:12:41.447759Z","iopub.status.idle":"2022-12-09T11:12:41.457909Z","shell.execute_reply.started":"2022-12-09T11:12:41.44772Z","shell.execute_reply":"2022-12-09T11:12:41.456835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Cancer distribution by age**","metadata":{}},{"cell_type":"code","source":"train_df[\"view\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:41.459386Z","iopub.execute_input":"2022-12-09T11:12:41.460479Z","iopub.status.idle":"2022-12-09T11:12:41.471596Z","shell.execute_reply.started":"2022-12-09T11:12:41.460444Z","shell.execute_reply":"2022-12-09T11:12:41.470545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df[\"age\"].describe()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:41.473078Z","iopub.execute_input":"2022-12-09T11:12:41.474058Z","iopub.status.idle":"2022-12-09T11:12:41.489741Z","shell.execute_reply.started":"2022-12-09T11:12:41.474023Z","shell.execute_reply":"2022-12-09T11:12:41.488472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sns\nplt.figure(figsize=(20,10))\nsns.histplot(data=train_df,x=\"age\",hue=\"cancer\",kde=\"true\")\n","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:41.491256Z","iopub.execute_input":"2022-12-09T11:12:41.492121Z","iopub.status.idle":"2022-12-09T11:12:42.572848Z","shell.execute_reply.started":"2022-12-09T11:12:41.49208Z","shell.execute_reply":"2022-12-09T11:12:42.571743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Bi-RADS (The Breast Imaging Reporting and Data System)\n[For the purposes of this challenge, the most important element of BI-RADS to understand is the BI-RADS score.](https://www.kaggle.com/competitions/rsna-breast-cancer-detection/discussion/369262)\n* There are 7 main BI-RADS scores, or categories:\n    *     0 - Need additional imaging evaluation\n    *     1 - Negative\n    *     2 - Benign\n    *     3 - Probably Benign\n    *     4 - Suspicious\n    *     5 - Highly Suggestive of Malignancy\n    *     6 - Known Biopsy-Proven Malignancy\n* Screening mammograms can only be assigned a BI-RADS score of 0, 1, or 2. This is why the dataset only contains 3 of the BI-RADS categories.","metadata":{}},{"cell_type":"code","source":"train_df.BIRADS.value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:42.574688Z","iopub.execute_input":"2022-12-09T11:12:42.57545Z","iopub.status.idle":"2022-12-09T11:12:42.587326Z","shell.execute_reply.started":"2022-12-09T11:12:42.575407Z","shell.execute_reply":"2022-12-09T11:12:42.585533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,10))\nsns.histplot(data=train_df,x=\"cancer\",hue=\"BIRADS\",kde=\"true\",palette=\"tab10\")","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:42.591065Z","iopub.execute_input":"2022-12-09T11:12:42.591381Z","iopub.status.idle":"2022-12-09T11:12:43.291477Z","shell.execute_reply.started":"2022-12-09T11:12:42.591353Z","shell.execute_reply":"2022-12-09T11:12:43.290228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Cancer/non cancer cases with wrt view**","metadata":{}},{"cell_type":"code","source":"train_df[\"view\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:43.29938Z","iopub.execute_input":"2022-12-09T11:12:43.30205Z","iopub.status.idle":"2022-12-09T11:12:43.318627Z","shell.execute_reply.started":"2022-12-09T11:12:43.301975Z","shell.execute_reply":"2022-12-09T11:12:43.317312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df[\"view\"].value_counts().plot(kind=\"bar\")","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:43.320625Z","iopub.execute_input":"2022-12-09T11:12:43.321112Z","iopub.status.idle":"2022-12-09T11:12:43.525732Z","shell.execute_reply.started":"2022-12-09T11:12:43.321075Z","shell.execute_reply":"2022-12-09T11:12:43.524708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.histplot(data=train_df,x=\"cancer\",hue=\"view\")","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:43.527092Z","iopub.execute_input":"2022-12-09T11:12:43.527551Z","iopub.status.idle":"2022-12-09T11:12:44.059243Z","shell.execute_reply.started":"2022-12-09T11:12:43.527516Z","shell.execute_reply":"2022-12-09T11:12:44.057558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Extract only cancerous cases and EDA","metadata":{}},{"cell_type":"code","source":"cancerous_df = train_df[train_df[\"cancer\"] == 1]\ncancerous_df.info()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.062384Z","iopub.execute_input":"2022-12-09T11:12:44.062671Z","iopub.status.idle":"2022-12-09T11:12:44.077905Z","shell.execute_reply.started":"2022-12-09T11:12:44.062644Z","shell.execute_reply":"2022-12-09T11:12:44.076852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**What is the r/n-ship between positive cancer cases, tissue density and age ??**\n- **density** - A rating for how dense the breast tissue is, with A being the least dense and D being the most dense. Extremely dense tissue can make diagnosis more difficult. Only provided for train.","metadata":{}},{"cell_type":"code","source":"fig,axs = plt.subplots(1,2,figsize=(20,10))\nsns.histplot(data=cancerous_df,x=\"density\",ax=axs[0])\nsns.histplot(data=cancerous_df,x=\"age\",ax=axs[1],kde=\"true\")","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.079218Z","iopub.execute_input":"2022-12-09T11:12:44.079834Z","iopub.status.idle":"2022-12-09T11:12:44.491734Z","shell.execute_reply.started":"2022-12-09T11:12:44.079765Z","shell.execute_reply":"2022-12-09T11:12:44.490744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**In depth analysis of the dataset can be found** [here](https://www.kaggle.com/competitions/rsna-breast-cancer-detection/discussion/370341)","metadata":{}},{"cell_type":"markdown","source":"## Process the dicom images","metadata":{}},{"cell_type":"code","source":"import pydicom","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.493611Z","iopub.execute_input":"2022-12-09T11:12:44.494546Z","iopub.status.idle":"2022-12-09T11:12:44.499228Z","shell.execute_reply.started":"2022-12-09T11:12:44.494506Z","shell.execute_reply":"2022-12-09T11:12:44.498028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"random_datapoints = train_df.sample(5)\nrandom_datapoints","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.500544Z","iopub.execute_input":"2022-12-09T11:12:44.501439Z","iopub.status.idle":"2022-12-09T11:12:44.524577Z","shell.execute_reply.started":"2022-12-09T11:12:44.501401Z","shell.execute_reply":"2022-12-09T11:12:44.523246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndicom_file = pydicom.read_file(f\"{train_image_path}/{str(random_datapoints.iloc[0].patient_id)}/{str(random_datapoints.iloc[0].image_id)}.dcm\")\nprint(dicom_file)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.526347Z","iopub.execute_input":"2022-12-09T11:12:44.527076Z","iopub.status.idle":"2022-12-09T11:12:44.703841Z","shell.execute_reply.started":"2022-12-09T11:12:44.527035Z","shell.execute_reply":"2022-12-09T11:12:44.702895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* [Lets use the version of the dataset wich is already processed and converted to png](https://www.kaggle.com/datasets/radek1/rsna-mammography-images-as-pngs?select=images_as_pngs_512)\n* [png and ROI with 768 pi extracted dataset](https://www.kaggle.com/datasets/remekkinas/rsna-breast-cancer-detection-poi-images)\n","metadata":{}},{"cell_type":"markdown","source":"### Evaluation metrics for the challenge is probablistic f1\n\n$ pF1=2 \\times \\frac{pPrecision \\times pRecall}{pPrecision+pRecall} $\n\nwith:\n\n$pPrecision=\\frac{pTP}{pTP+pFP}$\n\n$pRecall=\\frac{pTP}{pTP+pFN}$","metadata":{}},{"cell_type":"code","source":"# Python implementation of probablistic f1\n# https://www.kaggle.com/code/sohier/probabilistic-f-score\ndef pfbeta(labels, predictions, beta):\n    # beta = 1\n    y_true_count = 0\n    ctp = 0\n    cfp = 0\n\n    for idx in range(len(labels)):\n        prediction = min(max(predictions[idx], 0), 1)\n        if (labels[idx]):\n            y_true_count += 1\n            ctp += prediction\n        else:\n            cfp += prediction\n\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / y_true_count\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall)\n        return result\n    else:\n        return torch.tensor(0)\n    \ndef pfbeta_torch(labels, preds, beta=1):\n    preds = preds.clip(0, 1)\n    y_true_count = labels.sum()\n    ctp = preds[labels==1].sum()\n    cfp = preds[labels==0].sum()\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / y_true_count\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall)\n        return result\n    else:\n        return torch.tensor(0.0)\n","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.705409Z","iopub.execute_input":"2022-12-09T11:12:44.705798Z","iopub.status.idle":"2022-12-09T11:12:44.716028Z","shell.execute_reply.started":"2022-12-09T11:12:44.705746Z","shell.execute_reply":"2022-12-09T11:12:44.714824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Lets the get the ROI extracted png dataset","metadata":{}},{"cell_type":"code","source":"roi_dataset_path = \"/kaggle/input/rsna-breast-cancer-detection-poi-images/bc_768_roi\"\nos.listdir(roi_dataset_path)\nroi_train_path = os.path.join(roi_dataset_path,'train')\nroi_test_path = os.path.join(roi_dataset_path,'test')","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.717584Z","iopub.execute_input":"2022-12-09T11:12:44.718257Z","iopub.status.idle":"2022-12-09T11:12:44.727918Z","shell.execute_reply.started":"2022-12-09T11:12:44.718222Z","shell.execute_reply":"2022-12-09T11:12:44.726969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Lets visualize batch of the training data**","metadata":{}},{"cell_type":"code","source":"roi_train_path","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.729181Z","iopub.execute_input":"2022-12-09T11:12:44.730286Z","iopub.status.idle":"2022-12-09T11:12:44.739016Z","shell.execute_reply.started":"2022-12-09T11:12:44.730247Z","shell.execute_reply":"2022-12-09T11:12:44.737802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_batch_df = train_df.sample(6)\ntrain_batch_df","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.741427Z","iopub.execute_input":"2022-12-09T11:12:44.741837Z","iopub.status.idle":"2022-12-09T11:12:44.760985Z","shell.execute_reply.started":"2022-12-09T11:12:44.741806Z","shell.execute_reply":"2022-12-09T11:12:44.760052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.listdir(roi_train_path)[:2]","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.76237Z","iopub.execute_input":"2022-12-09T11:12:44.763114Z","iopub.status.idle":"2022-12-09T11:12:44.789685Z","shell.execute_reply.started":"2022-12-09T11:12:44.763073Z","shell.execute_reply":"2022-12-09T11:12:44.788673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig,axs = plt.subplots(1,6,figsize=(20,10))\nfor i in range(len(train_batch_df)):\n    row = train_batch_df.iloc[i]\n    image_path = os.path.join(roi_train_path,f\"{row.patient_id}_{row.image_id}.png\")\n    assert os.path.exists(image_path)\n    image = cv2.imread(image_path,cv2.IMREAD_GRAYSCALE)\n    axs[i].imshow(image,cmap=\"gray\")\n    axs[i].set_title(f\"Laterality = {row.laterality}\\n Cancer = {row.cancer}\\n{image.shape}\")","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:44.791244Z","iopub.execute_input":"2022-12-09T11:12:44.791585Z","iopub.status.idle":"2022-12-09T11:12:45.926689Z","shell.execute_reply.started":"2022-12-09T11:12:44.791553Z","shell.execute_reply":"2022-12-09T11:12:45.925645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare and train a resnet baseline on the dataset\n* First lets train on the unprocessed raw images, we will only resize the image to the dimension required by resnet.\n* We will also train our first baseline model on the undersampled version of the dataset.","metadata":{}},{"cell_type":"code","source":"count_dict = train_df.cancer.value_counts().to_dict()\ncount_dict","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:45.92822Z","iopub.execute_input":"2022-12-09T11:12:45.92864Z","iopub.status.idle":"2022-12-09T11:12:45.936869Z","shell.execute_reply.started":"2022-12-09T11:12:45.928603Z","shell.execute_reply":"2022-12-09T11:12:45.935886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# a simple function to undersample the majority class i.e cancer = 0.0\ndef under_sample(df,majority_class,minority_class,sample_amount):\n    return(pd.concat([df[df.cancer == majority_class]\n                      .sample(sample_amount)\n                      ,df[df.cancer == minority_class]]))","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:45.938915Z","iopub.execute_input":"2022-12-09T11:12:45.939627Z","iopub.status.idle":"2022-12-09T11:12:45.945889Z","shell.execute_reply.started":"2022-12-09T11:12:45.939591Z","shell.execute_reply":"2022-12-09T11:12:45.944841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_us_df = under_sample(train_df,0,1,count_dict[1])","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:45.947638Z","iopub.execute_input":"2022-12-09T11:12:45.948366Z","iopub.status.idle":"2022-12-09T11:12:45.967423Z","shell.execute_reply.started":"2022-12-09T11:12:45.948331Z","shell.execute_reply":"2022-12-09T11:12:45.966474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_us_df.cancer.value_counts().plot(kind=\"bar\")","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:45.968958Z","iopub.execute_input":"2022-12-09T11:12:45.969358Z","iopub.status.idle":"2022-12-09T11:12:46.158231Z","shell.execute_reply.started":"2022-12-09T11:12:45.969321Z","shell.execute_reply":"2022-12-09T11:12:46.156944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Over-sampling the minority class","metadata":{}},{"cell_type":"code","source":"# a simple function to undersample the majority class i.e cancer = 0.0\ndef over_sample(df,majority_class,minority_class,sample_amount):\n    return(pd.concat([df[df.cancer == minority_class]\n                      .sample(sample_amount,replace=True)\n                      ,df[df.cancer == majority_class]]))","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:46.159916Z","iopub.execute_input":"2022-12-09T11:12:46.165096Z","iopub.status.idle":"2022-12-09T11:12:46.175262Z","shell.execute_reply.started":"2022-12-09T11:12:46.165047Z","shell.execute_reply":"2022-12-09T11:12:46.173869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_os_df = over_sample(train_df,0,1,count_dict[0])\ntrain_os_df.cancer.value_counts().plot(kind=\"bar\")","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:46.18185Z","iopub.execute_input":"2022-12-09T11:12:46.184873Z","iopub.status.idle":"2022-12-09T11:12:46.413624Z","shell.execute_reply.started":"2022-12-09T11:12:46.18483Z","shell.execute_reply":"2022-12-09T11:12:46.412689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Split the over-sampled dataset into train and validation set","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nTRAIN_DF,VAL_DF = train_test_split(train_os_df,test_size=0.2)\n\n#VAL_DF = train_os_df.sample(frac=0.2,random_state=42)\n#TRAIN_DF = train_os_df.drop(VAL_DF.index)\n\nlen(TRAIN_DF),len(VAL_DF)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:46.417924Z","iopub.execute_input":"2022-12-09T11:12:46.42014Z","iopub.status.idle":"2022-12-09T11:12:46.469419Z","shell.execute_reply.started":"2022-12-09T11:12:46.420101Z","shell.execute_reply":"2022-12-09T11:12:46.468434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Create a pytorch dataset class","metadata":{}},{"cell_type":"code","source":"import torch\nimport torchvision\nfrom torch.utils.data import Dataset,DataLoader\nfrom torchvision import transforms\nimport torchvision.transforms.functional  as TF","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:46.473723Z","iopub.execute_input":"2022-12-09T11:12:46.476126Z","iopub.status.idle":"2022-12-09T11:12:46.4832Z","shell.execute_reply.started":"2022-12-09T11:12:46.476087Z","shell.execute_reply":"2022-12-09T11:12:46.482064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nclass RSNA_Dataset(Dataset):\n    def __init__(self,df,target_size=(224,244),train=True, data_path=None):\n        self.df = df\n        self.target_size = target_size\n        self.train = train\n        self.data_path = data_path\n        \n    def __len__(self):\n        return(len(self.df))\n    \n    def __getitem__(self,idx):\n        df_item = self.df.iloc[idx]\n        image_path = os.path.join(self.data_path,f\"{df_item.patient_id}_{df_item.image_id}.png\")\n        assert os.path.exists(image_path)\n        image = cv2.imread(image_path,cv2.IMREAD_GRAYSCALE)\n        class_label = int(df_item.cancer)\n        # Image to tensor\n        image = TF.to_tensor(image)\n        class_label = torch.tensor([class_label],dtype=torch.float32)\n        # Resize the image and mask to desired size\n        resize = transforms.Resize(size=self.target_size)\n        image = resize(image)\n        # Add image augmentation in training mode\n        if self.train:\n            # Randomly horizontally flip the images and masks\n            if random.random() < 0.5:\n                image = TF.hflip(image)\n            # Randomly vertically flip the images and masks\n            if random.random() < 0.5:\n                image = TF.vflip(image)\n            # Randomly rotate the images and masks\n            if random.random() < 0.5:\n                angle = random.randint(-20, 20)\n                image = TF.rotate(image, angle)\n        # Normalize image\n        image = (image - image.mean()) / image.std()\n        return image,class_label\n            \n        \n        ","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:46.488006Z","iopub.execute_input":"2022-12-09T11:12:46.491114Z","iopub.status.idle":"2022-12-09T11:12:46.50634Z","shell.execute_reply.started":"2022-12-09T11:12:46.491072Z","shell.execute_reply":"2022-12-09T11:12:46.505195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\n\n!pip install numpy==1.21.6\n!pip install pydicom==2.3.1\n!pip install pylibjpeg==1.4.0\n!pip install python_gdcm==3.0.20\n\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:12:46.511571Z","iopub.execute_input":"2022-12-09T11:12:46.514317Z","iopub.status.idle":"2022-12-09T11:13:39.319008Z","shell.execute_reply.started":"2022-12-09T11:12:46.514259Z","shell.execute_reply":"2022-12-09T11:13:39.317759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    import pylibjpeg\nexcept:\n    !rm -rf /root/.cache/torch/hub/checkpoints/\n    !mkdir -p /root/.cache/torch/hub/checkpoints/\n    !pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.whl,pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}\n    !pip install /kaggle/input/rsna-2022-whl/{torch-1.12.1-cp37-cp37m-manylinux1_x86_64.whl,torchvision-0.13.1-cp37-cp37m-manylinux1_x86_64.whl}","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:39.329092Z","iopub.execute_input":"2022-12-09T11:13:39.329413Z","iopub.status.idle":"2022-12-09T11:13:39.350269Z","shell.execute_reply.started":"2022-12-09T11:13:39.329382Z","shell.execute_reply":"2022-12-09T11:13:39.348958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check if the dataset class is working\nTRAIN_DATASET = RSNA_Dataset(TRAIN_DF,data_path=roi_train_path)\nVAL_DATASET = RSNA_Dataset(VAL_DF,train=False,data_path=roi_train_path)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:39.35193Z","iopub.execute_input":"2022-12-09T11:13:39.352541Z","iopub.status.idle":"2022-12-09T11:13:39.364063Z","shell.execute_reply.started":"2022-12-09T11:13:39.352491Z","shell.execute_reply":"2022-12-09T11:13:39.36295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Prepare the data loader","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 32\nTRAIN_LOADER = DataLoader(TRAIN_DATASET,batch_size = BATCH_SIZE,shuffle=True)\nVAL_LOADER = DataLoader(VAL_DATASET,batch_size = BATCH_SIZE)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:39.367596Z","iopub.execute_input":"2022-12-09T11:13:39.367883Z","iopub.status.idle":"2022-12-09T11:13:39.375451Z","shell.execute_reply.started":"2022-12-09T11:13:39.367857Z","shell.execute_reply":"2022-12-09T11:13:39.374349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get a sample batch\nbatch_image,batch_label = next(iter(TRAIN_LOADER))","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:39.377088Z","iopub.execute_input":"2022-12-09T11:13:39.377445Z","iopub.status.idle":"2022-12-09T11:13:39.726731Z","shell.execute_reply.started":"2022-12-09T11:13:39.377411Z","shell.execute_reply":"2022-12-09T11:13:39.725767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_image.shape,batch_label.shape","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:39.728348Z","iopub.execute_input":"2022-12-09T11:13:39.728729Z","iopub.status.idle":"2022-12-09T11:13:39.735403Z","shell.execute_reply.started":"2022-12-09T11:13:39.728689Z","shell.execute_reply":"2022-12-09T11:13:39.73426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig,axs = plt.subplots(1,8,figsize=(20,10))\nfor i,image in enumerate(batch_image):\n    image = image.permute(1,2,0).numpy()\n    axs[i].imshow(image,cmap=\"gray\")\n    if i > 6:break","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:39.737027Z","iopub.execute_input":"2022-12-09T11:13:39.737681Z","iopub.status.idle":"2022-12-09T11:13:40.596083Z","shell.execute_reply.started":"2022-12-09T11:13:39.737644Z","shell.execute_reply":"2022-12-09T11:13:40.594739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Prepare resnet18 baseline model","metadata":{}},{"cell_type":"code","source":"# replace the first Conv layer and the last fully connected layer\ndef get_model():\n    model = torchvision.models.resnet34()\n    model.conv1 = torch.nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n    in_features = model.fc.in_features\n    model.fc = torch.nn.Sequential(\n        torch.nn.Linear(in_features = in_features,out_features=1,bias=True),\n        torch.nn.Sigmoid()\n    )\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:40.599113Z","iopub.execute_input":"2022-12-09T11:13:40.599776Z","iopub.status.idle":"2022-12-09T11:13:40.60682Z","shell.execute_reply.started":"2022-12-09T11:13:40.599721Z","shell.execute_reply":"2022-12-09T11:13:40.605702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train the baseline model","metadata":{}},{"cell_type":"code","source":"model = get_model()\noptimizer = torch.optim.Adam(model.parameters(),lr=1e-3)\nbce_loss = torch.nn.BCELoss()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:40.608334Z","iopub.execute_input":"2022-12-09T11:13:40.608823Z","iopub.status.idle":"2022-12-09T11:13:40.977278Z","shell.execute_reply.started":"2022-12-09T11:13:40.608771Z","shell.execute_reply":"2022-12-09T11:13:40.976098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\ndevice","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:40.979024Z","iopub.execute_input":"2022-12-09T11:13:40.979432Z","iopub.status.idle":"2022-12-09T11:13:40.990766Z","shell.execute_reply.started":"2022-12-09T11:13:40.979392Z","shell.execute_reply":"2022-12-09T11:13:40.989712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Move model to device\nmodel = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:40.992744Z","iopub.execute_input":"2022-12-09T11:13:40.993365Z","iopub.status.idle":"2022-12-09T11:13:41.027703Z","shell.execute_reply.started":"2022-12-09T11:13:40.993329Z","shell.execute_reply":"2022-12-09T11:13:41.026684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nfrom tqdm.notebook import tqdm\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torchmetrics.classification import Accuracy\n\ndef train():\n    NUM_EPOCHS = 1\n    history = {\"train_loss\":[],\n               \"val_loss\":[],\n              \"train_pF\":[],\n              \"val_pF\":[],\n                \"train_acc\":[],\n               \"val_acc\":[],\n              \"epoch\":[]}\n    accuracy = Accuracy(task=\"binary\",threshold=0.5).to(device)\n    \n    scheduler = ReduceLROnPlateau(optimizer,\"min\",patience=5)\n    for epoch in range(NUM_EPOCHS):\n        model.train()\n        epoch_train_loss = []\n        epoch_train_pF_score = []\n        epoch_val_loss = []\n        epoch_val_pF_score = []\n        epoch_train_acc = []\n        epoch_val_acc = []\n        \n        progress_bar = tqdm(enumerate(TRAIN_LOADER),total=len(TRAIN_LOADER))\n        \n        for i,(images,labels) in progress_bar:\n            # Move the images and masks to device\n            images = images.to(device)\n            labels = labels.to(device)\n            # Perform forward pass and calculate the loss\n            preds = model(images) \n\n            # calculate BCEloss and pFscore\n            #preds = preds.squeeze(0)\n            loss = bce_loss(preds,labels)\n            pFscore = pfbeta(preds.detach().cpu(),labels.detach().cpu(),1)\n            acc = accuracy(preds.detach(),labels.detach())\n\n\n            # zero grand and back-propagate\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n\n            epoch_train_loss.append(loss.detach().data.item())\n            epoch_train_pF_score.append(pFscore.data.item())\n            epoch_train_acc.append(acc.data.item())\n\n            progress_bar.set_description(f\"Epoch {epoch+1}/{NUM_EPOCHS}\")\n            progress_bar.set_postfix(loss=np.array(epoch_train_loss).mean(),\n                                    pFscore=np.array(epoch_train_pF_score).mean(),\n                                    train_acc = np.array(epoch_train_acc).mean())\n            progress_bar.update()\n        if True:\n            with torch.no_grad():\n                model.eval()\n                eval_bar = tqdm(enumerate(VAL_LOADER),total=len(VAL_LOADER))\n                for i,(images,labels) in eval_bar:\n                    images = images.to(device)\n                    labels = labels.to(device)\n                    # Perform forward pass and calculate the loss\n                    preds = model(images) \n\n                    # calculate BCEloss and pFscore\n\n                    loss = bce_loss(preds,labels)\n                    pFscore = pfbeta(preds.detach().cpu(),labels.detach().cpu(),1)\n                    \n                    acc = accuracy(preds.detach(),labels.detach())\n\n                    epoch_val_loss.append(loss.detach().data.item())\n                    epoch_val_pF_score.append(pFscore.data.item())\n                    epoch_val_acc.append(acc.data.item())\n\n                    eval_bar.set_description(f\"Evaluating\")\n                    eval_bar.set_postfix(val_loss=np.array(epoch_val_loss).mean(),\n                                        val_pFscore=np.array(epoch_val_pF_score).mean(),\n                                        val_acc = np.array(epoch_val_acc).mean())\n                    eval_bar.update()\n            history[\"train_loss\"].append(np.array(epoch_train_loss).mean())\n            history[\"val_loss\"].append(np.array(epoch_val_loss).mean())\n            history[\"train_pF\"].append(np.array(epoch_train_pF_score).mean())\n            history[\"val_pF\"].append(np.array(epoch_val_pF_score).mean())\n            history[\"train_acc\"].append(np.array(epoch_train_acc).mean())\n            history[\"val_acc\"].append(np.array(epoch_val_acc).mean())\n            history[\"epoch\"].append(epoch)\n\n            scheduler.step(np.array(epoch_val_loss).mean())\n\n            # Save model\n            print(f\"[INFO] Saving model ... /kaggle/working/resnet_34_baseline.pth\")\n            torch.save(model.state_dict(),\"/kaggle/working/resnet_34_baseline.pth\")\n            # Save history\n            with open(\"history.json\", \"w\") as outfile:\n                json.dump(history, outfile)            ","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:41.029527Z","iopub.execute_input":"2022-12-09T11:13:41.029967Z","iopub.status.idle":"2022-12-09T11:13:41.050731Z","shell.execute_reply.started":"2022-12-09T11:13:41.029927Z","shell.execute_reply":"2022-12-09T11:13:41.049585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# start train\n#train()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:13:41.052418Z","iopub.execute_input":"2022-12-09T11:13:41.052904Z","iopub.status.idle":"2022-12-09T11:36:52.634245Z","shell.execute_reply.started":"2022-12-09T11:13:41.052864Z","shell.execute_reply":"2022-12-09T11:36:52.633155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_history(json_path):\n    history = json.loads(json_path)\n    fig,axs = plt.subplots(1,3,figsize=(20,5))\n    sns.lineplot(data=history,x=\"epoch\",y=\"train_loss\",label=\"Train\",ax=axs[0])\n    sns.lineplot(data=history,x=\"epoch\",y=\"val_loss\",label=\"Val\",ax=axs[0])\n\n    sns.lineplot(data=history,x=\"epoch\",y=\"train_pF\",label=\"Train\",ax=axs[1])\n    sns.lineplot(data=history,x=\"epoch\",y=\"val_pF\",label=\"Val\",ax=axs[1])\n\n    sns.lineplot(data=history,x=\"epoch\",y=\"train_acc\",label=\"Train\",ax=axs[2])\n    sns.lineplot(data=history,x=\"epoch\",y=\"val_acc\",label=\"Val\",ax=axs[2])","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:52.635998Z","iopub.execute_input":"2022-12-09T11:36:52.63639Z","iopub.status.idle":"2022-12-09T11:36:52.643758Z","shell.execute_reply.started":"2022-12-09T11:36:52.636353Z","shell.execute_reply":"2022-12-09T11:36:52.642829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def evaluate(weight_path):\n    model = get_model()\n    with torch.no_grad():\n        true_classes = []\n        pred_classes = []\n        model.eval()\n        eval_bar = tqdm(enumerate(VAL_LOADER),total=len(VAL_LOADER))\n        for i,(images,labels) in eval_bar:\n            images = images.to(device)\n            labels = labels.to(device)\n            # Perform forward pass and calculate the loss\n            preds = model(images)\n            preds_classes = (preds >= 0.5).float()\n            pred_classes += preds_classes.squeeze(1).tolist()\n            true_classes += labels.squeeze(1).tolist()    \n    return (predt_classes,true_classes)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:52.644971Z","iopub.execute_input":"2022-12-09T11:36:52.645978Z","iopub.status.idle":"2022-12-09T11:36:52.660575Z","shell.execute_reply.started":"2022-12-09T11:36:52.645943Z","shell.execute_reply":"2022-12-09T11:36:52.659603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#from sklearn.metrics import confusion_matrix,classification_report,roc_auc_score\n#print(classification_report(true_classes,pred_classes))","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:52.661653Z","iopub.execute_input":"2022-12-09T11:36:52.66194Z","iopub.status.idle":"2022-12-09T11:36:52.670867Z","shell.execute_reply.started":"2022-12-09T11:36:52.661915Z","shell.execute_reply":"2022-12-09T11:36:52.669933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\ncm = confusion_matrix(true_classes,pred_classes)\nsns.heatmap(cm,annot=True)\n\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:52.672025Z","iopub.execute_input":"2022-12-09T11:36:52.672311Z","iopub.status.idle":"2022-12-09T11:36:52.682534Z","shell.execute_reply.started":"2022-12-09T11:36:52.672287Z","shell.execute_reply":"2022-12-09T11:36:52.680815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Test images\nos.listdir(test_image_path)\n# Test df\ntest_df = pd.read_csv(csv_files[-1])\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:52.684532Z","iopub.execute_input":"2022-12-09T11:36:52.684982Z","iopub.status.idle":"2022-12-09T11:36:52.705841Z","shell.execute_reply.started":"2022-12-09T11:36:52.684946Z","shell.execute_reply":"2022-12-09T11:36:52.704838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inferance and test submission","metadata":{}},{"cell_type":"code","source":"data_path = \"/kaggle/input/rsna-breast-cancer-detection\"\ntest_image_path = os.path.join(data_path,\"test_images\")\ntest_csv = os.path.join(data_path,\"test.csv\")\nos.path.exists(test_image_path),os.path.exists(test_csv)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:52.707239Z","iopub.execute_input":"2022-12-09T11:36:52.707647Z","iopub.status.idle":"2022-12-09T11:36:52.717752Z","shell.execute_reply.started":"2022-12-09T11:36:52.707614Z","shell.execute_reply":"2022-12-09T11:36:52.716229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(test_csv)\nlen(test_df)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:52.719276Z","iopub.execute_input":"2022-12-09T11:36:52.719624Z","iopub.status.idle":"2022-12-09T11:36:52.73129Z","shell.execute_reply.started":"2022-12-09T11:36:52.719587Z","shell.execute_reply":"2022-12-09T11:36:52.730442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Conv test dicom images to png","metadata":{}},{"cell_type":"code","source":"from multiprocessing import Process\nfrom scipy import ndimage","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:52.732335Z","iopub.execute_input":"2022-12-09T11:36:52.732934Z","iopub.status.idle":"2022-12-09T11:36:52.737602Z","shell.execute_reply.started":"2022-12-09T11:36:52.732899Z","shell.execute_reply":"2022-12-09T11:36:52.736617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def crop_roi(image):\n    \"\"\" Crop ROI \"\"\"\n    if(len(image.shape)==3):\n        image = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)\n    (ori_h,ori_w)=image.shape                   \n    \n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (9,9))\n    image = cv2.dilate(image, kernel, iterations=3)\n    \n    treshIm=np.uint8((np.int32(image)>10)*255)\n\n    breast_mask=ndimage.binary_fill_holes(treshIm).astype(np.uint8)\n    labeled, nr_objects = ndimage.label(breast_mask)\n    sizes = ndimage.sum(breast_mask, labeled, range(nr_objects + 1))\n    breast_mask=(labeled==np.argmax(sizes)).astype(np.uint8())\n    coords = cv2.findNonZero(breast_mask)\n    off_x,off_y,new_w,new_h = cv2.boundingRect(coords)\n    off_x=max(0,off_x-10)\n    off_y=max(0,off_y-10)\n    new_w=min(new_w+10,ori_w)\n    new_h=min(new_h+10,ori_h)\n    new_img=image[off_y:(off_y+new_h),off_x:(off_x+new_w)]   \n    return new_img\n\ndef dicom_to_png(dcm_file):\n    dicom = pydicom.dcmread(dcm_file)\n    img = dicom.pixel_array\n    img = (img - img.min()) / (img.max() - img.min())\n    if dicom.PhotometricInterpretation == 'MONOCHROME1':\n        img = 1 - img\n    roi_img = crop_roi(img)\n    #img = cv2.resize(img, image_size, interpolation=cv2.INTER_LINEAR)\n    roi_img = (roi_img * 255).astype(np.uint8)\n    \n    return roi_img\n\n#dcm_file = test_image_path + '/' + test_df.patient_id.astype(str) + '/'  + test_df.image_id.astype(str) + '.dcm'\n","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:52.739045Z","iopub.execute_input":"2022-12-09T11:36:52.740121Z","iopub.status.idle":"2022-12-09T11:36:52.752324Z","shell.execute_reply.started":"2022-12-09T11:36:52.740095Z","shell.execute_reply":"2022-12-09T11:36:52.751285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TEST_RSNA_Dataset(Dataset):\n    def __init__(self,df,target_size=(224,244),train=True, data_path=None):\n        self.df = df\n        self.target_size = target_size\n        self.train = train\n        self.data_path = data_path\n        \n    def __len__(self):\n        return(len(self.df))\n    \n    def __getitem__(self,idx):\n        df_item = self.df.iloc[idx]\n        image_path = os.path.join(self.data_path,f\"{df_item.patient_id}/{df_item.image_id}.dcm\")\n        assert os.path.exists(image_path)\n        image = dicom_to_png(image_path)\n        \n        # Image to tensor\n        image = TF.to_tensor(image)\n        # Resize the image and mask to desired size\n        resize = transforms.Resize(size=self.target_size)\n        image = resize(image)\n        image = (image - image.mean()) / image.std()\n        return image","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:52.75446Z","iopub.execute_input":"2022-12-09T11:36:52.755471Z","iopub.status.idle":"2022-12-09T11:36:52.764101Z","shell.execute_reply.started":"2022-12-09T11:36:52.755397Z","shell.execute_reply":"2022-12-09T11:36:52.763179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TEST_DATASET = TEST_RSNA_Dataset(test_df,data_path=test_image_path)\nTEST_LOADER = DataLoader(TEST_DATASET,batch_size = 16)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T12:17:08.176406Z","iopub.execute_input":"2022-12-09T12:17:08.176801Z","iopub.status.idle":"2022-12-09T12:17:08.181524Z","shell.execute_reply.started":"2022-12-09T12:17:08.176748Z","shell.execute_reply":"2022-12-09T12:17:08.180568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = get_model()\nmodel.load_state_dict(torch.load(\"/kaggle/input/resnet-14-checkpoint/resnet_34_baseline(1).pth\",map_location=device))\nmodel = model.to(device)\n\npredictions = []\nwith torch.no_grad():\n    true_classes = []\n    pred_classes = []\n    model.eval()\n    eval_bar = tqdm(enumerate(TEST_LOADER),total=len(TEST_LOADER))\n    for i,images in eval_bar:\n        images = images.to(device)\n        # Perform forward pass and calculate the loss\n        preds = model(images)\n        predictions+=preds.detach().cpu().view(-1).tolist()\n","metadata":{"execution":{"iopub.status.busy":"2022-12-09T12:23:06.117918Z","iopub.execute_input":"2022-12-09T12:23:06.118375Z","iopub.status.idle":"2022-12-09T12:23:12.661201Z","shell.execute_reply.started":"2022-12-09T12:23:06.118332Z","shell.execute_reply":"2022-12-09T12:23:12.660102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame(data={'prediction_id': test_df['prediction_id'], 'cancer': np.array(predictions)}).drop_duplicates(subset='prediction_id')","metadata":{"execution":{"iopub.status.busy":"2022-12-09T12:23:20.16382Z","iopub.execute_input":"2022-12-09T12:23:20.164724Z","iopub.status.idle":"2022-12-09T12:23:20.172539Z","shell.execute_reply.started":"2022-12-09T12:23:20.164656Z","shell.execute_reply":"2022-12-09T12:23:20.171544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T11:36:59.861692Z","iopub.status.idle":"2022-12-09T11:36:59.862474Z","shell.execute_reply.started":"2022-12-09T11:36:59.862228Z","shell.execute_reply":"2022-12-09T11:36:59.862253Z"},"trusted":true},"execution_count":null,"outputs":[]}]}