{"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":"!conda install -y pyvips","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BAD_IMAGE_IDS = ['5adc4c_0', '7b9aaa_0', 'bb06a5_0', 'e26a04_0', '280c26_0'] + \\\n                ['4ae44b_0', '53e66f_0', '7c2c2f_0', '74a450_1']\n\nBLOCK_SIZE = 28\nBLOCKS_PER_CROP = 8\nCROP_SIZE = BLOCK_SIZE * BLOCKS_PER_CROP\nBLOCK_THR = 90\nCROP_THR = 0.6\nMAX_CROPS_PER_IMAGE = 20\nIMAGES_PER_SAMPLE = 4\nEPOCHS_NUM = 10\nSCALE_FACTOR = 24","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nimport os\nfrom time import time\nfrom typing import List, Tuple\n\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm import tqdm\n\nimport pyvips\nimport cv2\n\n\nclass DataPreparation:\n    def __init__(self, visualize: bool = False, seed: int = 42):\n        self.visualize = visualize\n        self.seed = seed\n\n        train_metadata = pd.read_csv('/kaggle/input/mayo-clinic-strip-ai/train.csv')\n        train_metadata = list(zip(\n            train_metadata['image_id'].tolist(),\n            train_metadata['label'].tolist(),\n            train_metadata['center_id'].tolist(),\n        ))\n        self.train = self._filter_bad_images(train_metadata)\n        self.all_center_ids = sorted(list({center_id for _, _, center_id in self.train}))\n\n        other_metadata = pd.read_csv('/kaggle/input/mayo-clinic-strip-ai/other.csv').query('label == \\'Other\\'')\n        other_metadata = list(zip(\n            other_metadata['image_id'].tolist(),\n            ['LAA' for _ in range(other_metadata.shape[0])],\n            [-1 for _ in range(other_metadata.shape[0])],\n        ))\n        self.other = self._filter_bad_images(other_metadata)\n\n    @staticmethod\n    def _filter_bad_images(data: List[Tuple]) -> List[Tuple]:\n        return [\n            (image_id, label, center_id)\n            for image_id, label, center_id in data\n            if image_id not in BAD_IMAGE_IDS\n        ]\n\n    @staticmethod\n    def _add_rect_to_numpy(image: np.ndarray, x: int, y: int, size: int, thickness: int) -> None:\n        image[x:x + size, y:y + thickness] = (0, 0, 0)\n        image[x:x + thickness, y:y + size] = (0, 0, 0)\n        image[x:x + size, y + size:y + size + thickness] = (0, 0, 0)\n        image[x + size:x + size + thickness, y:y + size] = (0, 0, 0)\n\n    @staticmethod\n    def _get_blocks_map(image: np.ndarray) -> np.ndarray:\n        pixels_diff = np.sum((image[:-1, :, :] - image[1:, :, :]) ** 2, axis=2)\n        pixels_diff = np.cumsum(np.cumsum(pixels_diff, axis=0), axis=1)\n        blocks_map = np.zeros((\n            (image.shape[0] + BLOCK_SIZE - 1) // BLOCK_SIZE,\n            (image.shape[1] + BLOCK_SIZE - 1) // BLOCK_SIZE,\n        ))\n        for x in range(0, pixels_diff.shape[0], BLOCK_SIZE):\n            for y in range(0, pixels_diff.shape[1], BLOCK_SIZE):\n                nx = min(x + BLOCK_SIZE, pixels_diff.shape[0])\n                ny = min(y + BLOCK_SIZE, pixels_diff.shape[1])\n                block_sum = int(pixels_diff[nx - 1, ny - 1])\n                if x:\n                    block_sum -= int(pixels_diff[x - 1, ny - 1])\n                if y:\n                    block_sum -= int(pixels_diff[nx - 1, y - 1])\n                if x and y:\n                    block_sum += int(pixels_diff[x - 1, y - 1])\n                blocks_map[x // BLOCK_SIZE][y // BLOCK_SIZE] = \\\n                    (block_sum / BLOCK_SIZE / BLOCK_SIZE) > BLOCK_THR\n        return blocks_map\n\n    def _generate_crops_positions(\n            self,\n            image: np.ndarray,\n            crop_thr: float,\n    ) -> Tuple[List[Tuple[int, int]], np.ndarray, np.ndarray, np.ndarray, np.ndarray]:\n        blocks_map = self._get_blocks_map(image)\n\n        if self.visualize:\n            for i in range(blocks_map.shape[0]):\n                for j in range(blocks_map.shape[1]):\n                    if blocks_map[i][j]:\n                        self._add_rect_to_numpy(\n                            image,\n                            i * BLOCK_SIZE,\n                            j * BLOCK_SIZE,\n                            BLOCK_SIZE,\n                            1,\n                        )\n\n        good_crops_starts = []\n        for x in range(0, image.shape[0] - CROP_SIZE + 1, BLOCK_SIZE):\n            for y in range(0, image.shape[1] - CROP_SIZE + 1, BLOCK_SIZE):\n                _x, _y = x // BLOCK_SIZE, y // BLOCK_SIZE\n                crop_sum = blocks_map[_x:_x + BLOCKS_PER_CROP, _y:_y + BLOCKS_PER_CROP].sum()\n                if crop_sum > BLOCKS_PER_CROP * BLOCKS_PER_CROP * crop_thr:\n                    good_crops_starts.append((x, y))\n\n        if self.visualize:\n            for x, y in good_crops_starts:\n                self._add_rect_to_numpy(image, x, y, CROP_SIZE, 1)\n\n        return good_crops_starts\n\n    @staticmethod\n    def _process_crop(crop: np.ndarray) -> np.ndarray:\n        return crop\n\n    def _create_crops(\n        self,\n        image: np.ndarray,\n        crops_starts: List[Tuple[int]],\n    ) -> List[np.ndarray]:\n        return [\n            Image.fromarray(\n                self._process_crop(\n                    image[x:x + CROP_SIZE, y:y + CROP_SIZE],\n                )\n            )\n            for x, y in crops_starts\n        ]\n\n    @staticmethod\n    def _get_unique_crops(crop_starts: List[Tuple[int, int]], order) -> List[Tuple[int, int]]:\n        def inter_size_1d(a: int, b: int, c: int, d: int) -> int:\n            return max(0, min(b, d) - max(a, c))\n\n        def inter_size_2d(crop_start_1: Tuple[int, int], crop_start_2: Tuple[int, int]) -> int:\n            return inter_size_1d(\n                crop_start_1[0], crop_start_1[0] + CROP_SIZE,\n                crop_start_2[0], crop_start_2[0] + CROP_SIZE,\n            ) * inter_size_1d(\n                crop_start_1[1], crop_start_1[1] + CROP_SIZE,\n                crop_start_2[1], crop_start_2[1] + CROP_SIZE,\n            )\n\n        crop_starts_sorted = sorted(crop_starts, key=order)\n        final_crop_starts = []\n        for crop_start in crop_starts_sorted:\n            if any(\n                    inter_size_2d(crop_start, crop_start_prev) > CROP_SIZE * CROP_SIZE // 2\n                    for crop_start_prev in final_crop_starts\n            ):\n                continue\n            final_crop_starts.append(crop_start)\n        return final_crop_starts\n    \n    @staticmethod\n    def _read_and_resize_image(image_id: str, base_image_path: str) -> np.ndarray:       \n        image_path = os.path.join(base_image_path, f'{image_id}.tif')\n        image = pyvips.Image.new_from_file(image_path, access='sequential')\n        return image.resize(1.0 / SCALE_FACTOR).numpy()\n    \n    def prepare_crops(\n            self,\n            image_ids: List[int],\n            base_image_path: str,\n    ) -> Tuple[List[List[np.ndarray]], List[List[Tuple[int]]], List[Tuple[np.ndarray, np.ndarray]]]:\n        np.random.seed(self.seed)\n        image_crops = []\n        image_crops_indices = []\n        for image_id in tqdm(image_ids):\n            start_time = time()\n            image = self._read_and_resize_image(image_id, base_image_path)\n            gc.collect()\n            print(f'Rescaling done in {time() - start_time} seconds. Image shape is {image.shape}')\n            found_flag = False\n            for crop_thr in np.arange(CROP_THR, -0.1, -0.1):\n                good_crops_starts = self._generate_crops_positions(image, crop_thr)\n                if len(good_crops_starts) < IMAGES_PER_SAMPLE:\n                    print('Bad image', image_id, 'crop_thr', crop_thr, 'only', len(good_crops_starts))\n                    continue\n\n                good_crops_starts_unique = []\n                for order in [\n                    lambda x: (x[0], x[1]),\n                    lambda x: (-x[0], -x[1]),\n                ]:\n                    good_crops_starts_unique.extend(self._get_unique_crops(good_crops_starts, order))\n                good_crops_starts_unique = list(set(good_crops_starts_unique))\n\n                if len(good_crops_starts_unique) < IMAGES_PER_SAMPLE:\n                    print('Bad image', image_id, 'crop_thr', crop_thr, 'only', len(good_crops_starts_unique))\n                    continue\n\n                good_crops_starts_sample_ids = np.random.choice(\n                    list(range(len(good_crops_starts_unique))),\n                    min(len(good_crops_starts_unique), MAX_CROPS_PER_IMAGE),\n                    replace=False,\n                )\n                good_crops_starts_sample = np.array(good_crops_starts_unique)[good_crops_starts_sample_ids]\n                image_crops_indices.append(good_crops_starts_sample)\n                image_crops.append(self._create_crops(image, good_crops_starts_sample))\n                found_flag = True\n                break\n            if not found_flag:\n                image_crops_indices.append([])\n                image_crops.append([])\n                print('No crops was found')\n            print(f'Done {image_id} in {time() - start_time} seconds')\n        gc.collect()\n        return image_crops, image_crops_indices\n\n    def process_train(\n            self\n    ) -> Tuple[List[List[np.ndarray]], List[List[Tuple[int]]]]:\n        return self.prepare_crops(\n            [image_id for image_id, _, _ in self.train],\n            '/kaggle/input/mayo-clinic-strip-ai/train/',\n        )\n\n    def process_other(\n            self\n    ) -> Tuple[List[List[np.ndarray]], List[List[Tuple[int]]]]:\n        return self.prepare_crops(\n            [image_id for image_id, _, _ in self.other],\n            '/kaggle/input/mayo-clinic-strip-ai/other/',\n        )","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from __future__ import print_function, division\n\nimport os\nimport pickle\nimport sys\nfrom collections import Counter\n\nimport cv2\nimport numpy as np\nimport ssl\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.backends.cudnn as cudnn\nfrom PIL import Image\nfrom sklearn.metrics import roc_auc_score\nfrom tqdm import tqdm\nfrom torchvision import transforms\n\n\nssl._create_default_https_context = ssl._create_unverified_context\ncudnn.benchmark = True\n\n\nDUMPED_DATALOADER_PATH = '/kaggle/working/data_loaders.pkl'\nDUMPED_DATALOADER_OTHER_PATH = '/kaggle/working/data_loaders_other.pkl'\n\n\ndef get_sub_data(data, image_crops, image_crops_indices, sample_ids):\n    return [data[i][0] for i in sample_ids], \\\n        [data[i][1] for i in sample_ids], \\\n        [data[i][2] for i in sample_ids], \\\n        [image_crops[i] for i in sample_ids], \\\n        [image_crops_indices[i] for i in sample_ids]\n\n\ndata_prep = DataPreparation()\n\nimage_crops, image_crops_indices = data_prep.process_train()\nwith open(DUMPED_DATALOADER_PATH, 'wb') as file:\n    pickle.dump([image_crops, image_crops_indices], file)\n\nimage_crops_other, image_crops_indices_other = data_prep.process_other()\nwith open(DUMPED_DATALOADER_OTHER_PATH, 'wb') as file:\n    pickle.dump([image_crops_other, image_crops_indices_other], file)","metadata":{},"execution_count":null,"outputs":[]}]}