{"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":"# Oversampling positive cancer cases to get balanced training batches","metadata":{}},{"cell_type":"markdown","source":"Only ~2% of the images in the training data of the RSNA Breast cancer detection competition contain cancer. Because of this cancer positive cases will be very rare in batches of training samples, thus preventing a model from learning features to classify positive from negativ cases. To balance positive and negative cases in training batches, oversampling can be used.","metadata":{}},{"cell_type":"code","source":"# install libraries to read dicom files\n!pip install /kaggle/input/gdcm-dicomsdl-dali/{pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.21-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}","metadata":{"execution":{"iopub.status.busy":"2023-02-20T17:57:22.491154Z","iopub.execute_input":"2023-02-20T17:57:22.492014Z","iopub.status.idle":"2023-02-20T17:57:36.043021Z","shell.execute_reply.started":"2023-02-20T17:57:22.491915Z","shell.execute_reply":"2023-02-20T17:57:36.0416Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom PIL import Image\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.transforms import ToTensor\nfrom collections import Counter\nimport glob\nimport os\nimport pydicom\nimport cv2\nfrom joblib import Parallel, delayed\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-02-20T17:57:36.045443Z","iopub.execute_input":"2023-02-20T17:57:36.04591Z","iopub.status.idle":"2023-02-20T17:57:38.328037Z","shell.execute_reply.started":"2023-02-20T17:57:36.045874Z","shell.execute_reply":"2023-02-20T17:57:38.32685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Prepare some data","metadata":{"execution":{"iopub.status.busy":"2023-02-19T20:30:32.287051Z","iopub.execute_input":"2023-02-19T20:30:32.287492Z","iopub.status.idle":"2023-02-19T20:30:32.292825Z","shell.execute_reply.started":"2023-02-19T20:30:32.287456Z","shell.execute_reply":"2023-02-19T20:30:32.291534Z"}}},{"cell_type":"code","source":"def preprocess(dcm_file, out_dir):\n    \"\"\"Preprocess image data.\n    \n    Preprocessed images are save in the out_dir folder.\n    \"\"\"\n    ds = pydicom.dcmread(dcm_file)\n    img = ds.pixel_array\n    img = (img - img.min()) / (img.max() - img.min())\n    if ds.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n    img = (img * 255).astype(np.uint8)\n    #\n    # do some preprocessing stuff\n    # just doing resizing for demo\n    img = cv2.resize(img, (224, 224))\n    patient_id = dcm_file.split('/')[-2]\n    image_id = dcm_file.split('/')[-1][:-4]\n    out_file = os.path.join(out_dir, f\"{patient_id}_{image_id}.jpg\")\n    cv2.imwrite(out_file, img)","metadata":{"execution":{"iopub.status.busy":"2023-02-20T17:57:38.329668Z","iopub.execute_input":"2023-02-20T17:57:38.330245Z","iopub.status.idle":"2023-02-20T17:57:38.339827Z","shell.execute_reply.started":"2023-02-20T17:57:38.330212Z","shell.execute_reply":"2023-02-20T17:57:38.337936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# preprocessed images are saved here\nout_dir = '../tmp/output'\nos.makedirs(out_dir, exist_ok=True)\n\n# Just processing first 128 image files to demo\ndcm_files = glob.glob(\"../input/rsna-breast-cancer-detection/train_images/*/*.dcm\")[:128]\n\n# create a lookup dict with image_id as key and cancer indicator as value\ndf_train = pd.read_csv(\"../input/rsna-breast-cancer-detection/train.csv\")\nkeys = df_train.image_id.astype(str).values\nvals = df_train.cancer.values\nlabel_dict = {k: v for k, v in zip(keys, vals)}","metadata":{"execution":{"iopub.status.busy":"2023-02-20T17:57:38.34203Z","iopub.execute_input":"2023-02-20T17:57:38.342371Z","iopub.status.idle":"2023-02-20T17:58:38.733523Z","shell.execute_reply.started":"2023-02-20T17:57:38.342328Z","shell.execute_reply":"2023-02-20T17:58:38.732521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# do the preprocessing\n_ = Parallel(n_jobs=8)(delayed(preprocess)(dcm, out_dir=out_dir) for dcm in tqdm(dcm_files))","metadata":{"execution":{"iopub.status.busy":"2023-02-20T17:58:38.735216Z","iopub.execute_input":"2023-02-20T17:58:38.735563Z","iopub.status.idle":"2023-02-20T17:59:42.224393Z","shell.execute_reply.started":"2023-02-20T17:58:38.735535Z","shell.execute_reply":"2023-02-20T17:59:42.223399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset implementing oversampling to balance training batches","metadata":{"execution":{"iopub.status.busy":"2023-02-19T21:03:32.535811Z","iopub.execute_input":"2023-02-19T21:03:32.536347Z","iopub.status.idle":"2023-02-19T21:03:32.543021Z","shell.execute_reply.started":"2023-02-19T21:03:32.536303Z","shell.execute_reply":"2023-02-19T21:03:32.541693Z"}}},{"cell_type":"code","source":"class ImageDatasetBalanced(Dataset):\n    \"\"\"Oversampling cancer cases to produce balanced training batches.\n    \n    Preprocessed files are split into a positive and negative cancer list\n    according to the provided label dict.\n    \n    Getting an item indexes alternately into the pos and neg lists, so that\n    even indexes retrieve a sample from the pos list and odd indexes from the\n    neg list.\n    \n    If e.g. the pos list contains only a few samples, these are sampled over\n    and over again to keep the 1:1 sampling ratio.\n    \n    If the ratio parameter is not used, indexing into the dataset is done\n    as usual. This is useful e.g. in validation, when batches do not need to \n    be balanced.\n    \"\"\"\n    \n    def __init__(self, image_dir, label_dict=dict(), ratio=None, transform=None):\n        image_dir = os.path.join(image_dir, '')\n        samples = glob.glob(f\"{image_dir}*.jpg\")\n        self.samples = samples\n        \n        # lists for positive and negative cancer cases\n        pos_list = []\n        neg_list = []\n        for f in samples:\n            if label_dict.setdefault(self._get_key(f), 0) == 1:\n                pos_list.append(f)\n            else:\n                neg_list.append(f)\n        np.random.shuffle(pos_list)\n        np.random.shuffle(neg_list)\n        self.pos_list = pos_list\n        self.neg_list = neg_list\n        \n        self.image_dir = image_dir\n        self.ratio = ratio\n        self.label_dict = label_dict\n        self.transform = transform\n            \n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n        \n        if self.ratio:\n            # index in pos list\n            pos_idx = idx // (self.ratio + 1)\n            # every other sample is taken from neg list\n            if idx % (self.ratio + 1):\n                neg_idx = idx - 1 - pos_idx\n                # if list is exhausted, wrap around\n                neg_idx = neg_idx % len(self.neg_list)\n                file = self.neg_list[neg_idx]\n            else:\n                # wrap around if necessary\n                pos_idx = pos_idx % len(self.pos_list)\n                file = self.pos_list[pos_idx]\n        else:\n            file = self.samples[idx]\n        \n        key = self._get_key(file)\n        label = self.label_dict.setdefault(key, 0)\n        \n        im_arr = cv2.imread(file)\n        image = Image.fromarray(im_arr).convert(\"RGB\")\n        if self.transform:\n            image = self.transform(image)\n        return image, label\n    \n    def _get_key(self, filename):\n        return os.path.basename(filename).split(\"_\")[1][:-4]","metadata":{"execution":{"iopub.status.busy":"2023-02-20T17:59:42.225877Z","iopub.execute_input":"2023-02-20T17:59:42.226301Z","iopub.status.idle":"2023-02-20T17:59:42.242668Z","shell.execute_reply.started":"2023-02-20T17:59:42.226258Z","shell.execute_reply":"2023-02-20T17:59:42.241491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create a dataset and dataloader\n# With a parameter ratio=1, the batches are balanced in a 1:1 ratio\nds = ImageDatasetBalanced(\"../tmp/output\", label_dict=label_dict, ratio=1, transform=ToTensor())\ndl = DataLoader(ds, batch_size=32, num_workers=2, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-02-20T17:59:42.245175Z","iopub.execute_input":"2023-02-20T17:59:42.246155Z","iopub.status.idle":"2023-02-20T17:59:42.269604Z","shell.execute_reply.started":"2023-02-20T17:59:42.246108Z","shell.execute_reply":"2023-02-20T17:59:42.268466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# The dataloader generates batches with postive and negative cases\n# in a 1:1 ratio\nfor batch, label in dl:\n    label = label.numpy()\n    print(label)","metadata":{"execution":{"iopub.status.busy":"2023-02-20T17:59:42.271113Z","iopub.execute_input":"2023-02-20T17:59:42.271655Z","iopub.status.idle":"2023-02-20T17:59:42.613963Z","shell.execute_reply.started":"2023-02-20T17:59:42.271598Z","shell.execute_reply":"2023-02-20T17:59:42.612566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create a dataset and dataloader\n# This time without giving the ratio parameter\nds = ImageDatasetBalanced(\"../tmp/output\", label_dict=label_dict, transform=ToTensor())\ndl = DataLoader(ds, batch_size=32, num_workers=2, shuffle=False)\nfor batch, label in dl:\n    label = label.numpy()\n    print(label)","metadata":{"execution":{"iopub.status.busy":"2023-02-20T17:59:42.615619Z","iopub.execute_input":"2023-02-20T17:59:42.615974Z","iopub.status.idle":"2023-02-20T17:59:42.861221Z","shell.execute_reply.started":"2023-02-20T17:59:42.615942Z","shell.execute_reply":"2023-02-20T17:59:42.860106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}