{"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 --offline /kaggle/input/track4-train/*.tar.bz2","metadata":{"execution":{"iopub.status.busy":"2022-10-05T18:54:27.826719Z","iopub.execute_input":"2022-10-05T18:54:27.827113Z","iopub.status.idle":"2022-10-05T18:55:10.693046Z","shell.execute_reply.started":"2022-10-05T18:54:27.827078Z","shell.execute_reply":"2022-10-05T18:55:10.691594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !ls /kaggle/input/track4-train/models","metadata":{"execution":{"iopub.status.busy":"2022-10-05T18:55:36.793463Z","iopub.execute_input":"2022-10-05T18:55:36.793849Z","iopub.status.idle":"2022-10-05T18:55:37.742929Z","shell.execute_reply.started":"2022-10-05T18:55:36.793816Z","shell.execute_reply":"2022-10-05T18:55:37.741769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# [0.6382310415028003, 0.6189348573669843, 0.594305478150372, 0.6456368418815037, 0.650946396469057, 0.6610530600887038]\n# 0.6348512792432368\n# Full validation metric: 0.6373838583629989\n\n# Counter({1: 47225, 0: 42655})\n# ROC AUC metric: 0.6831401875473259\n# Accuracy: 0.6177681352914998\n# Full target metric: 0.5901076765112288\n# Full target metric fixed: 0.5898723615115555\n# Target metric by center_id:\n# Center_id 1: 0.5837030889916592\n# Center_id 2: 0.5128716159021234\n# Center_id 3: 0.6234127532907505\n# Center_id 4: 0.5803120005321822\n# Center_id 5: 0.633916901390537\n# Center_id 6: 0.6509644227059665\n# Center_id 7: 0.5673793620392944\n# Center_id 8: 0.6490039016757818\n# Center_id 9: 0.529810187493138\n# Center_id 10: 0.5777616326658479\n# Center_id 11: 0.5891695748590473","metadata":{"execution":{"iopub.status.busy":"2022-10-04T21:22:37.573893Z","iopub.execute_input":"2022-10-04T21:22:37.574791Z","iopub.status.idle":"2022-10-04T21:22:37.580942Z","shell.execute_reply.started":"2022-10-04T21:22:37.574749Z","shell.execute_reply":"2022-10-04T21:22:37.579689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import pandas as pd\n# from collections import Counter\n\n\n# data = pd.read_csv('/kaggle/input/mayo-clinic-strip-ai/train.csv')\n# print(Counter(data['center_id'].tolist()))","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:03:20.850969Z","iopub.execute_input":"2022-10-05T08:03:20.851359Z","iopub.status.idle":"2022-10-05T08:03:20.87424Z","shell.execute_reply.started":"2022-10-05T08:03:20.85132Z","shell.execute_reply":"2022-10-05T08:03:20.873355Z"},"trusted":true},"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\nTEST_SAMPLE_DUPL_RATE = 20\nTRAIN_SAMPLE_DUPL_RATE = 4","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-10-05T18:55:10.695789Z","iopub.execute_input":"2022-10-05T18:55:10.69618Z","iopub.status.idle":"2022-10-05T18:55:10.706201Z","shell.execute_reply.started":"2022-10-05T18:55:10.696139Z","shell.execute_reply":"2022-10-05T18:55:10.7052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nfrom collections import defaultdict\nfrom typing import List\n\nimport numpy as np\nimport torch\nfrom PIL import Image\n\n\nclass ClotImageDataset(torch.utils.data.Dataset):\n    def __init__(\n            self,\n            image_ids: List[str],\n            labels: List[str],\n            image_crops: List[List[np.ndarray]],\n            seed: int,\n            is_test: bool,\n            transformations,\n    ):\n        self.image_ids = image_ids\n        self.labels = [float(label == 'CE') for label in labels]\n        self.image_crops = image_crops\n        self.seed = seed\n        self.is_test = is_test\n        self.transformations = transformations\n\n        if not self.is_test:\n            np.random.seed(self.seed)\n\n            label_to_indices = defaultdict(list)\n            for i, (label, crops) in enumerate(zip(self.labels, self.image_crops)):\n                if len(crops) > 0:\n                    label_to_indices[label].append(i)\n\n            max_size = TRAIN_SAMPLE_DUPL_RATE * max(len(indices) for indices in label_to_indices.values())\n\n            self.sample_ids = []\n            for i, indices in enumerate(label_to_indices.values()):\n                np.random.shuffle(indices)\n                while len(self.sample_ids) < max_size * (i + 1):\n                    req_size = min(len(indices), max_size * (i + 1) - len(self.sample_ids))\n                    self.sample_ids += indices[:req_size]\n        else:\n            self.sample_ids = []\n            for _ in range(TEST_SAMPLE_DUPL_RATE):\n                self.sample_ids.extend(list(range(len(self.image_ids))))\n\n        self.image_index_ids = []\n        sample_id_to_image_index = defaultdict(int)\n        for sample_id in self.sample_ids:\n            self.image_index_ids.append(sample_id_to_image_index[sample_id])\n            image_crops_cnt = len(self.image_crops[sample_id])\n            if image_crops_cnt:\n                sample_id_to_image_index[sample_id] = (sample_id_to_image_index[sample_id] + 1) % image_crops_cnt\n\n    def __len__(self):\n        return len(self.sample_ids)\n\n    def __getitem__(self, idx):\n        if self.is_test:\n            np.random.seed(self.seed + idx)\n            random.seed(self.seed + idx)\n            torch.manual_seed(self.seed + idx)\n        idx, image_index = self.sample_ids[idx], self.image_index_ids[idx]\n        if len(self.image_crops[idx]) == 0:\n            return (\n                self.transformations(Image.fromarray(np.zeros((224, 224, 3)).astype(np.uint8))),\n                torch.tensor(self.labels[idx]),\n                self.image_ids[idx],\n            )\n        # image_index = np.random.randint(0, len(self.image_crops[idx]))\n        return (\n            self.transformations(self.image_crops[idx][image_index]),\n            torch.tensor(self.labels[idx]),\n            self.image_ids[idx],\n        )\n\n\ndef get_loader(\n        image_ids: List[str],\n        labels: List[str],\n        image_crops: List[List[np.ndarray]],\n        seed: int,\n        is_test: bool,\n        transformations,\n        shuffle: bool,\n        batch_size: int,\n        num_workers: int\n):\n    dataset = ClotImageDataset(\n        image_ids, labels, image_crops, seed, is_test, transformations,\n    )\n    return torch.utils.data.DataLoader(dataset, shuffle=shuffle, batch_size=batch_size, num_workers=num_workers)","metadata":{"execution":{"iopub.status.busy":"2022-10-05T18:55:10.7085Z","iopub.execute_input":"2022-10-05T18:55:10.709319Z","iopub.status.idle":"2022-10-05T18:55:11.440129Z","shell.execute_reply.started":"2022-10-05T18:55:10.709218Z","shell.execute_reply":"2022-10-05T18:55:11.438706Z"},"trusted":true},"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\nimport matplotlib.pyplot as plt\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#         self.test = self.train\n        test_metadata = pd.read_csv('/kaggle/input/mayo-clinic-strip-ai/test.csv')\n        self.test = list(zip(\n            test_metadata['image_id'].tolist(),\n            ['Unknown' for _ in range(test_metadata.shape[0])],\n            test_metadata['center_id'].tolist(),\n        ))\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        )\n\n    def process_test(\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.test],\n            '/kaggle/input/mayo-clinic-strip-ai/test',\n        )","metadata":{"execution":{"iopub.status.busy":"2022-10-05T18:55:11.443031Z","iopub.execute_input":"2022-10-05T18:55:11.443666Z","iopub.status.idle":"2022-10-05T18:55:11.725547Z","shell.execute_reply.started":"2022-10-05T18:55:11.443629Z","shell.execute_reply":"2022-10-05T18:55:11.724436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import transforms\n\n\ntrain_transforms = transforms.Compose([\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    # transforms.RandomResizedCrop((224, 224), scale=(0.5, 1.0), ratio=(1.0, 1.0)),\n    transforms.RandomAdjustSharpness(sharpness_factor=2, p=1.0),\n    transforms.RandomAdjustSharpness(sharpness_factor=2, p=0.5),\n    transforms.ColorJitter(brightness=0.2, saturation=0.5, hue=0.5),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])\n\ntest_transforms = transforms.Compose([\n    # transforms.RandomResizedCrop((224, 224), scale=(0.5, 1.0), ratio=(1.0, 1.0)),\n    transforms.RandomAdjustSharpness(sharpness_factor=2, p=1.0),\n    transforms.RandomAdjustSharpness(sharpness_factor=2, p=0.5),\n    transforms.ColorJitter(brightness=0.2, saturation=0.5, hue=0.5),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])","metadata":{"execution":{"iopub.status.busy":"2022-10-05T18:55:11.72682Z","iopub.execute_input":"2022-10-05T18:55:11.727381Z","iopub.status.idle":"2022-10-05T18:55:11.908056Z","shell.execute_reply.started":"2022-10-05T18:55:11.727344Z","shell.execute_reply":"2022-10-05T18:55:11.907014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import defaultdict\n\nimport numpy as np\n\n\ndef _group_by_patients(y_true, y_pred, image_ids):\n    patients = [image_id.split('_')[0] for image_id in image_ids]\n    patient_to_y_true, patient_to_y_pred = defaultdict(list), defaultdict(list)\n    for y, y_hat, patient in zip(y_true, y_pred, patients):\n        patient_to_y_true[patient].append(y)\n        patient_to_y_pred[patient].append(y_hat)\n    patient_to_y_true = {\n        patient: np.mean(y_true)\n        for patient, y_true in patient_to_y_true.items()\n    }\n    patient_to_y_pred = {\n        patient: np.mean(y_pred).tolist()\n        for patient, y_pred in patient_to_y_pred.items()\n    }\n    y_true, y_pred, patients = [], [], []\n    for patient, y in patient_to_y_true.items():\n        y_true.append(y)\n        y_pred.append(patient_to_y_pred[patient])\n        patients.append(patient)\n        \n    return y_true, np.array([[1 - p, p] for p in y_pred]), patients\n\n\ndef get_target_metric(y_true, y_pred, image_ids):\n    return _weighted_mc_log_loss(*_group_by_patients(y_true, y_pred, image_ids)[:2])\n\n\ndef _weighted_mc_log_loss(y_true, y_pred, epsilon=1e-15):\n    class_cnt = [sum(int(val == cl) for val in y_true) for cl in range(2)]\n    w = [0.5 for _ in range(2)]\n    return -sum(\n        w[cl] * sum(\n            (y == cl) / class_cnt[cl] * np.log(max(min(y_hat, 1 - epsilon), epsilon))\n            for y, y_hat in zip(y_true, y_pred[:, cl])\n        )\n        for cl in range(2)\n    ) / sum(w[cl] for cl in range(2))\n","metadata":{"execution":{"iopub.status.busy":"2022-10-05T18:55:11.909725Z","iopub.execute_input":"2022-10-05T18:55:11.910069Z","iopub.status.idle":"2022-10-05T18:55:11.923648Z","shell.execute_reply.started":"2022-10-05T18:55:11.910034Z","shell.execute_reply":"2022-10-05T18:55:11.922355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torchvision import models\n\n\nclass ClotModelSingle(nn.Module):\n    def __init__(self, encoder_model):\n        super().__init__()\n\n        if encoder_model == 'effnet_b0':\n            base_model = models.efficientnet_b0(pretrained=True)\n            self.model = base_model.features\n            in_features_cnt = base_model.classifier[1].in_features\n        elif encoder_model == 'resnet18':\n            base_model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)\n            self.model = nn.Sequential(*list(base_model.children())[:-2])\n            in_features_cnt = list(base_model.children())[-1].in_features\n        elif encoder_model == 'regnet_x_1_6gf':\n            base_model = models.regnet_x_1_6gf(weights=models.RegNet_X_1_6GF_Weights.IMAGENET1K_V2)\n            self.model = nn.Sequential(base_model.stem, base_model.trunk_output)\n            in_features_cnt = base_model.fc.in_features\n        else:\n            raise Exception('Incorrect encoder name')\n\n        self.head = nn.Sequential(\n            nn.AdaptiveAvgPool2d(output_size=1),\n            nn.Flatten(),\n            nn.Linear(in_features_cnt, 1),\n            nn.Sigmoid(),\n        )\n\n    def freeze_encoder(self, flag):\n        for param in self.model.parameters():\n            param.requires_grad = not flag\n\n    def forward(self, x):\n        return self.head(self.model(x))\n\n    def save(self, model_path):\n        weights = self.state_dict()\n        torch.save(weights, model_path)\n\n    def load(self, model_path):\n        weights = torch.load(model_path, map_location='cpu')\n        self.load_state_dict(weights)\n","metadata":{"execution":{"iopub.status.busy":"2022-10-05T18:55:11.925808Z","iopub.execute_input":"2022-10-05T18:55:11.926332Z","iopub.status.idle":"2022-10-05T18:55:11.938753Z","shell.execute_reply.started":"2022-10-05T18:55:11.926243Z","shell.execute_reply":"2022-10-05T18:55:11.937623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def update_final_prediction_with_train_data(submission: pd.DataFrame) -> pd.DataFrame:\n    train_data = pd.read_csv('/kaggle/input/mayo-clinic-strip-ai/train.csv')\n    patient_id_to_label = {\n        image_id.split('_')[0]: label\n        for image_id, label in zip(train_data['image_id'].tolist(), train_data['label'].tolist())\n    }\n    submission['CE'] = [\n        pred if patient_id not in patient_id_to_label else float(patient_id_to_label[patient_id] == 'CE')\n        for patient_id, pred in zip(submission['patient_id'].tolist(), submission['CE'].tolist())\n    ]\n    submission['LAA'] = [\n        pred if patient_id not in patient_id_to_label else float(patient_id_to_label[patient_id] == 'LAA')\n        for patient_id, pred in zip(submission['patient_id'].tolist(), submission['LAA'].tolist())\n    ]\n    return submission","metadata":{"execution":{"iopub.status.busy":"2022-10-05T18:55:11.940309Z","iopub.execute_input":"2022-10-05T18:55:11.940848Z","iopub.status.idle":"2022-10-05T18:55:11.95115Z","shell.execute_reply.started":"2022-10-05T18:55:11.940812Z","shell.execute_reply":"2022-10-05T18:55:11.950171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from __future__ import print_function, division\n\nimport os\nimport pickle\nfrom collections import Counter\nfrom typing import List\n\nimport ssl\nimport torch\nimport torch.backends.cudnn as cudnn\nfrom sklearn.metrics import roc_auc_score\nfrom tqdm import tqdm\n\n\nssl._create_default_https_context = ssl._create_unverified_context\ncudnn.benchmark = True\n\nTEST_BATCH_SIZE = 16\nDUMPED_DATALOADER_PATH = '/kaggle/input/track-4-dataprep/data_loaders.pkl'\n\n\ndef _get_model_path_by_center_id(folder_name: str, center_id: str) -> str:\n    for file_name in os.listdir(folder_name):\n        if not file_name.endswith('.h5'):\n            continue\n        if file_name.split('_')[2] == center_id:\n            return os.path.join(folder_name, file_name)\n    raise Exception(f'Model for center id {center_id} was not found in folder {folder_name}')\n\n\ndef _explode_image_ids(image_ids: List[str], factor: int) -> List[str]:\n    exploded_image_ids = []\n    for start_id in range(0, len(image_ids), TEST_BATCH_SIZE):\n        end_id = min(start_id + TEST_BATCH_SIZE, len(image_ids))\n        for _ in range(factor):\n            exploded_image_ids.extend(image_ids[start_id:end_id])\n    return exploded_image_ids\n\n\ndata_prep = DataPreparation()\n\nimage_crops, _ = data_prep.process_test()\n# with open(DUMPED_DATALOADER_PATH, 'rb') as file:\n#     image_crops, _ = pickle.load(file)\n\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\ndataloader = get_loader(\n    [image_id for image_id, _, _ in data_prep.test],\n    [label for _, label, _ in data_prep.test],\n    image_crops,\n    seed=41,\n    is_test=True,\n    transformations=test_transforms,\n    shuffle=False,\n    batch_size=TEST_BATCH_SIZE,\n    num_workers=2,\n)\n\nmodels = [\n    torch.load(_get_model_path_by_center_id(\n        '/kaggle/input/track4-train/models/',\n        center_id,\n    ), map_location=torch.device('cpu')).to(device)\n    for center_id in ['1.5', '10.3', '11', '4', '6.2.8.9', '7']\n]\nfor model in models:\n    model.eval()\n\nwith torch.no_grad():\n    y_hat, y, image_ids = [], [], []\n    for image, label, image_id in tqdm(dataloader):\n        image = image.to(device)\n        label = label.cpu().detach().numpy().tolist()\n        for model in models:\n            y_hat.extend(model.forward(image).squeeze().cpu().detach().numpy().tolist())\n            y.extend(label)\n            image_ids.extend(image_id)\n            \nbad_image_ids = {\n    image_id\n    for image_id, crops in zip([image_id for image_id, _, _ in data_prep.test], image_crops)\n    if len(crops) == 0 \n}\ny_hat_fixed = [\n    0.5 if image_id in bad_image_ids else p\n    for p, image_id in zip(y_hat, image_ids)\n]\n\nlabels, preds, patients = _group_by_patients(y, y_hat_fixed, image_ids)\nresult = pd.DataFrame({\n    'patient_id': patients,\n    'CE': [pair[1].round(6) for pair in preds],\n    'LAA': [pair[0].round(6) for pair in preds],\n})\nprint(result)\nresult = update_final_prediction_with_train_data(result)\nprint(result)\nresult.to_csv('submission.csv', index=False)\n# print(Counter([int(p > 0.5) for p in y_hat_fixed]))\n# print('ROC AUC metric:', roc_auc_score(y, y_hat_fixed))\n# accuracy = sum([int(int(p > 0.5) == label) for p, label in zip(y_hat_fixed, y)]) / len(y)\n# print('Accuracy:', accuracy)\n# image_id_to_center_id = {image_id: center_id for image_id, _, center_id in data_prep.test}\n# target_metric = get_target_metric(\n#     y,\n#     y_hat,\n#     image_ids,\n# )\n# print('Full target metric:', target_metric)\n# target_metric = get_target_metric(\n#     y,\n#     y_hat_fixed,\n#     image_ids,\n# )\n# print('Full target metric fixed:', target_metric)\n# print('Target metric by center_id:')\n# for center_id in range(1, 12):\n#     sub_y = [label for label, image_id in zip(y, image_ids) if image_id_to_center_id[image_id] == center_id]\n#     sub_y_hat = [pred for pred, image_id in zip(y_hat_fixed, image_ids) if image_id_to_center_id[image_id] == center_id]\n#     sub_image_ids = [image_id for image_id in image_ids if image_id_to_center_id[image_id] == center_id]\n#     print(f'Center_id {center_id}: {get_target_metric(sub_y, sub_y_hat, sub_image_ids)}')","metadata":{"execution":{"iopub.status.busy":"2022-10-05T18:56:25.775989Z","iopub.execute_input":"2022-10-05T18:56:25.77638Z","iopub.status.idle":"2022-10-05T18:58:56.681393Z","shell.execute_reply.started":"2022-10-05T18:56:25.776347Z","shell.execute_reply":"2022-10-05T18:58:56.68007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}