{"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 config --add pkgs_dirs ./ # Set the location where conda package will be downloaded\n!conda install --download-only -y pyvips # Download pyvips and dependencies","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-10-05T08:20:36.810194Z","iopub.execute_input":"2022-10-05T08:20:36.810589Z","iopub.status.idle":"2022-10-05T08:23:37.221991Z","shell.execute_reply.started":"2022-10-05T08:20:36.810486Z","shell.execute_reply":"2022-10-05T08:23:37.220765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!conda install *.tar.bz2 ","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:23:37.224479Z","iopub.execute_input":"2022-10-05T08:23:37.224867Z","iopub.status.idle":"2022-10-05T08:24:11.233772Z","shell.execute_reply.started":"2022-10-05T08:23:37.224832Z","shell.execute_reply":"2022-10-05T08:24:11.232575Z"},"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# res = dict(Counter(data['center_id'].tolist()))\n# print(res)","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:45:16.749036Z","iopub.execute_input":"2022-10-05T08:45:16.749432Z","iopub.status.idle":"2022-10-05T08:45:16.760767Z","shell.execute_reply.started":"2022-10-05T08:45:16.749395Z","shell.execute_reply":"2022-10-05T08:45:16.759689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# groups = [(11,), (4,), (7,), (1, 5,), (10, 3), (6, 2, 8, 9,)]\n# for gr in groups:\n#     print(gr, sum(res[i] for i in gr))","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:46:17.484332Z","iopub.execute_input":"2022-10-05T08:46:17.484717Z","iopub.status.idle":"2022-10-05T08:46:17.493437Z","shell.execute_reply.started":"2022-10-05T08:46:17.484685Z","shell.execute_reply":"2022-10-05T08:46:17.492189Z"},"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","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:24:11.236683Z","iopub.execute_input":"2022-10-05T08:24:11.237856Z","iopub.status.idle":"2022-10-05T08:24:11.246285Z","shell.execute_reply.started":"2022-10-05T08:24:11.237811Z","shell.execute_reply":"2022-10-05T08:24:11.245248Z"},"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\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":{"iopub.status.busy":"2022-10-05T08:24:11.252887Z","iopub.execute_input":"2022-10-05T08:24:11.253766Z","iopub.status.idle":"2022-10-05T08:24:11.889771Z","shell.execute_reply.started":"2022-10-05T08:24:11.253715Z","shell.execute_reply":"2022-10-05T08:24:11.888806Z"},"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 = 4 * 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(20):\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-05T08:24:11.891434Z","iopub.execute_input":"2022-10-05T08:24:11.891825Z","iopub.status.idle":"2022-10-05T08:24:12.646162Z","shell.execute_reply.started":"2022-10-05T08:24:11.891789Z","shell.execute_reply":"2022-10-05T08:24:12.64501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import defaultdict\n\nimport numpy as np\n\n\ndef get_target_metric(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 = [], []\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    return _weighted_mc_log_loss(y_true, np.array([[1 - p, p] for p in y_pred]))\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))","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:24:12.651286Z","iopub.execute_input":"2022-10-05T08:24:12.651953Z","iopub.status.idle":"2022-10-05T08:24:12.670538Z","shell.execute_reply.started":"2022-10-05T08:24:12.651915Z","shell.execute_reply":"2022-10-05T08:24:12.669553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n# import pretrainedmodels as pm\nfrom torchvision import models\n\n\nclass ClotModelMIL(nn.Module):\n    def __init__(self, num_crops=None):\n        super().__init__()\n        self.num_crops = num_crops\n\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        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        # x: bs x N x C x W x W\n        bs, _, ch, w, h = x.shape\n        x = x.view(bs * self.num_crops, ch, w, h)  # x: N bs x C x W x W\n        x = self.model(x)  # x: N bs x C' x W' x W'\n\n        # Concat and pool\n        bs2, ch2, w2, h2 = x.shape\n        x = x \\\n            .view(-1, self.num_crops, ch2, w2, h2) \\\n            .permute(0, 2, 1, 3, 4) \\\n            .contiguous() \\\n            .view(bs, ch2, self.num_crops * w2, h2)  # x: bs x C' x N W'' x W''\n        return self.head(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\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) # weights=models.EfficientNet_B0_Weights.IMAGENET1K_V1)\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)","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:24:12.675923Z","iopub.execute_input":"2022-10-05T08:24:12.678399Z","iopub.status.idle":"2022-10-05T08:24:12.992967Z","shell.execute_reply.started":"2022-10-05T08:24:12.67836Z","shell.execute_reply":"2022-10-05T08:24:12.987179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\n\nos.makedirs('/kaggle/working/models/', exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:24:12.998436Z","iopub.execute_input":"2022-10-05T08:24:12.998841Z","iopub.status.idle":"2022-10-05T08:24:13.009736Z","shell.execute_reply.started":"2022-10-05T08:24:12.9988Z","shell.execute_reply":"2022-10-05T08:24:13.007877Z"},"trusted":true},"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/input/track-4-dataprep/data_loaders.pkl'\nDUMPED_DATALOADER_OTHER_PATH = '/kaggle/input/track-4-dataprep/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\n# image_crops, image_crops_indices = data_prep.process_train()\n# with open(DUMPED_DATALOADER_PATH, 'wb') as file:\n#     pickle.dump([image_crops, image_crops_indices], file)\n\n# image_crops_other, image_crops_indices_other = data_prep.process_other()\n# with open(DUMPED_DATALOADER_OTHER_PATH, 'wb') as file:\n#     pickle.dump([image_crops_other, image_crops_indices_other], file)\n    \nwith open(DUMPED_DATALOADER_PATH, 'rb') as file:\n    image_crops, image_crops_indices = pickle.load(file)\nwith open(DUMPED_DATALOADER_OTHER_PATH, 'rb') as file:\n    image_crops_other, image_crops_indices_other = pickle.load(file)    ","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:24:13.011381Z","iopub.execute_input":"2022-10-05T08:24:13.013151Z","iopub.status.idle":"2022-10-05T08:24:25.323342Z","shell.execute_reply.started":"2022-10-05T08:24:13.013114Z","shell.execute_reply":"2022-10-05T08:24:25.322357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transforms = {\n    'train': 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    'test': 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    ]),\n}\n\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\ntrain_data = data_prep.train + data_prep.other\nimage_crops += image_crops_other\nimage_crops_indices += image_crops_indices_other\n\nall_metrics = []\nall_y, all_y_hat, all_image_ids = [], [], []\n#for test_center_id in {center_id for _, _, center_id in train_data if center_id != -1}:\nfor test_centers_group in [(11,), (4,), (7,), (1, 5,), (10, 3), (6, 2, 8, 9,)]:\n    best_validation_metric = None\n    best_y, best_y_hat, best_image_ids = [], [], []\n    for iteration in range(3):\n        for _ in range(3):\n            print('-' * 80)\n        test_centers_group_str = '.'.join(map(str, test_centers_group))\n        print(f'CV with {test_centers_group_str} as test')\n        train_sample_ids = [i for i, (_, _, center_id) in enumerate(train_data) if center_id not in test_centers_group]\n        test_sample_ids = [i for i, (_, _, center_id) in enumerate(train_data) if center_id in test_centers_group]\n\n        train_image_ids, train_labels, train_center_ids, train_crops, train_crop_indices = get_sub_data(\n            train_data,\n            image_crops,\n            image_crops_indices,\n            train_sample_ids\n        )\n        test_image_ids, test_labels, test_center_ids, test_crops, test_crop_indices = get_sub_data(\n            train_data,\n            image_crops,\n            image_crops_indices,\n            test_sample_ids\n        )\n        print(f'Train/Test sizes: {len(train_labels)}/{len(test_labels)}')\n        print('Train/Test label distribution:')\n        print({key: value / len(train_labels) for key, value in dict(Counter(train_labels)).items()})\n        print({key: value / len(test_labels) for key, value in dict(Counter(test_labels)).items()})\n\n        dataloaders = {\n            'train': get_loader(\n                train_image_ids,\n                train_labels,\n                train_crops,\n                seed=42,\n                is_test=False,\n                transformations=data_transforms['train'],\n                shuffle=True,\n                batch_size=64,\n                num_workers=2,\n            ),\n            'test': get_loader(\n                test_image_ids,\n                test_labels,\n                test_crops,\n                seed=42,\n                is_test=True,\n                transformations=data_transforms['test'],\n                shuffle=False,\n                batch_size=64,\n                num_workers=2,\n            ),\n        }\n\n        model = ClotModelSingle(encoder_model='effnet_b0').to(device)\n        model.freeze_encoder(True)\n        criterion = nn.BCELoss()\n        optimizer = optim.Adam(model.head.parameters(), lr=5e-3, weight_decay=5e-4)\n\n        train_loss, val_loss = [], []\n        for epoch in range(EPOCHS_NUM):\n            np.random.seed(epoch)\n\n            #if epoch == 10:\n            #    optimizer = optim.Adam(model.parameters(), lr=1e-4, weight_decay=5e-4)\n\n            print('*' * 80)\n            print(\"epoch {}/{}\".format(epoch + 1, EPOCHS_NUM))\n\n            model.train()\n            running_loss, running_score = 0.0, 0.0\n            y_hat, y, image_ids = [], [], []\n            for image, label, image_id in tqdm(dataloaders['train']):\n                image = image.to(device)\n                label = label.to(device)\n                optimizer.zero_grad()\n                y_pred = model.forward(image).squeeze()\n                loss = criterion(y_pred, label)\n                running_loss += loss.item()\n                loss.backward()\n                optimizer.step()\n\n                y_pred = y_pred.cpu().detach().numpy().tolist()\n                label = label.cpu().detach().numpy().tolist()\n                y_hat.extend(y_pred)\n                y.extend(label)\n                image_ids.extend(image_id)\n\n                running_score += sum([int(int(y_hat > 0.5) == y) for y_hat, y in zip(y_pred, label)])\n\n            print('Train:')\n            print(Counter([int(p > 0.5) for p in y_hat]))\n            print('ROC AUC metric:', roc_auc_score(y, y_hat))\n            print('target metric:', get_target_metric(y, y_hat, image_ids))\n\n            epoch_score = running_score / len(dataloaders['train'].dataset)\n            epoch_loss = running_loss / len(dataloaders['train'])\n            train_loss.append(epoch_loss)\n            print(\"loss: {}, accuracy: {}\".format(epoch_loss, epoch_score))\n\n            with torch.no_grad():\n                model.eval()\n                running_loss, running_score = 0.0, 0.0\n                y_hat, y, image_ids = [], [], []\n                for image, label, image_id in tqdm(dataloaders['test']):\n                    image = image.to(device)\n                    label = label.to(device)\n                    optimizer.zero_grad()\n                    y_pred = model.forward(image).squeeze()\n                    loss = criterion(y_pred, label)\n                    running_loss += loss.item()\n\n                    y_pred = y_pred.cpu().detach().numpy().tolist()\n                    label = label.cpu().detach().numpy().tolist()\n                    y_hat.extend(y_pred)\n                    y.extend(label)\n                    image_ids.extend(image_id)\n\n                    running_score += sum([int(int(y_hat > 0.5) == y) for y_hat, y in zip(y_pred, label)])\n\n                bad_image_ids = {\n                    image_id\n                    for image_id, crops in zip(test_image_ids, test_crops)\n                    if len(crops) == 0 \n                }\n                y_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\n                print('Validation:')\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\n                target_metric = get_target_metric(y, y_hat, image_ids)\n                print('target metric:', target_metric)\n                target_metric = get_target_metric(y, y_hat_fixed, image_ids)\n                print('target metric fixed:', target_metric)            \n                if best_validation_metric is None or target_metric < best_validation_metric:\n                    best_validation_metric = target_metric\n                    best_y, best_y_hat, best_image_ids = y, y_hat_fixed, image_ids\n                    torch.save(\n                        model,\n                        os.path.join(\n                            '/kaggle/working/models',\n                            f'center_id_{test_centers_group_str}_epoch_{epoch}_target_{round(target_metric, 3)}.h5',\n                        ),\n                    )\n\n                epoch_score = running_score / len(dataloaders['test'].dataset)\n                epoch_loss = running_loss / len(dataloaders['test'])\n                val_loss.append(epoch_loss)\n                print(\"loss: {}, accuracy: {}\".format(epoch_loss, epoch_score))\n\n    print(f'Best validation metric: {best_validation_metric}')\n    all_metrics.append(best_validation_metric)\n    all_y.extend(best_y)\n    all_y_hat.extend(best_y_hat)\n    all_image_ids.extend(best_image_ids)\n\nprint(all_metrics)\nprint(np.mean(all_metrics))\nfinal_metric = get_target_metric(\n    all_y,\n    all_y_hat,\n    all_image_ids,\n)\nnp.save('/kaggle/working/all_y.npy', all_y)\nnp.save('/kaggle/working/all_y_hat.npy', all_y_hat)\nnp.save('/kaggle/working/all_image_ids.npy', all_image_ids)\nprint('Full validation metric:', final_metric)","metadata":{"execution":{"iopub.status.busy":"2022-10-05T14:09:22.652241Z","iopub.execute_input":"2022-10-05T14:09:22.652684Z","iopub.status.idle":"2022-10-05T14:52:02.700726Z","shell.execute_reply.started":"2022-10-05T14:09:22.652646Z","shell.execute_reply":"2022-10-05T14:52:02.699482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# [0.653830656159886, 0.6238455839150563, 0.605778830856576, 0.6647636130357559, 0.6544317172083116, 0.6532475492188892]\n# 0.6426496583990792\n# Full validation metric: 0.6421152903205289\n\n\n# [0.6455449287371509, 0.6218588657257058, 0.5979860325098842, 0.6420611577083772, 0.675211714053326, 0.6503222829996482]\n# 0.6388308302890154\n# Full validation metric: 0.6404691832962497\n\n#64 5e3\n# [0.6476080393945753, 0.6079157320890402, 0.6050800143556536, 0.6378247315988896, 0.6661450085369098, 0.6504032066047876]\n# 0.635829455429976\n# Full validation metric: 0.6342736651767453\n\n#64 1e3\n# Best validation metric: 0.6557763951171189\n# [0.651014691249416, 0.6347137845717545, 0.6147587091636446, 0.6665997338703076, 0.6587241317996946, 0.6557763951171189]\n# 0.6469312409619894\n# Full validation metric: 0.6467345027007692","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\n\nmodel_files = [file_name for file_name in list(os.listdir('/kaggle/working/models')) if file_name.endswith('.h5')]\nprint(model_files[:5])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import defaultdict\n\n\ncenters_to_models = defaultdict(list)\nfor model_file in model_files:\n    center_id = model_file.split('_')[2]\n    epoch = int(model_file.split('_')[4])\n    metric = float(model_file.split('_')[6][:-3])\n    \n    centers_to_models[center_id].append((metric, epoch, model_file))\n    \ncenter_id_to_best_model_file_name = {\n    center_id: sorted(model_files, key=lambda x: (x[0], x[1]))[0][2]\n    for center_id, model_files in centers_to_models.items()\n}\nprint(center_id_to_best_model_file_name)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"good_models = list(center_id_to_best_model_file_name.values())\nprint(good_models)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for model_file in model_files:\n    if model_file not in good_models:\n        os.remove(os.path.join('/kaggle/working/models', model_file))","metadata":{},"execution_count":null,"outputs":[]}]}