{"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":"## import","metadata":{}},{"cell_type":"code","source":"!pip install dicomsdl pytorch_lightning timm --no-index --find-links=../input/rbcd-downloads","metadata":{"execution":{"iopub.status.busy":"2023-02-16T05:24:06.25148Z","iopub.execute_input":"2023-02-16T05:24:06.252444Z","iopub.status.idle":"2023-02-16T05:24:19.98151Z","shell.execute_reply.started":"2023-02-16T05:24:06.252331Z","shell.execute_reply":"2023-02-16T05:24:19.980363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm\nimport cv2\nimport os\nimport dicomsdl\nimport random\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import transforms\n# from torchvision.models import resnet50, googlenet\n\nimport timm","metadata":{"execution":{"iopub.status.busy":"2023-02-16T05:24:19.984263Z","iopub.execute_input":"2023-02-16T05:24:19.984688Z","iopub.status.idle":"2023-02-16T05:24:23.443937Z","shell.execute_reply.started":"2023-02-16T05:24:19.984647Z","shell.execute_reply":"2023-02-16T05:24:23.442828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('torch: ', torch.__version__, ' timm: ', timm.__version__)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T05:24:23.445715Z","iopub.execute_input":"2023-02-16T05:24:23.446353Z","iopub.status.idle":"2023-02-16T05:24:23.4585Z","shell.execute_reply.started":"2023-02-16T05:24:23.446315Z","shell.execute_reply":"2023-02-16T05:24:23.457272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"random.seed(0)\nnp.random.seed(0)\ntorch.manual_seed(0)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T05:24:23.460304Z","iopub.execute_input":"2023-02-16T05:24:23.461837Z","iopub.status.idle":"2023-02-16T05:24:23.481651Z","shell.execute_reply.started":"2023-02-16T05:24:23.4618Z","shell.execute_reply":"2023-02-16T05:24:23.480709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_PATH = \"/kaggle/input/rsna-mammography-images-as-pngs/images_as_pngs_512/train_images_processed_512\"\nTEST_PATH = \"/kaggle/input/rsna-breast-cancer-detection/test_images\"","metadata":{"execution":{"iopub.status.busy":"2023-02-16T05:24:23.48725Z","iopub.execute_input":"2023-02-16T05:24:23.488217Z","iopub.status.idle":"2023-02-16T05:24:23.495227Z","shell.execute_reply.started":"2023-02-16T05:24:23.488175Z","shell.execute_reply":"2023-02-16T05:24:23.493972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## train dataset","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(\"/kaggle/input/rsna-breast-cancer-detection/train.csv\")\n# add image filename columns\ndf[\"img_name\"] = df[\"patient_id\"].astype(str) + \"/\" + df[\"image_id\"].astype(str) + \".png\"\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:05.769453Z","iopub.execute_input":"2023-02-16T06:46:05.770535Z","iopub.status.idle":"2023-02-16T06:46:05.900193Z","shell.execute_reply.started":"2023-02-16T06:46:05.770495Z","shell.execute_reply":"2023-02-16T06:46:05.899097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[\"cancer\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:05.90233Z","iopub.execute_input":"2023-02-16T06:46:05.902725Z","iopub.status.idle":"2023-02-16T06:46:05.911321Z","shell.execute_reply.started":"2023-02-16T06:46:05.902688Z","shell.execute_reply":"2023-02-16T06:46:05.9102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# shuffle it\ndf = df.sample(frac=1).reset_index(drop=True)\n\n# undersample according to the cancer patients since they are minority\nrate = 10\nundersample_amount = int(rate*len(df[df[\"cancer\"]==1]))\n\ndfnotcancer = df[df[\"cancer\"]==0].sample(undersample_amount).reset_index(drop=True)\ndfcancer = df[df[\"cancer\"]==1].reset_index(drop=True)\n\n# concat and then shuffle, reset index\ndff = pd.concat([dfcancer, dfnotcancer]).sample(frac=1).reset_index(drop=True)\n\nprint(f\"Old data shape is {df.shape} and new data shape is: {dff.shape}\")\n\ndff.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:05.91305Z","iopub.execute_input":"2023-02-16T06:46:05.913593Z","iopub.status.idle":"2023-02-16T06:46:05.975281Z","shell.execute_reply.started":"2023-02-16T06:46:05.913545Z","shell.execute_reply":"2023-02-16T06:46:05.974092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dff[\"cancer\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:05.978051Z","iopub.execute_input":"2023-02-16T06:46:05.978951Z","iopub.status.idle":"2023-02-16T06:46:05.98887Z","shell.execute_reply.started":"2023-02-16T06:46:05.978912Z","shell.execute_reply":"2023-02-16T06:46:05.98761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def crop_coords(img):\n    \"\"\"\n    Crop ROI from image.\n    \"\"\"\n    # Otsu's thresholding after Gaussian filtering\n    blur = cv2.GaussianBlur(img, (5, 5), 0)\n    _, breast_mask = cv2.threshold(blur, 0, 255, cv2.THRESH_BINARY+cv2.THRESH_OTSU)\n    cnts, _ = cv2.findContours(breast_mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    cnt = max(cnts, key = cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(cnt)\n    return (x, y, w, h)\n\ndef truncation_normalization(img):\n    \"\"\"\n    Clip and normalize pixels in the breast ROI.\n    @img : numpy array image\n    return: numpy array of the normalized image\n    \"\"\"\n    Pmin = np.percentile(img[img!=0], 5)\n    Pmax = np.percentile(img[img!=0], 99)\n    truncated = np.clip(img,Pmin, Pmax)  \n    normalized = (truncated - Pmin)/(Pmax - Pmin)\n    normalized[img==0] = 0\n    return normalized\n\ndef clahe(img, clip):\n    \"\"\"\n    Image enhancement.\n    @img : numpy array image\n    @clip : float, clip limit for CLAHE algorithm\n    return: numpy array of the enhanced image\n    \"\"\"\n    clahe = cv2.createCLAHE(clipLimit=clip)\n    cl = clahe.apply(np.array(img*255, dtype=np.uint8))\n    return cl\n\ndef img2roi(img, is_dicom=False, is_bgr=False):\n    \"\"\"\n    Returns ROI area in other words \n    cuts the image to a desired one\n    \n    Because there are machine label tags,\n    undesired details out of the breast image.\n    \"\"\"\n    if is_dicom:\n        img = np.array(img * 255, dtype = np.uint8)\n    elif is_bgr:\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n\n    (x, y, w, h) = crop_coords(img)\n    img_cropped = img[y:y+h, x:x+w]\n    img_normalized = truncation_normalization(img_cropped)\n    cl1 = clahe(img_normalized, 1.0)\n    cl2 = clahe(img_normalized, 2.0)\n    img_final = cv2.merge((np.array(img_normalized*255, dtype=np.uint8), cl1, cl2))\n\n    return img_final\n\ndef img2fft(img):\n    f = np.fft.fft2(img)\n    fshift = np.fft.fftshift(f)\n    result = np.log(np.abs(fshift))\n    return result","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:05.991323Z","iopub.execute_input":"2023-02-16T06:46:05.991836Z","iopub.status.idle":"2023-02-16T06:46:06.006051Z","shell.execute_reply.started":"2023-02-16T06:46:05.991797Z","shell.execute_reply":"2023-02-16T06:46:06.005104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = cv2.imread(TRAIN_PATH+\"/\"+dff.img_name[0], cv2.IMREAD_GRAYSCALE)\nroi = img2roi(img, is_bgr=False)\nprint(roi.shape)\nplt.imshow(roi, cmap=\"bone\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:06.007166Z","iopub.execute_input":"2023-02-16T06:46:06.008159Z","iopub.status.idle":"2023-02-16T06:46:06.500518Z","shell.execute_reply.started":"2023-02-16T06:46:06.008121Z","shell.execute_reply":"2023-02-16T06:46:06.499571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = cv2.imread(TRAIN_PATH+\"/\"+dff.img_name[0], cv2.IMREAD_GRAYSCALE)\nimg = cv2.resize(img, (224, 224))\nres = img2fft(img)\nprint(img.shape, res.shape)\nplt.subplot(121)\nplt.imshow(img, 'gray')\nplt.title('Original Image')\nplt.axis('off')\nplt.subplot(122)\nplt.imshow(res, 'gray')\nplt.title('Fourier Image')\nplt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:06.501988Z","iopub.execute_input":"2023-02-16T06:46:06.506141Z","iopub.status.idle":"2023-02-16T06:46:06.701895Z","shell.execute_reply.started":"2023-02-16T06:46:06.506102Z","shell.execute_reply":"2023-02-16T06:46:06.700936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = cv2.imread(TRAIN_PATH+\"/\"+dff.img_name[0], cv2.IMREAD_GRAYSCALE)\n(x, y, w, h) = crop_coords(img)\nimg = img[y:y+h, x:x+w]\nimg = cv2.resize(img, (224, 224))\nimg_normalized = truncation_normalization(img)\nf = np.fft.fft2(img)\nfshift = np.fft.fftshift(f)\nres = np.log(np.abs(fshift))\nstack = np.array([img, img_normalized, res])\nprint(stack.shape)\nplt.imshow(stack.transpose((1, 2, 0)))\nplt.title('Stack Image')\nplt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:06.703304Z","iopub.execute_input":"2023-02-16T06:46:06.704171Z","iopub.status.idle":"2023-02-16T06:46:06.857236Z","shell.execute_reply.started":"2023-02-16T06:46:06.704141Z","shell.execute_reply":"2023-02-16T06:46:06.856223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## test dataset","metadata":{}},{"cell_type":"code","source":"df_test = pd.read_csv(\"/kaggle/input/rsna-breast-cancer-detection/test.csv\")\ndf_test[\"dcm_path\"] = df_test[\"patient_id\"].astype(str) + \"/\" + df_test[\"image_id\"].astype(str) + \".dcm\"\ndf_test.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:06.858745Z","iopub.execute_input":"2023-02-16T06:46:06.859131Z","iopub.status.idle":"2023-02-16T06:46:06.882929Z","shell.execute_reply.started":"2023-02-16T06:46:06.859089Z","shell.execute_reply":"2023-02-16T06:46:06.881821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dicom2img(path):\n#     dicom = pydicom.dcmread(path)\n#     img = dicom.pixel_array\n    dicom = dicomsdl.open(path)\n    img = dicom.pixelData()\n    img = (img - img.min()) / (img.max() - img.min())\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1.0 - img\n    # img = (img * 255).astype(np.uint8)\n    return img","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:06.887377Z","iopub.execute_input":"2023-02-16T06:46:06.888886Z","iopub.status.idle":"2023-02-16T06:46:06.894622Z","shell.execute_reply.started":"2023-02-16T06:46:06.888849Z","shell.execute_reply":"2023-02-16T06:46:06.893825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = dicom2img(os.path.join(TEST_PATH, df_test[\"dcm_path\"][0]))\nroi = img2roi(img, is_dicom=True)\nbgr = cv2.cvtColor(np.array(img * 255, dtype = np.uint8), cv2.COLOR_GRAY2BGR)\nprint(img.shape, roi.shape, bgr.shape)\nplt.subplot(131)\nplt.imshow(img, cmap=\"bone\")\nplt.title('Original Image')\nplt.axis('off')\nplt.subplot(132)\nplt.imshow(roi, cmap=\"bone\")\nplt.title('ROI Image')\nplt.axis('off')\nplt.subplot(133)\nplt.imshow(bgr, cmap=\"bone\")\nplt.title('bgr Image')\nplt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:06.896141Z","iopub.execute_input":"2023-02-16T06:46:06.897041Z","iopub.status.idle":"2023-02-16T06:46:09.073437Z","shell.execute_reply.started":"2023-02-16T06:46:06.896999Z","shell.execute_reply":"2023-02-16T06:46:09.072538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = dicom2img(os.path.join(TEST_PATH, df_test[\"dcm_path\"][0]))\nimg = np.array(img * 255, dtype = np.uint8)\nimg = cv2.resize(img, (224, 224))\nf = np.fft.fft2(img)\nfshift = np.fft.fftshift(f)\nres = np.log(np.abs(fshift))\n# res = np.array([res, res, res])\n# res = res.transpose((1, 2, 0))\nprint(img.shape, res.shape, '\\n', res)\nplt.subplot(121)\nplt.imshow(img, 'gray')\nplt.title('Original Image')\nplt.axis('off')\nplt.subplot(122)\nplt.imshow(res, 'gray')\nplt.title('Fourier Image')\nplt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:09.075243Z","iopub.execute_input":"2023-02-16T06:46:09.076248Z","iopub.status.idle":"2023-02-16T06:46:09.901667Z","shell.execute_reply.started":"2023-02-16T06:46:09.076208Z","shell.execute_reply":"2023-02-16T06:46:09.900563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = dicom2img(os.path.join(TEST_PATH, df_test[\"dcm_path\"][0]))\nimg = np.array(img * 255, dtype = np.uint8)\n(x, y, w, h) = crop_coords(img)\nimg = img[y:y+h, x:x+w]\nimg = cv2.resize(img, (224, 224))\nimg_normalized = truncation_normalization(img)\nf = np.fft.fft2(img)\nfshift = np.fft.fftshift(f)\nres = np.log(np.abs(fshift))\nstack = np.array([img, img_normalized, res])\nprint(stack.shape)\nplt.imshow(stack.transpose((1, 2, 0)))\nplt.title('Stack Image')\nplt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:09.903452Z","iopub.execute_input":"2023-02-16T06:46:09.904086Z","iopub.status.idle":"2023-02-16T06:46:10.723188Z","shell.execute_reply.started":"2023-02-16T06:46:09.904047Z","shell.execute_reply":"2023-02-16T06:46:10.722188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## building dataset and dataloader","metadata":{}},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    def __init__(self, df, img_folder, transform=None, is_test=False, size=(224, 224), img_mode=None):\n        self.df = df\n        self.img_folder = img_folder\n        self.transform = transform\n        self.is_test = is_test\n        self.size = size\n        self.img_mode = img_mode  # roi fft all\n\n    def __getitem__(self, idx):\n        if self.is_test:\n            dcm_path = os.path.join(self.img_folder, self.df[\"dcm_path\"][idx])\n            img = dicom2img(dcm_path)\n            if self.img_mode == 'roi':\n                img = img2roi(img, is_dicom=True)\n                img = cv2.resize(img, self.size)\n            elif self.img_mode == 'fft':\n                img = np.array(img * 255)\n                img = cv2.resize(img, self.size)\n                img = img2fft(img)\n                normal_img = (img - np.min(img)) / (np.max(img) - np.min(img))\n                img = (normal_img*255).astype(np.uint8)\n                # img = np.array([img, img, img])\n                # img = img.transpose((1, 2, 0))\n            elif self.img_mode == 'all':\n                img = np.array(img * 255, dtype = np.uint8)\n                (x, y, w, h) = crop_coords(img)\n                img = img[y:y+h, x:x+w]\n                img = cv2.resize(img, self.size)\n                img_normalized = truncation_normalization(img)\n                f = np.fft.fft2(img)\n                fshift = np.fft.fftshift(f)\n                res = np.log(np.abs(fshift))\n                img = np.array([img, img_normalized, res]).transpose((1, 2, 0))\n            else:\n                img = (img * 255).astype(np.uint8)\n                img = cv2.resize(img, self.size)\n                img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)\n\n        else:\n            img_path = os.path.join(self.img_folder, self.df[\"img_name\"][idx])\n            if self.img_mode == 'roi':\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                img = img2roi(img, is_dicom=False)\n                img = cv2.resize(img, self.size)\n            elif self.img_mode == 'fft':\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                img = cv2.resize(img, self.size)\n                img = img2fft(img)\n                normal_img = (img - np.min(img)) / (np.max(img) - np.min(img))\n                img = (normal_img*255).astype(np.uint8)\n                # img = np.array([img, img, img])\n                # img = img.transpose((1, 2, 0))\n            elif self.img_mode == 'all':\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                (x, y, w, h) = crop_coords(img)\n                img = img[y:y+h, x:x+w]\n                img = cv2.resize(img, self.size)\n                img_normalized = truncation_normalization(img)\n                f = np.fft.fft2(img)\n                fshift = np.fft.fftshift(f)\n                res = np.log(np.abs(fshift))\n                img = np.array([img, img_normalized, res]).transpose((1, 2, 0))\n            else:\n                img = cv2.imread(img_path)\n                img = cv2.resize(img, self.size)\n\n        if self.transform is not None:\n            img = self.transform(img)\n            img = torch.as_tensor(img, dtype=torch.float)\n        else:\n            img = torch.tensor(img, dtype=torch.float)\n\n        #img = img.permute(2, 1, 0)\n        if not self.is_test:\n            target = self.df[\"cancer\"][idx]\n            target = torch.tensor(target, dtype=torch.long)\n            return img, target\n\n        # img = img.unsqueeze(0)\n        return img\n\n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:10.725412Z","iopub.execute_input":"2023-02-16T06:46:10.726327Z","iopub.status.idle":"2023-02-16T06:46:10.747413Z","shell.execute_reply.started":"2023-02-16T06:46:10.726286Z","shell.execute_reply":"2023-02-16T06:46:10.746322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.2),\n            transforms.GaussianBlur(kernel_size=3, sigma=(0.1, 2.0)),\n            transforms.RandomVerticalFlip(),\n            transforms.RandomEqualize(),\n            transforms.ToTensor()\n])\n\nval_transform = transforms.Compose([\n            # transforms.ToPILImage(),\n            # transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.2),\n            # transforms.GaussianBlur(kernel_size=3, sigma=(0.1, 2.0)),\n            # transforms.RandomVerticalFlip(),\n            # transforms.RandomEqualize(),\n            transforms.ToTensor()\n])","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:10.749375Z","iopub.execute_input":"2023-02-16T06:46:10.749979Z","iopub.status.idle":"2023-02-16T06:46:10.762561Z","shell.execute_reply.started":"2023-02-16T06:46:10.749934Z","shell.execute_reply":"2023-02-16T06:46:10.76167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = dff.sample(frac=1).reset_index(drop=True)\ndf[\"cancer\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:10.764339Z","iopub.execute_input":"2023-02-16T06:46:10.764801Z","iopub.status.idle":"2023-02-16T06:46:10.78553Z","shell.execute_reply.started":"2023-02-16T06:46:10.764764Z","shell.execute_reply":"2023-02-16T06:46:10.784845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"split = 0.9\ntrain_samples = int(len(df) * split)\ntrain_df = df[:train_samples+1].reset_index(drop=True)\nval_df = df[train_samples:].reset_index(drop=True)\nprint(train_df[\"cancer\"].value_counts(), '\\n', val_df[\"cancer\"].value_counts())","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:10.786945Z","iopub.execute_input":"2023-02-16T06:46:10.78753Z","iopub.status.idle":"2023-02-16T06:46:10.798746Z","shell.execute_reply.started":"2023-02-16T06:46:10.787492Z","shell.execute_reply":"2023-02-16T06:46:10.797666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = RSNADataset(\n    df=train_df, img_folder=TRAIN_PATH,\n    transform=train_transform, is_test=False, size=(224, 224),\n    #img_mode='fft'\n)\nval_dataset   = RSNADataset(\n    df=val_df, img_folder=TRAIN_PATH,\n    transform=val_transform, is_test=False, size=(224, 224),\n    #img_mode='fft'\n)\n\ntrain_loader  = DataLoader(train_dataset, batch_size=32, shuffle=True)\nval_loader    = DataLoader(val_dataset, batch_size=32, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:10.800129Z","iopub.execute_input":"2023-02-16T06:46:10.80172Z","iopub.status.idle":"2023-02-16T06:46:10.808452Z","shell.execute_reply.started":"2023-02-16T06:46:10.801692Z","shell.execute_reply":"2023-02-16T06:46:10.80757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_aug(inputs, targets=None, nrows=4, ncols=4):\n    plt.figure(figsize=(10, 10))\n    plt.subplots_adjust(wspace=0.2, hspace=0.2)\n    i_ = 0\n\n    for idx in range(len(inputs)):\n        img = inputs[idx].numpy().astype(np.float32)\n        # img = img[0,:,:]\n        plt.subplot(nrows, ncols, i_+1)\n        if targets is not None:\n            plt.title(f\"Label: {targets[idx].item()}\")\n        # plt.imshow(img, cmap=\"bone\"); \n        plt.imshow(img.transpose((1, 2, 0)))\n        plt.axis('off')\n        i_ += 1\n\n    return plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:10.810818Z","iopub.execute_input":"2023-02-16T06:46:10.811702Z","iopub.status.idle":"2023-02-16T06:46:10.819369Z","shell.execute_reply.started":"2023-02-16T06:46:10.811662Z","shell.execute_reply":"2023-02-16T06:46:10.818638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, targets = next(iter(train_loader))\nprint(images.shape, targets.shape)\nshow_aug(images, targets, nrows=4, ncols=8)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:10.820904Z","iopub.execute_input":"2023-02-16T06:46:10.82159Z","iopub.status.idle":"2023-02-16T06:46:12.935545Z","shell.execute_reply.started":"2023-02-16T06:46:10.821554Z","shell.execute_reply":"2023-02-16T06:46:12.934567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, targets = next(iter(val_loader))\nprint(images.shape, targets.shape)\nshow_aug(images, targets, nrows=4, ncols=8)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:12.937276Z","iopub.execute_input":"2023-02-16T06:46:12.938101Z","iopub.status.idle":"2023-02-16T06:46:15.228272Z","shell.execute_reply.started":"2023-02-16T06:46:12.938062Z","shell.execute_reply":"2023-02-16T06:46:15.227144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model","metadata":{}},{"cell_type":"code","source":"# model = timm.create_model('seresnext50_32x4d', num_classes=1, pretrained=False, in_chans=1)\n# model = timm.create_model('seresnext50_32x4d', num_classes=1, pretrained=False, in_chans=3)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:15.230166Z","iopub.execute_input":"2023-02-16T06:46:15.230559Z","iopub.status.idle":"2023-02-16T06:46:15.641113Z","shell.execute_reply.started":"2023-02-16T06:46:15.230518Z","shell.execute_reply":"2023-02-16T06:46:15.640099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def gem(x, p=3, eps=1e-6):\n    return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1.0 / p)\n\n\nclass GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6, p_trainable=False):\n        super(GeM, self).__init__()\n        if p_trainable:\n            self.p = nn.Parameter(torch.ones(1) * p)\n        else:\n            self.p = p\n        self.eps = eps\n\n    def forward(self, x):\n        ret = gem(x, p=self.p, eps=self.eps)\n        return ret\n\n    def __repr__(self):\n        if type(self.p) is int:\n            return (self.__class__.__name__  + f\"(p={self.p:.4f},eps={self.eps})\")\n        return (self.__class__.__name__  + f\"(p={self.p.data.tolist()[0]:.4f},eps={self.eps})\")\n        \n\nclass Net(nn.Module):\n    def __init__(self, backbone='seresnext50_32x4d', pretrained=False, in_chans=3, n_classes=1):\n        super(Net, self).__init__()\n\n        self.backbone = timm.create_model(\n            backbone, \n            pretrained=pretrained, \n            num_classes=0, \n            global_pool=\"\", \n            in_chans=in_chans\n        )\n        backbone_out = self.backbone.feature_info[-1]['num_chs']\n        self.global_pool = GeM(p_trainable=False)\n        self.cls = torch.nn.Linear(backbone_out, n_classes)\n\n    def forward(self, x):\n        x = self.backbone(x)\n        x = self.global_pool(x)\n        x = x[:,:,0,0]\n        logits = self.cls(x)\n        return logits\n    ","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:15.642726Z","iopub.execute_input":"2023-02-16T06:46:15.643119Z","iopub.status.idle":"2023-02-16T06:46:15.655882Z","shell.execute_reply.started":"2023-02-16T06:46:15.643079Z","shell.execute_reply":"2023-02-16T06:46:15.654396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Net()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:15.657715Z","iopub.execute_input":"2023-02-16T06:46:15.658105Z","iopub.status.idle":"2023-02-16T06:46:16.072753Z","shell.execute_reply.started":"2023-02-16T06:46:15.65807Z","shell.execute_reply":"2023-02-16T06:46:16.071683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():\n    # pic = torch.ones(size=(2, 1, 224, 224))\n    pic = torch.ones(size=(2, 3, 224, 224))\n    outputs = model(pic)\n    print(outputs, '\\n', torch.sigmoid(outputs), '\\n', torch.round(torch.sigmoid(outputs)))","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:16.074186Z","iopub.execute_input":"2023-02-16T06:46:16.074537Z","iopub.status.idle":"2023-02-16T06:46:16.4143Z","shell.execute_reply.started":"2023-02-16T06:46:16.074503Z","shell.execute_reply":"2023-02-16T06:46:16.413071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## train","metadata":{}},{"cell_type":"code","source":"epochs = 20\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\noptimizer = Adam(model.parameters(), lr=1e-4)\nscheduler = ReduceLROnPlateau(optimizer=optimizer, mode='max', patience=5, verbose=True, factor=0.5)\n# criterion = nn.CrossEntropyLoss()\ncriterion = nn.BCEWithLogitsLoss()\n# criterion = criterion.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:16.416124Z","iopub.execute_input":"2023-02-16T06:46:16.416654Z","iopub.status.idle":"2023-02-16T06:46:16.464947Z","shell.execute_reply.started":"2023-02-16T06:46:16.416582Z","shell.execute_reply":"2023-02-16T06:46:16.463967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_loss = []\ntotal_acc = []\nmax_acc = 0.0\nfor epoch_num in range(epochs):\n\n    model.train()\n    total_loss_train = 0\n    total_acc_train = 0\n    for inputs, label in train_loader:\n        inputs = inputs.to(device)\n        # label = label.to(device)\n        # Unsqueeze(1): shape=[3] -> shape=[3, 1]\n        label = label.unsqueeze(1).float().to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, label)\n        total_loss_train += loss.item()\n        # acc_num = (outputs.argmax(1) == label).sum()\n        preds = torch.round(torch.sigmoid(outputs)) # 0 and 1\n        acc_num = (preds.squeeze(1) == label.squeeze(1)).sum()\n        total_acc_train += acc_num.cpu().item()\n        loss.backward()\n        optimizer.step()\n\n    model.eval()\n    total_loss_val = 0\n    total_acc_val = 0\n    with torch.no_grad():\n        for inputs, label in val_loader:\n            inputs = inputs.to(device)\n            # label = label.to(device)\n            # Unsqueeze(1): shape=[3] -> shape=[3, 1]\n            label = label.unsqueeze(1).float().to(device)\n            outputs = model(inputs)\n            loss = criterion(outputs, label)\n            total_loss_val += loss.item()\n            # acc_num = (outputs.argmax(1) == label).sum()\n            preds = torch.round(torch.sigmoid(outputs)) # 0 and 1\n            acc_num = (preds.squeeze(1) == label.squeeze(1)).sum()\n            total_acc_val += acc_num.cpu().item()\n    scheduler.step(total_acc_val)\n\n    print(\n        \"Epoch: [{:0>2d} / {:0>2d}] | Train Loss: {:.4f} | Val Loss: {:.4f} | Train Acc: {:.4f} | Val Acc: {:.4f}\".format(\n            epoch_num + 1,\n            epochs,\n            total_loss_train / len(train_dataset),\n            total_loss_val / len(val_dataset),\n            total_acc_train / len(train_dataset),\n            total_acc_val / len(val_dataset)\n        )\n    )\n    \n    if total_acc_val / len(val_dataset) > max_acc:\n        max_acc = total_acc_val / len(val_dataset)\n        torch.save(model.state_dict(), \"/kaggle/working/best_demo_nn.pt\")\n\n    total_loss.append([total_loss_train / len(train_dataset), total_loss_val / len(val_dataset)])\n    total_acc.append([(total_acc_train / len(train_dataset)), (total_acc_val / len(val_dataset))])","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:46:16.466847Z","iopub.execute_input":"2023-02-16T06:46:16.467244Z","iopub.status.idle":"2023-02-16T06:55:01.26559Z","shell.execute_reply.started":"2023-02-16T06:46:16.467205Z","shell.execute_reply":"2023-02-16T06:55:01.262524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# save\n# torch.save(model.state_dict(), \"/kaggle/working/demo_nn.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:01.266971Z","iopub.status.idle":"2023-02-16T06:55:01.267772Z","shell.execute_reply.started":"2023-02-16T06:55:01.26749Z","shell.execute_reply":"2023-02-16T06:55:01.267514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load\n# model = timm.create_model('resnet50', num_classes=1, pretrained=False, in_chans=1)\n# model.load_state_dict(torch.load(\"/kaggle/working/best_demo_nn.pt\", map_location=device))\n# model = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:01.269328Z","iopub.status.idle":"2023-02-16T06:55:01.270264Z","shell.execute_reply.started":"2023-02-16T06:55:01.269944Z","shell.execute_reply":"2023-02-16T06:55:01.269971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_loss = np.array(total_loss, dtype=np.float32)\ntotal_acc = np.array(total_acc, dtype=np.float32)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:01.27167Z","iopub.status.idle":"2023-02-16T06:55:01.272434Z","shell.execute_reply.started":"2023-02-16T06:55:01.272175Z","shell.execute_reply":"2023-02-16T06:55:01.272199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure()\nplt.xlabel(\"epoch\")\nplt.ylabel(\"loss\")\nplt.plot(total_loss[:,0], 'b-', label=\"train\")\nplt.plot(total_loss[:,1], 'r-', label=\"val\")\nplt.legend()\nplt.grid(True)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:01.273816Z","iopub.status.idle":"2023-02-16T06:55:01.274574Z","shell.execute_reply.started":"2023-02-16T06:55:01.274325Z","shell.execute_reply":"2023-02-16T06:55:01.274349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure()\nplt.xlabel(\"epoch\")\nplt.ylabel(\"acc\")\nplt.plot(total_acc[:,0], 'b-', label=\"train\")\nplt.plot(total_acc[:,1], 'r-', label=\"val\")\nplt.legend()\nplt.grid(True)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:01.275945Z","iopub.status.idle":"2023-02-16T06:55:01.276711Z","shell.execute_reply.started":"2023-02-16T06:55:01.276443Z","shell.execute_reply":"2023-02-16T06:55:01.276467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## inference","metadata":{}},{"cell_type":"code","source":"test_transform = transforms.Compose([\n            # transforms.ToPILImage(),\n            transforms.ToTensor()\n])","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:09.005701Z","iopub.execute_input":"2023-02-16T06:55:09.006073Z","iopub.status.idle":"2023-02-16T06:55:09.013467Z","shell.execute_reply.started":"2023-02-16T06:55:09.006042Z","shell.execute_reply":"2023-02-16T06:55:09.012342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = RSNADataset(\n    df=df_test, img_folder=TEST_PATH,\n    transform=test_transform, is_test=True, size=(224, 224),\n    # img_mode='fft'\n)\ntest_loader  = DataLoader(test_dataset, batch_size=1, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:12.095798Z","iopub.execute_input":"2023-02-16T06:55:12.096169Z","iopub.status.idle":"2023-02-16T06:55:12.102113Z","shell.execute_reply.started":"2023-02-16T06:55:12.096131Z","shell.execute_reply":"2023-02-16T06:55:12.100833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images = next(iter(test_loader))\nprint(images.shape)\nshow_aug(images, nrows=1, ncols=4)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:15.509925Z","iopub.execute_input":"2023-02-16T06:55:15.510788Z","iopub.status.idle":"2023-02-16T06:55:16.261709Z","shell.execute_reply.started":"2023-02-16T06:55:15.510735Z","shell.execute_reply":"2023-02-16T06:55:16.257741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prediction_model(model, data_loader, is_test=False):    \n    model.eval()\n    with torch.no_grad():\n        preds = []\n        if is_test:\n            # just images\n            for images in tqdm(data_loader, total=len(test_loader)):\n                images = images.to(device)\n                pred = model(images)\n                # label = pred.argmax(1).cpu().item()\n                preds.append(torch.sigmoid(pred).cpu().item())\n        else:\n            # there are images and targets in loader in batches\n            for images, targets in tqdm(data_loader, total=len(test_loader)):\n                images = images.to(device)\n                pred = model(images)\n                # label = pred.argmax(1).cpu().item()\n                preds.append(torch.sigmoid(pred).cpu().item())\n        return preds","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:19.139194Z","iopub.execute_input":"2023-02-16T06:55:19.139587Z","iopub.status.idle":"2023-02-16T06:55:19.147537Z","shell.execute_reply.started":"2023-02-16T06:55:19.139554Z","shell.execute_reply":"2023-02-16T06:55:19.146284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = prediction_model(model, test_loader, is_test=True)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:22.793766Z","iopub.execute_input":"2023-02-16T06:55:22.794876Z","iopub.status.idle":"2023-02-16T06:55:25.197497Z","shell.execute_reply.started":"2023-02-16T06:55:22.794826Z","shell.execute_reply":"2023-02-16T06:55:25.196327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:27.55931Z","iopub.execute_input":"2023-02-16T06:55:27.560136Z","iopub.status.idle":"2023-02-16T06:55:27.567579Z","shell.execute_reply.started":"2023-02-16T06:55:27.560087Z","shell.execute_reply":"2023-02-16T06:55:27.56666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.DataFrame(data = {\"prediction_id\": df_test['prediction_id'], \"cancer\": preds})\nsub_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:32.350107Z","iopub.execute_input":"2023-02-16T06:55:32.350496Z","iopub.status.idle":"2023-02-16T06:55:32.362497Z","shell.execute_reply.started":"2023-02-16T06:55:32.350463Z","shell.execute_reply":"2023-02-16T06:55:32.361085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sub_df.drop_duplicates(\"prediction_id\", keep='first', inplace=True)\nsub_df = sub_df.groupby(\"prediction_id\").max().reset_index()\nsub_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:55:37.534034Z","iopub.execute_input":"2023-02-16T06:55:37.534421Z","iopub.status.idle":"2023-02-16T06:55:37.551137Z","shell.execute_reply.started":"2023-02-16T06:55:37.534389Z","shell.execute_reply":"2023-02-16T06:55:37.549453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df['cancer'] = sub_df['cancer'].apply(lambda x: np.int8(x > 0.5))\ndisplay(sub_df.info())\ndisplay(sub_df.head())","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:56:38.493388Z","iopub.execute_input":"2023-02-16T06:56:38.493803Z","iopub.status.idle":"2023-02-16T06:56:38.519796Z","shell.execute_reply.started":"2023-02-16T06:56:38.49377Z","shell.execute_reply":"2023-02-16T06:56:38.518672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-02-16T06:56:47.789493Z","iopub.execute_input":"2023-02-16T06:56:47.790206Z","iopub.status.idle":"2023-02-16T06:56:47.797155Z","shell.execute_reply.started":"2023-02-16T06:56:47.790167Z","shell.execute_reply":"2023-02-16T06:56:47.796224Z"},"trusted":true},"execution_count":null,"outputs":[]}]}