{"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":"%%writefile peak_backprop.py\nfrom copy import deepcopy\n\nimport torch\nimport torch.nn.functional as F\nfrom torch.autograd import Function\n\n\nclass PreHook(Function):\n    \n    @staticmethod\n    def forward(ctx, input, offset):\n        ctx.save_for_backward(input, offset)\n        return input.clone()\n    \n    @staticmethod\n    def backward(ctx, grad_output):\n        input, offset = ctx.saved_variables\n        return (input - offset) * grad_output, None\n\nclass PostHook(Function):\n    \n    @staticmethod\n    def forward(ctx, input, norm_factor):\n        ctx.save_for_backward(norm_factor)\n        return input.clone()\n    \n    @staticmethod\n    def backward(ctx, grad_output):\n        norm_factor, = ctx.saved_variables\n        eps = 1e-10\n        zero_mask = norm_factor < eps\n        grad_input = grad_output / (torch.abs(norm_factor) + eps)\n        grad_input[zero_mask.detach()] = 0\n        return None, grad_input\n\n\ndef pr_conv2d(self, input):\n    offset = input.min().detach()\n    input = PreHook.apply(input, offset)\n    resp = F.conv2d(input, self.weight, self.bias, self.stride, self.padding, self.dilation, self.groups).detach()\n    pos_weight = F.relu(self.weight).detach()\n    norm_factor = F.conv2d(input - offset, pos_weight, None, self.stride, self.padding, self.dilation, self.groups)\n    output = PostHook.apply(resp, norm_factor)\n    return output","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.034829,"end_time":"2021-03-25T15:21:02.671955","exception":false,"start_time":"2021-03-25T15:21:02.637126","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile peak_stimulation.py\nimport torch\nimport torch.nn.functional as F\nfrom torch.autograd import Function\n\n\nclass PeakStimulation(Function):\n\n    @staticmethod\n    def forward(ctx, input, return_aggregation, win_size, peak_filter):\n        ctx.num_flags = 4\n\n        # peak finding\n        assert win_size % 2 == 1, 'Window size for peak finding must be odd.'\n        offset = (win_size - 1) // 2\n        padding = torch.nn.ConstantPad2d(offset, float('-inf'))\n        padded_maps = padding(input)\n        batch_size, num_channels, h, w = padded_maps.size()\n        element_map = torch.arange(0, h * w).long().view(1, 1, h, w)[:, :, offset: -offset, offset: -offset]\n        element_map = element_map.to(input.device)\n        _, indices  = F.max_pool2d(\n            padded_maps,\n            kernel_size = win_size,\n            stride = 1, \n            return_indices = True)\n        peak_map = (indices == element_map)\n\n        # peak filtering\n        if peak_filter:\n            mask = input >= peak_filter(input)\n            peak_map = (peak_map & mask)\n        peak_list = torch.nonzero(peak_map)\n        ctx.mark_non_differentiable(peak_list)\n        \n        # peak aggregation\n        if return_aggregation:\n            peak_map = peak_map.float()\n            ctx.save_for_backward(input, peak_map)\n            return peak_list, (input * peak_map).view(batch_size, num_channels, -1).sum(2) / \\\n                peak_map.view(batch_size, num_channels, -1).sum(2)\n        else:\n            return peak_list\n\n    @staticmethod\n    def backward(ctx, grad_peak_list, grad_output):\n        input, peak_map, = ctx.saved_tensors\n        batch_size, num_channels, _, _ = input.size()\n        grad_input = peak_map * grad_output.view(batch_size, num_channels, 1, 1)\n        return (grad_input,) + (None,) * ctx.num_flags\n\n\ndef peak_stimulation(input, return_aggregation=True, win_size=3, peak_filter=None):\n    return PeakStimulation.apply(input, return_aggregation, win_size, peak_filter)","metadata":{"papermill":{"duration":0.03205,"end_time":"2021-03-25T15:21:02.72829","exception":false,"start_time":"2021-03-25T15:21:02.69624","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile fc_resnet.py\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass FC_ResNet(nn.Module):\n\n    def __init__(self, model, num_classes):\n        super(FC_ResNet, self).__init__()\n\n        # feature encoding\n        self.features = nn.Sequential(\n            model.conv1,\n            model.bn1,\n            model.relu,\n            model.maxpool,\n            model.layer1,\n            model.layer2,\n            model.layer3,\n            model.layer4)\n\n        # classifier\n        num_features = model.layer4[1].conv1.in_channels\n        self.classifier = nn.Sequential(\n            nn.Conv2d(num_features, num_classes, kernel_size=1, bias=True))\n\n    def forward(self, x):\n        x = self.features(x)\n        x = self.classifier(x)\n        return x","metadata":{"papermill":{"duration":0.031946,"end_time":"2021-03-25T15:21:02.784545","exception":false,"start_time":"2021-03-25T15:21:02.752599","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile peak_response_mapping.py\nfrom types import MethodType\n\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\n\nfrom peak_backprop import pr_conv2d\nfrom peak_stimulation import peak_stimulation\n\n\nclass PeakResponseMapping(nn.Sequential):\n\n    def __init__(self, *args, **kargs):\n        super(PeakResponseMapping, self).__init__(*args)\n\n        self.inferencing = False\n        # use global average pooling to aggregate responses if peak stimulation is disabled\n        self.enable_peak_stimulation = kargs.get('enable_peak_stimulation', True)\n        # return only the class response maps in inference mode if peak backpropagation is disabled\n        self.enable_peak_backprop = kargs.get('enable_peak_backprop', True)\n        # window size for peak finding\n        self.win_size = kargs.get('win_size', 3)\n        # sub-pixel peak finding\n        self.sub_pixel_locating_factor = kargs.get('sub_pixel_locating_factor', 1)\n        # peak filtering\n        self.filter_type = kargs.get('filter_type', 'median')\n        if self.filter_type == 'median':\n            self.peak_filter = self._median_filter\n        elif self.filter_type == 'mean':\n            self.peak_filter = self._mean_filter\n        elif self.filter_type == 'max':\n            self.peak_filter = self._max_filter\n        elif isinstance(self.filter_type, (int, float)):\n            self.peak_filter = lambda x: self.filter_type\n        else:\n            self.peak_filter = None\n\n    @staticmethod\n    def _median_filter(input):\n        batch_size, num_channels, h, w = input.size()\n        threshold, _ = torch.median(input.view(batch_size, num_channels, h * w), dim=2)\n        return threshold.contiguous().view(batch_size, num_channels, 1, 1)\n    \n    @staticmethod\n    def _mean_filter(input):\n        batch_size, num_channels, h, w = input.size()\n        threshold = torch.mean(input.view(batch_size, num_channels, h * w), dim=2)\n        return threshold.contiguous().view(batch_size, num_channels, 1, 1)\n    \n    @staticmethod\n    def _max_filter(input):\n        batch_size, num_channels, h, w = input.size()\n        threshold, _ = torch.max(input.view(batch_size, num_channels, h * w), dim=2)\n        return threshold.contiguous().view(batch_size, num_channels, 1, 1)\n\n    def _patch(self):\n        for module in self.modules():\n            if isinstance(module, nn.Conv2d):\n                module._original_forward = module.forward\n                module.forward = MethodType(pr_conv2d, module)\n\n    def _recover(self):\n        for module in self.modules():\n            if isinstance(module, nn.Conv2d) and hasattr(module, '_original_forward'):\n                module.forward = module._original_forward\n\n    def instance_nms(self, instance_list, threshold=0.3, merge_peak_response=True):\n        selected_instances = []\n        while len(instance_list) > 0:\n            instance = instance_list.pop(0)\n            selected_instances.append(instance)\n            src_mask = instance[2].astype(bool)\n            src_peak_response = instance[3]\n            def iou_filter(x):\n                dst_mask = x[2].astype(bool)\n                # IoU\n                intersection = np.logical_and(src_mask, dst_mask).sum()\n                union = np.logical_or(src_mask, dst_mask).sum()\n                iou = intersection / (union + 1e-10)\n                if iou < threshold:\n                    return x\n                else:\n                    if merge_peak_response:\n                        nonlocal src_peak_response\n                        src_peak_response += x[3]\n                    return None\n            instance_list = list(filter(iou_filter, instance_list))\n        return selected_instances\n\n    def instance_seg(self, class_response_maps, peak_list, peak_response_maps, peak_response_scores, retrieval_cfg):        \n        # cast tensors to numpy array\n        class_response_maps = class_response_maps.squeeze().cpu().numpy()\n        peak_list = peak_list.cpu().numpy()\n        peak_response_maps = peak_response_maps.cpu().numpy()\n        peak_response_scores = peak_response_scores.detach().cpu().numpy()\n\n        img_height, img_width = peak_response_maps.shape[1], peak_response_maps.shape[2]\n        \n        # image size\n        img_area = img_height * img_width\n\n        # segment proposals off-the-shelf\n        proposals = retrieval_cfg['proposals']\n\n        # proposal contour width\n        contour_width = retrieval_cfg.get('contour_width', 5)\n\n        # limit range of proposal size\n        proposal_size_limit = retrieval_cfg.get('proposal_size_limit', (0.00002, 0.85))\n\n        # selected number of proposals\n        proposal_count = retrieval_cfg.get('proposal_count', 100)\n\n        # nms threshold\n        nms_threshold = retrieval_cfg.get('nms_threshold', 0.3)\n        \n        # merge peak response during nms\n        merge_peak_response = retrieval_cfg.get('merge_peak_response', False)\n\n        # metric free parameters\n        param = retrieval_cfg.get('param', None)\n\n        # process each peak\n        instance_list = []\n        for i in range(len(peak_response_maps)):\n            class_idx = peak_list[i, 1]\n\n            # extract hyper-params\n            if isinstance(param, tuple):\n                # shared param\n                bg_threshold_factor, penalty_factor, balance_factor = param\n            elif isinstance(param, list):\n                # independent params between classes\n                bg_threshold_factor, penalty_factor, balance_factor = param[class_idx]\n            else:\n                raise TypeError('Invalid hyper-params \"%s\".' % param)\n            \n            class_response = cv2.resize(src=class_response_maps[class_idx], dsize=(img_height, img_width), interpolation=cv2.INTER_CUBIC)\n            bg_response = (class_response < bg_threshold_factor * class_response.mean()).astype(np.float32)\n            peak_response_map = peak_response_maps[i]\n            peak_response_score = peak_response_scores[i]\n            # select proposal\n            max_val = -np.inf\n            instance_mask = None\n\n            for j in range(min(proposal_count, len(proposals))):\n                raw_mask = cv2.resize(src=proposals[j].astype(int), dsize=peak_response_map.shape, interpolation=cv2.INTER_NEAREST)\n                # get contour of the proposal\n                contour_mask = cv2.morphologyEx(raw_mask.astype(\"uint8\"), cv2.MORPH_GRADIENT, np.ones((contour_width, contour_width), np.uint8)).astype(bool)\n                mask = raw_mask.astype(bool)\n                # metric\n                mask_area = mask.sum()\n                if (mask_area >= proposal_size_limit[1] * img_area) or \\\n                    (mask_area < proposal_size_limit[0] * img_area):\n                    continue\n                else:\n                    val = balance_factor * peak_response_map[mask].sum() + \\\n                        peak_response_map[contour_mask].sum() - \\\n                        penalty_factor * bg_response[mask].sum()\n                    if val > max_val:\n                        max_val = val\n                        instance_mask = mask\n            \n            if instance_mask is not None:\n                instance_list.append((max_val, class_idx, instance_mask, peak_response_map, peak_response_score))\n\n        instance_list = sorted(instance_list, key=lambda x: x[0], reverse=True)\n        if nms_threshold is not None:\n            instance_list = self.instance_nms(sorted(instance_list, key=lambda x: x[0], reverse=True), nms_threshold, merge_peak_response)\n        return [dict(category=v[1], mask=v[2], prm=v[3], score=v[4]) for v in instance_list]\n\n    def forward(self, input, class_threshold=0, peak_threshold=0.0, retrieval_cfg=None):\n        assert input.dim() == 4, 'PeakResponseMapping layer only supports batch mode.'\n        if self.inferencing:\n            input.requires_grad_()\n\n        # classification network forwarding\n        class_response_maps = super(PeakResponseMapping, self).forward(input)\n        if self.enable_peak_stimulation:\n            # sub-pixel peak finding\n            if self.sub_pixel_locating_factor > 1:\n                class_response_maps = F.upsample(class_response_maps, scale_factor=self.sub_pixel_locating_factor, mode='bilinear', align_corners=True)\n            # aggregate responses from informative receptive fields estimated via class peak responses\n            peak_list, aggregation = peak_stimulation(class_response_maps, win_size=self.win_size, peak_filter=self.peak_filter)\n        else:\n            # aggregate responses from all receptive fields\n            peak_list, aggregation = None, F.adaptive_avg_pool2d(class_response_maps, 1).squeeze(2).squeeze(2)\n\n        if self.inferencing:\n            if not self.enable_peak_backprop:\n                # extract only class-aware visual cues\n                return aggregation, class_response_maps\n            \n            # extract instance-aware visual cues, i.e., peak response maps\n            assert class_response_maps.size(0) == 1, 'Currently inference mode (with peak backpropagation) only supports one image at a time.'\n            if peak_list is None:\n                peak_list = peak_stimulation(class_response_maps, return_aggregation=False, win_size=self.win_size, peak_filter=self.peak_filter)\n\n            peak_response_maps = []\n            peak_response_scores = []\n            valid_peak_list = []\n            # peak backpropagation\n            grad_output = class_response_maps.new_empty(class_response_maps.size())\n            for idx in range(peak_list.size(0)):\n                if aggregation[peak_list[idx, 0], peak_list[idx, 1]] >= class_threshold:\n                    peak_val = class_response_maps[peak_list[idx, 0], peak_list[idx, 1], peak_list[idx, 2], peak_list[idx, 3]]\n#                     if peak_val > peak_threshold:\n                    grad_output.zero_()\n                    # starting from the peak\n                    grad_output[peak_list[idx, 0], peak_list[idx, 1], peak_list[idx, 2], peak_list[idx, 3]] = 1\n                    if input.grad is not None:\n                        input.grad.zero_()\n                    class_response_maps.backward(grad_output, retain_graph=True)\n                    prm = input.grad.detach().sum(1).clone().clamp(min=0)\n                    peak_response_maps.append(prm / prm.sum())\n                    peak_response_scores.append(torch.sigmoid(peak_val))\n                    valid_peak_list.append(peak_list[idx, :])\n\n            # return results\n            class_response_maps = class_response_maps.detach()\n            aggregation = aggregation.detach()\n\n            if len(peak_response_maps) > 0:\n                valid_peak_list = torch.stack(valid_peak_list)\n                peak_response_maps = torch.cat(peak_response_maps, 0)\n                peak_response_scores = torch.stack(peak_response_scores)\n                if retrieval_cfg is None:\n                    # classification confidence scores, class-aware and instance-aware visual cues\n                    return aggregation, class_response_maps, valid_peak_list, peak_response_maps, peak_response_scores\n                else:\n                    # instance segmentation using build-in proposal retriever\n                    return self.instance_seg(class_response_maps, valid_peak_list, peak_response_maps, peak_response_scores, retrieval_cfg)\n            else:\n                return []\n        else:\n            # classification confidence scores\n            return aggregation\n\n    def train(self, mode=True):\n        super(PeakResponseMapping, self).train(mode)\n        if self.inferencing:\n            self._recover()\n            self.inferencing = False\n        return self\n\n    def inference(self):\n        super(PeakResponseMapping, self).train(False)\n        self._patch()\n        self.inferencing = True\n        return self","metadata":{"papermill":{"duration":0.034739,"end_time":"2021-03-25T15:21:02.844416","exception":false,"start_time":"2021-03-25T15:21:02.809677","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile prm.py\nfrom typing import Union, Optional, List, Tuple\n\nimport os\nimport cv2\nimport torch\nimport shutil\nimport numpy as np\nimport torch.nn as nn\nfrom matplotlib.colors import hsv_to_rgb\nfrom scipy.ndimage import center_of_mass\n\nfrom fc_resnet import FC_ResNet \nfrom peak_response_mapping import PeakResponseMapping\n\n\ndef fc_resnest(cache_root=\".\", repo=\"zhanghang1989/ResNeSt\", model=\"resnest101\", num_classes: int = 20, pretrained: bool = True) -> nn.Module:\n    dir_name = \"_\".join(repo.split(\"/\")+[\"master\"])\n    if not os.path.exists(dir_name) and os.path.exists(os.path.join(cache_root, f\"{dir_name}.zip\")):\n        shutil.unpack_archive(os.path.join(cache_root, f\"{dir_name}.zip\"), dir_name, \"zip\")\n    torch.hub.set_dir(\".\")\n    model = FC_ResNet(torch.hub.load(repo, \"resnest101\", pretrained=pretrained), num_classes)\n    shutil.make_archive(dir_name, \"zip\", dir_name)\n    shutil.rmtree(dir_name, ignore_errors=False)\n    return model\n\ndef peak_response_mapping(\n    backbone: nn.Module,\n    enable_peak_stimulation: bool = True,\n    enable_peak_backprop: bool = True,\n    win_size: int = 3,\n    sub_pixel_locating_factor: int = 1,\n    filter_type: Union[str, int, float] = 'median') -> nn.Module:\n    \"\"\"Peak Response Mapping.\n    \"\"\"\n\n    model = PeakResponseMapping(\n        backbone, \n        enable_peak_stimulation = enable_peak_stimulation,\n        enable_peak_backprop = enable_peak_backprop, \n        win_size = win_size, \n        sub_pixel_locating_factor = sub_pixel_locating_factor, \n        filter_type = filter_type)\n    return model\n\ndef prm_visualize(\n    instance_list: List[dict], \n    class_names: Optional[List[str]]=None,\n    font_scale: Union[int, float] = 1) -> Tuple[np.ndarray, np.ndarray]:\n    \"\"\"Prediction visualization.\n    \"\"\"\n\n    # helper functions\n    def rgb2hsv(r, g, b):\n        mx = max(r, g, b)\n        mn = min(r, g, b)\n        df = mx - mn\n        if mx == mn:\n            h = 0\n        elif mx == r:\n            h = (60 * ((g - b) / df) + 360) % 360\n        elif mx == g:\n            h = (60 * ((b - r) / df) + 120) % 360\n        elif mx == b:\n            h = (60 * ((r - g) / df) + 240) % 360\n        s = 0 if mx == 0 else df / mx\n        v = mx\n        return h / 360.0, s, v\n\n    def color_palette(N):\n        cmap = np.zeros((N, 3))\n        for i in range(0, N):\n            uid = i\n            r, g, b = 0, 0, 0\n            for j in range(0, 8):\n                r = np.bitwise_or(r, (((uid & (1 << 0)) != 0) << 7 - j))\n                g = np.bitwise_or(g, (((uid & (1 << 1)) != 0) << 7 - j))\n                b = np.bitwise_or(b, (((uid & (1 << 2)) != 0) << 7 - j))\n                uid = (uid >> 3)\n            cmap[i, 0] = min(r + 86, 255)\n            cmap[i, 1] = min(g + 86, 255)\n            cmap[i, 2] = b\n        cmap = cmap.astype(np.float32) / 255\n        return cmap\n\n    if len(instance_list) > 0:\n        palette = color_palette(len(instance_list) + 1)\n        height, width = instance_list[0]['mask'].shape[0], instance_list[0]['mask'].shape[1]\n        instance_mask = np.zeros((height, width, 3), dtype=np.float32)\n        peak_response_map = np.zeros((height, width, 3), dtype=np.float32)\n        for idx, pred in enumerate(instance_list):\n            category, mask, prm = pred['category'], pred['mask'], pred['prm']\n            # instance masks\n            instance_mask[mask, 0] = palette[idx + 1][0]\n            instance_mask[mask, 1] = palette[idx + 1][1]\n            instance_mask[mask, 2] = palette[idx + 1][2]\n            if class_names is not None:\n                y, x = center_of_mass(mask)\n                y, x = int(y), int(x)\n                text = class_names[category]\n                font_face = cv2.FONT_HERSHEY_SIMPLEX\n                thickness = 2\n                text_size, _ = cv2.getTextSize(text, font_face, font_scale, thickness)\n                cv2.putText(\n                    instance_mask,\n                    text,\n                    (x - text_size[0] // 2, y),\n                    font_face,\n                    font_scale,\n                    (1., 1., 1.),\n                    thickness)\n            # peak response map\n            peak_response = (prm - prm.min()) / (prm.max() - prm.min())\n            mask = peak_response > 0.01\n            h, s, _ = rgb2hsv(palette[idx + 1][0], palette[idx + 1][1], palette[idx + 1][2])\n            peak_response_map[mask, 0] = h\n            peak_response_map[mask, 1] = s\n            peak_response_map[mask, 2] = np.power(peak_response[mask], 0.5)\n\n        peak_response_map =  hsv_to_rgb(peak_response_map)\n        return instance_mask, peak_response_map","metadata":{"papermill":{"duration":0.032871,"end_time":"2021-03-25T15:21:02.902395","exception":false,"start_time":"2021-03-25T15:21:02.869524","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport re\nimport prm\nimport cv2\nimport zlib\nimport torch\nimport shutil\nimport base64\nimport skimage\nimport ipywidgets\nimport numpy as np\nimport pandas as pd\nfrom torch import nn\nimport albumentations as A\nfrom IPython import display\nimport scipy.ndimage as ndi\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data.sampler import SubsetRandomSampler\nfrom skmultilearn.model_selection import IterativeStratification","metadata":{"papermill":{"duration":3.595422,"end_time":"2021-03-25T15:21:06.523077","exception":false,"start_time":"2021-03-25T15:21:02.927655","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training = 0\nnum_folds = 3\nnum_epoch = 20\nbatch_size = 32\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nkernel_data_root = os.path.join(\"/\", \"kaggle\", \"input\", \"human-protein-atlas-prm\")\ncompetition_root = os.path.join(\"/\", \"kaggle\", \"input\", \"hpa-single-cell-image-classification\")","metadata":{"papermill":{"duration":0.483282,"end_time":"2021-03-25T15:21:07.031558","exception":false,"start_time":"2021-03-25T15:21:06.548276","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pip_install(pkg, cache_root=kernel_data_root):\n    dir_name = f\"{pkg}_master\"\n    if not os.path.exists(dir_name) and os.path.exists(os.path.join(cache_root, f\"{dir_name}.zip\")):\n        shutil.unpack_archive(os.path.join(cache_root, f\"{dir_name}.zip\"), dir_name, \"zip\")\n    else:\n        os.system(f\"pip wheel {pkg} -w {dir_name}\")\n    os.system(f\"pip install {pkg} --no-index --find-links {dir_name}\")\n    shutil.make_archive(dir_name, \"zip\", dir_name)\n    shutil.rmtree(dir_name, ignore_errors=False)\n\ndef git_install(user, repo, cache_root=kernel_data_root):\n    dir_name = f\"{user}_{repo}_master\"\n    if not os.path.exists(dir_name) and os.path.exists(os.path.join(cache_root, f\"{dir_name}.zip\")):\n        shutil.unpack_archive(os.path.join(cache_root, f\"{dir_name}.zip\"), dir_name, \"zip\")\n    else:\n        os.system(f\"git clone https://github.com/{user}/{repo} {dir_name}\")\n        with open(f\"./{dir_name}/setup.py\", \"r\") as f:\n            setup = re.sub(\"@https://github.com[a-z0-9/_.]+\", \"\", f.read())\n        with open(f\"./{dir_name}/setup.py\", \"w\") as f:\n            f.write(setup)\n    os.system(f\"pip install ./{dir_name}\")\n    shutil.make_archive(dir_name, \"zip\", dir_name)\n    shutil.rmtree(dir_name, ignore_errors=False)\n\npip_install(\"pycocotools\")\ngit_install(\"facebookresearch\", \"iopath\")\ngit_install(\"facebookresearch\", \"fvcore\")\ngit_install(\"haoxusci\", \"pytorch_zoo\")\ngit_install(\"CellProfiling\", \"HPA-Cell-Segmentation\")\nfrom pycocotools import _mask as coco_mask\nfrom hpacellseg.cellsegmentator import CellSegmentator","metadata":{"papermill":{"duration":380.591087,"end_time":"2021-03-25T15:27:27.664255","exception":false,"start_time":"2021-03-25T15:21:07.073168","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, alpha=0.75, gamma=2, reduction=\"mean\"):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n        self.function = nn.BCEWithLogitsLoss(reduction=\"none\")\n\n    def forward(self, inputs, targets):\n        assert inputs.size() == targets.size() and targets.numel() > 0\n\n        p = torch.sigmoid(inputs)\n        ce_loss = self.function(inputs, targets)\n        p_t = p * targets + (1 - p) * (1 - targets)\n        loss = ce_loss * ((1 - p_t) ** self.gamma)\n\n        if self.alpha >= 0:\n            alpha_t = self.alpha * targets + (1 - self.alpha) * (1 - targets)\n            loss = alpha_t * loss\n\n        if self.reduction == \"mean\":\n            loss = loss.mean()\n        elif self.reduction == \"sum\":\n            loss = loss.sum()\n\n        return loss\n\nclass F1Score(nn.Module):\n    def __init__(self, epsilon=1e-7):\n        super().__init__()\n        self.epsilon = epsilon\n\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n\n        tp = (targets * probs).sum(dim=0).to(torch.float32)\n        tn = ((1 - targets) * (1 - probs)).sum(dim=0).to(torch.float32)\n        fp = ((1 - targets) * probs).sum(dim=0).to(torch.float32)\n        fn = (targets * (1 - probs)).sum(dim=0).to(torch.float32)\n\n        precision = tp / (tp + fp + self.epsilon)\n        recall = tp / (tp + fn + self.epsilon)\n\n        f1 = 2* (precision*recall) / (precision + recall + self.epsilon)\n        f1 = f1.clamp(min=self.epsilon, max=1-self.epsilon)\n        return f1.mean()\n","metadata":{"papermill":{"duration":0.040313,"end_time":"2021-03-25T15:27:27.729759","exception":false,"start_time":"2021-03-25T15:27:27.689446","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LabelTensor(nn.Module):\n    def __init__(self, structure=None):\n        super().__init__()\n        s = torch.tensor([\n            [ 1,  1,  1],\n            [ 1,  0, -1],\n            [-1, -1, -1]\n        ]) if structure is None else structure\n        self.register_buffer(\"s\", s)\n#         assert self.s.size(0) == self.s.size(1) == 3, \"Structure should be (3, 3) in size\"\n        self.pad = nn.ConstantPad2d((self.s.size(0)-1) // 2, 0)\n\n    def forward(self, x):\n        x = self.pad(x)\n        num_objs = 0\n        for i in range(x.size(-2)-self.s.size(-2)+1):\n            for j in range(x.size(-1)-self.s.size(-1)+1):\n                mat = x[..., i:i+self.s.size(-2), j:j+self.s.size(-1)] * self.s\n                ci, cj = i+(self.s.size(-2)-1)//2, j+(self.s.size(-1)-1)//2\n                num_objs += mat.max().eq(0) * x[..., ci, cj]\n                x[..., ci, cj] *= num_objs * mat.max().eq(0) + mat.max()\n        l, r, t, b = self.pad.padding\n        x = x[..., l: -r, t: -b]\n        return x\n\ndef visualize(image, labels, ax=None, title='', **kwargs):\n    if ax is None: _, ax = plt.subplots(1, 1, figsize = (10, 10))\n    if title: ax.set_title(title)\n    ax.set_title(title)\n    ax.imshow(image[:3].transpose(1, 2, 0) if len(image.shape) > 2 else image, **kwargs)\n    ax.set_xlabel(f\"Labels: {'|'.join(map(str,labels.nonzero()[0])) if isinstance(labels, np.ndarray) else labels}\")\n\ndef plot_image(competition_root, folder, image_id, ax=None):\n    if ax is None: _, ax = plt.subplots(1, 1, figsize = (10, 10))\n    colors = [\"red\", \"green\", \"blue\", \"yellow\"]\n    image_path = os.path.join(competition_root, folder, f\"{image_id}_{{}}.png\")\n    image = [cv2.imread(image_path.format(c), cv2.IMREAD_GRAYSCALE) for c in colors]\n#     image = [np.array(Image.open(image_path.format(c).convert(\"L\"))) for c in colors]\n    ax.imshow(np.stack(image, axis=-1)[..., [0, 1, 2]])\n    ax.set_xlabel(folder)","metadata":{"papermill":{"duration":0.042221,"end_time":"2021-03-25T15:27:27.796882","exception":false,"start_time":"2021-03-25T15:27:27.754661","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_, axs = plt.subplots(1, 2, figsize=(10, 5))\nplot_image(competition_root, \"train\", \"000a6c98-bb9b-11e8-b2b9-ac1f6b6435d0\", axs[0])\nplot_image(competition_root, \"test\", \"004a270d-34a2-4d60-bbe4-365fca868193\", axs[1])\nplt.show()","metadata":{"papermill":{"duration":2.20573,"end_time":"2021-03-25T15:27:30.027753","exception":false,"start_time":"2021-03-25T15:27:27.822023","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Oversampler:\n    def __init__(self, num_classes=19):\n        self.num_classes = num_classes\n\n    def __call__(self, df):\n        factor = np.zeros((self.num_classes,))\n        for idx, (_, row)  in enumerate(df.iterrows(), start=1):\n            factor[list(map(int, row.Label.split(\"|\")))] += 1\n        factor = np.clip(((idx / self.num_classes) / factor).round(), a_min=1, a_max=None)\n        return factor.astype(np.int)","metadata":{"papermill":{"duration":0.039919,"end_time":"2021-03-25T15:27:30.098857","exception":false,"start_time":"2021-03-25T15:27:30.058938","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainDataset:\n    def __init__(self, competition_root, kernel_data_root, num_class=19, resize=(256, 256)):\n        if os.path.exists(os.path.join(kernel_data_root, \"train.zip\")) and not os.path.exists(\"train\"):\n            shutil.unpack_archive(os.path.join(kernel_data_root, \"train.zip\"), \"train\", \"zip\")\n\n        self.data = []\n        colors = ['red', 'green', 'blue', 'yellow']\n        df = pd.read_csv(os.path.join(competition_root, \"train.csv\"))\n        oversampler = Oversampler(num_class)\n        factor = oversampler(df)\n\n        for i, row in tqdm(df.iterrows(), total = len(df)):\n            labels = np.zeros((num_class,))\n            labels[list(map(int, row.Label.split(\"|\")))] = 1\n            if not os.path.exists('train'): os.makedirs('train')\n            if not os.path.exists(os.path.join(\"train\", f\"{row.ID}.npy\")):\n                image_path = os.path.join(competition_root, \"train\", f\"{row.ID}_{{}}.png\")\n                image = [cv2.imread(image_path.format(c), cv2.IMREAD_GRAYSCALE) for c in colors]\n                image = np.stack([cv2.resize(c, resize) for c in image], axis=-1)\n                np.save(os.path.join(\"train\", f\"{row.ID}.npy\"), image)\n            self.data += [{\n                \"image_id\": row.ID,\n                \"labels\": labels,\n                \"image_path\": os.path.join(\"train\", f\"{row.ID}.npy\")\n            }] * factor[labels.astype(np.bool)].max()\n\n        shutil.make_archive(\"train\", \"zip\", \"train\")\n\n        self.transforms = A.Compose([\n            A.Resize(width=resize[1], height=resize[0]),\n            A.HorizontalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.RandomBrightnessContrast(p=0.2)\n        ])\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        image_path = self.data[idx][\"image_path\"]\n        labels = self.data[idx][\"labels\"]\n        image = np.load(image_path, allow_pickle=True)\n        image = self.transforms(image=image)[\"image\"]#[..., :3]\n        image = image.transpose(2, 0, 1).astype(np.float32)\n        image = np.divide(image, 255.0)\n        return image, labels","metadata":{"papermill":{"duration":0.147279,"end_time":"2021-03-25T15:27:30.276338","exception":false,"start_time":"2021-03-25T15:27:30.129059","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntrain_dataset = TrainDataset(competition_root, kernel_data_root)","metadata":{"papermill":{"duration":139.914581,"end_time":"2021-03-25T15:29:50.221663","exception":false,"start_time":"2021-03-25T15:27:30.307082","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx = np.random.randint(len(train_dataset))\nvisualize(*train_dataset[idx])","metadata":{"papermill":{"duration":0.357042,"end_time":"2021-03-25T15:29:50.611769","exception":false,"start_time":"2021-03-25T15:29:50.254727","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if os.path.exists(os.path.join(kernel_data_root, 'checkpoint.pth')):\n    checkpoint = torch.load(os.path.join(kernel_data_root, 'checkpoint.pth'), map_location = device)\n    torch.save(checkpoint, \"checkpoint.pth\")\nelse:\n    checkpoint = {}","metadata":{"papermill":{"duration":22.621346,"end_time":"2021-03-25T15:30:13.273265","exception":false,"start_time":"2021-03-25T15:29:50.651919","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loaders = {}\nvalid_loaders = {}\ntrain_folds = checkpoint.get('train_folds', {})\nvalid_folds = checkpoint.get('valid_folds', {})\niss = IterativeStratification(n_splits=num_folds, order=1)\nsplitter = iss.split(train_dataset, np.array([x[\"labels\"] for x in train_dataset.data]))\nfor fold, (train_indices, valid_indices) in enumerate(splitter):\n    train_folds[fold] = train_folds.get(fold, train_indices)\n    valid_folds[fold] = valid_folds.get(fold, valid_indices)\n    # Creating PT data samplers and loaders\n    train_sampler = SubsetRandomSampler(train_folds[fold])\n    valid_sampler = SubsetRandomSampler(valid_folds[fold])\n    train_loaders[fold] = DataLoader(train_dataset, batch_size=batch_size, sampler=train_sampler, drop_last=True)\n    valid_loaders[fold] = DataLoader(train_dataset, batch_size=batch_size, sampler=valid_sampler, drop_last=True)","metadata":{"papermill":{"duration":2.894675,"end_time":"2021-03-25T15:30:16.207643","exception":false,"start_time":"2021-03-25T15:30:13.312968","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = {}\noptimizers = {}\nschedulers = {}\n\nmetric = F1Score()\ncriterion = FocalLoss()\n\nfor fold in tqdm(range(num_folds)):\n    backbone = prm.fc_resnest(kernel_data_root, num_classes=19, pretrained=False)\n    models[fold] = prm.peak_response_mapping(backbone).to(device)\n    models[fold].load_state_dict(checkpoint.get(\"models\", {}).get(fold, models[fold].state_dict()))\n    # Prepare optimizer and schedule (linear warmup and decay)    \n    params = [p for n, p in models[fold].named_parameters() if p.requires_grad]\n    optimizers[fold] = torch.optim.Adam(params, lr=1e-3, weight_decay=1e-6)\n    optimizers[fold].load_state_dict(checkpoint.get(\"optimizers\", {}).get(fold, optimizers[fold].state_dict()))\n    schedulers[fold] = torch.optim.lr_scheduler.StepLR(optimizers[fold], step_size=50, gamma=0.1, last_epoch=-1)\n    schedulers[fold].load_state_dict(checkpoint.get(\"schedulers\", {}).get(fold, schedulers[fold].state_dict()))","metadata":{"papermill":{"duration":13.910966,"end_time":"2021-03-25T15:30:30.159206","exception":false,"start_time":"2021-03-25T15:30:16.24824","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(fold):\n    total_loss, total_score = 0, 0\n    models[fold].train()\n    loader = tqdm(train_loaders[fold], desc = f\"Training fold {fold+1}\")\n    for idx, (images, labels) in enumerate(loader, start=1):\n        images, labels = images[:, :3].to(device), labels.to(device)\n        # Execute\n        logits = models[fold](images)\n        loss = criterion(logits, labels)\n        score = metric(logits, labels)\n        total_loss += loss.item(); total_score += score.item()\n        # Optimize + Backward\n        optimizers[fold].zero_grad()\n        loss.backward()\n        optimizers[fold].step()\n        # print statistics\n        loader.set_postfix_str(f\"Score: {score:.4f} | Loss: {loss:.4f}\")\n        loader.update()\n        # Clear variable\n        del images; del labels; del logits; del loss; del score\n        torch.cuda.empty_cache()\n    print(f\"Trained fold {fold+1} | Score: {total_score/idx:.4f} | Loss: {total_loss/idx:.4f}\")\n    schedulers[fold].step()\n    return total_score/idx, total_loss/idx","metadata":{"papermill":{"duration":0.050934,"end_time":"2021-03-25T15:30:30.251183","exception":false,"start_time":"2021-03-25T15:30:30.200249","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def valid(fold):\n    total_loss, total_score = 0, 0\n    models[fold].eval()\n    loader = tqdm(valid_loaders[fold], desc = f\"Validating fold {fold+1}\")\n    for idx, (images, labels) in enumerate(loader, start=1):\n        images, labels = images[:, :3].to(device), labels.to(device)\n        # Execute\n        with torch.no_grad():\n            logits = models[fold](images, labels)\n        loss = criterion(logits, labels)\n        score = metric(logits, labels)\n        total_loss += loss.item(); total_score += score.item()\n        # print statistics\n        loader.set_postfix_str(f\"Score: {score:.4f} | Loss: {loss:.4f}\")\n        loader.update()\n        # Clear variable\n        del images; del labels; del logits; del loss; del score\n        torch.cuda.empty_cache()\n    print(f\"Validated fold {fold+1} | Score: {total_score/idx:.4f} | Loss: {total_loss/idx:.4f}\")\n    return total_score/idx, total_loss/idx","metadata":{"papermill":{"duration":0.050218,"end_time":"2021-03-25T15:30:30.341139","exception":false,"start_time":"2021-03-25T15:30:30.290921","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = checkpoint.get('train_data', {})\nvalid_data = checkpoint.get('valid_data', {})\nepoch_data = checkpoint.get('epoch_data', [])\n# del checkpoint; torch.cuda.empty_cache()\nif epoch_data: fig, axs = plt.subplots(num_folds, 2, figsize=(10*2, 5*num_folds))\nfor fold in range(num_folds):\n    if fold in valid_data and fold in train_data:\n        # Visualize\n        fold_axs = axs if num_folds <= 1 else axs[fold]\n        fold_axs[0].clear(); fold_axs[1].clear()\n        fold_axs[0].plot(epoch_data, train_data[fold][:, 0], label = f\"Train fold {fold+1} score {train_data[fold][-1, 0]:.4f}\")\n        fold_axs[0].plot(epoch_data, valid_data[fold][:, 0], label = f\"Valid fold {fold+1} score {valid_data[fold][-1, 0]:.4f}\")\n        fold_axs[1].plot(epoch_data, train_data[fold][:, 1], label = f\"Train fold {fold+1} loss {train_data[fold][-1, 1]:.4f}\")\n        fold_axs[1].plot(epoch_data, valid_data[fold][:, 1], label = f\"Valid fold {fold+1} loss {valid_data[fold][-1, 1]:.4f}\")\n        fold_axs[0].legend(); fold_axs[1].legend()\nplt.show()","metadata":{"papermill":{"duration":1.003537,"end_time":"2021-03-25T15:30:31.384458","exception":false,"start_time":"2021-03-25T15:30:30.380921","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loader = tqdm(range(len(epoch_data), len(epoch_data) + num_epoch * training), desc = \"Epoch\")\nboard = ipywidgets.Output()\nif training: display.display(board)\ngraph = display.display(ipywidgets.widgets.HTML(f\"<b>Training statred: {bool(training)}</b>\"), display_id = True)\nfor i in loader:\n    with board:\n        epoch_data.append(i+1)\n        # Make grid\n        fig, axs = plt.subplots(num_folds, 2, figsize=(10*2, 5*num_folds))\n        # Close figure\n        plt.close(fig)\n        for fold in range(num_folds):\n            train_fold, valid_fold = train(fold), valid(fold)\n            train_data[fold] = np.append(train_data.get(fold, np.empty((0, 2))), [train_fold], axis = 0)\n            valid_data[fold] = np.append(valid_data.get(fold, np.empty((0, 2))), [valid_fold], axis = 0)\n            # Visualize\n            fold_axs = axs if num_folds <= 1 else axs[fold]\n            fold_axs[0].clear(); fold_axs[1].clear()\n            fold_axs[0].plot(epoch_data, train_data[fold][:, 0], label = f\"Train fold {fold+1} score {train_data[fold][-1, 0]:.4f}\")\n            fold_axs[0].plot(epoch_data, valid_data[fold][:, 0], label = f\"Valid fold {fold+1} score {valid_data[fold][-1, 0]:.4f}\")\n            fold_axs[1].plot(epoch_data, train_data[fold][:, 1], label = f\"Train fold {fold+1} loss {train_data[fold][-1, 1]:.4f}\")\n            fold_axs[1].plot(epoch_data, valid_data[fold][:, 1], label = f\"Valid fold {fold+1} loss {valid_data[fold][-1, 1]:.4f}\")\n            fold_axs[0].legend(); fold_axs[1].legend()\n            graph.update(fig)\n        # Clear all progress bar with in board widget\n        display.clear_output()\n        graph = display.display(fig, display_id = True)\n    # Save model\n    params = {\n        'models': dict([(fold, models[fold].state_dict()) for fold in models]),\n        'optimizers': dict([(fold, optimizers[fold].state_dict()) for fold in optimizers]),\n        'schedulers': dict([(fold, schedulers[fold].state_dict()) for fold in schedulers]),\n        'train_folds': train_folds,\n        'valid_folds': valid_folds,\n        'train_data': train_data,\n        'valid_data': valid_data,\n        'epoch_data': epoch_data\n    }\n    torch.save(params, \"checkpoint.pth\")\nloader.write(\"Done!\")","metadata":{"papermill":{"duration":0.112404,"end_time":"2021-03-25T15:30:31.542786","exception":false,"start_time":"2021-03-25T15:30:31.430382","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx = np.random.randint(len(train_dataset))\nfold = np.random.randint(num_folds)\nimage, labels = train_dataset[idx]\nmodels[fold].eval()\nlogits = models[fold](torch.tensor(image[:3]).unsqueeze(0).to(device))\npred_labels = torch.sigmoid(logits[0]).gt(0.5).nonzero(as_tuple=True)[0].cpu().numpy()\n_, axs = plt.subplots(1, 2, figsize=(15, 5))\nvisualize(image, labels, axs[0], \"Original\")\nvisualize(image, pred_labels, axs[1], \"Predicted\")","metadata":{"papermill":{"duration":1.02575,"end_time":"2021-03-25T15:30:32.614344","exception":false,"start_time":"2021-03-25T15:30:31.588594","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models[fold].inference()\nvisual_cues =  models[fold](torch.tensor(image[:3]).unsqueeze(0).to(device))\nif visual_cues is None:\n    print('No class peak response detected')\nelse:\n    confidence, class_response_maps, class_peak_responses, peak_response_maps, peak_response_scores = visual_cues\n    _, class_idx = torch.max(confidence, dim=1)\n    class_idx = class_idx.item()\n    num_plots = 2 + len(peak_response_maps)\n    fig, axs = plt.subplots(num_plots//4+1, 4, figsize=(5 * 4, 5*(num_plots//4)+1))\n    visualize(image, labels, axs[0, 0], \"Image\")\n    visualize(class_response_maps[0, class_idx].cpu(), class_idx, axs[0, 1], 'Class Response Map', interpolation='bicubic')\n    for idx, (peak_response_map, peak) in enumerate(sorted(zip(peak_response_maps, class_peak_responses), key=lambda v: v[-1][-1])):\n        visualize(peak_response_map.cpu(), peak[1].item(), axs[(idx + 2)//4, (idx + 2)%4], \"Peak Response Map\", cmap=\"jet\")","metadata":{"papermill":{"duration":2.276936,"end_time":"2021-03-25T15:30:34.940763","exception":false,"start_time":"2021-03-25T15:30:32.663827","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"crms = models[fold][0](torch.tensor(image[:3]).unsqueeze(0).to(device))\nscores = torch.sigmoid(logits)\ncrms = nn.functional.interpolate(crms, (256, 256), mode='bilinear', align_corners=False)\n# Plot prms grid with original image\n_, axs = plt.subplots(4, 5, figsize=(25, 20))\nvisualize(image, labels, axs[0, 0], \"Original image\")\nfor i, (crm, cscore) in enumerate(zip(crms.squeeze(0), scores.squeeze(0))):\n    visualize(crm.detach().cpu().unsqueeze(0).cpu().numpy(), i, axs[(i+1)//5, (i+1)%5], f\"Class score{cscore:.2f}\", cmap=\"jet\")","metadata":{"papermill":{"duration":0.143086,"end_time":"2021-03-25T15:30:38.148378","exception":false,"start_time":"2021-03-25T15:30:38.005292","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CellSegment(CellSegmentator):\n    def __init__(self, device, cache_root=\".\", *args, **kwargs):\n        if os.path.exists(os.path.join(kernel_data_root, \"cell_segmentator.zip\")) and not os.path.exists(\"cell_segmentator\"):\n            shutil.unpack_archive(os.path.join(kernel_data_root, \"cell_segmentator.zip\"), \"cell_segmentator\", \"zip\")\n\n        if not os.path.exists('cell_segmentator'): os.makedirs('cell_segmentator')\n        nuclei_model = 'cell_segmentator/nuclei_model.pth'\n        cell_model = 'cell_segmentator/cell_model.pth'\n        super().__init__(nuclei_model=nuclei_model,\n                         cell_model=cell_model,\n                         scale_factor=1,\n                         device=device,\n                         padding=False,\n                         multi_channel_model=True)\n\n        shutil.make_archive(\"cell_segmentator\", \"zip\", \"cell_segmentator\")\n        shutil.rmtree(\"cell_segmentator\", ignore_errors=False)\n\n    def build_image_names(self, image_root, image_ids) -> list:\n        colors = [\"red\", \"yellow\", \"blue\"]\n        image_paths = [os.path.join(image_root, f\"{image_id}_{{}}.png\") for image_id in image_ids]\n        image = [[image_path.format(c) for image_path in image_paths] for c in colors]\n        return image\n\n    def __fill_holes(self, image):\n        \"\"\"Fill_holes for labelled image, with a unique number.\"\"\"\n        boundaries = skimage.segmentation.find_boundaries(image)\n        image = np.multiply(image, np.invert(boundaries))\n        image = ndi.binary_fill_holes(image > 0)\n        image = ndi.label(image)[0]\n        return image\n\n    def __wsh(self,\n            mask_img,\n            threshold,\n            border_img,\n            seeds,\n            threshold_adjustment=0.35,\n            small_object_size_cutoff=10,):\n\n        img_copy = np.copy(mask_img)\n        m = seeds * border_img  # * dt\n        img_copy[m <= threshold + threshold_adjustment] = 0\n        img_copy[m > threshold + threshold_adjustment] = 1\n        img_copy = img_copy.astype(np.bool)\n        img_copy = skimage.morphology.remove_small_objects(img_copy, small_object_size_cutoff).astype(\n            np.uint8\n        )\n\n        mask_img[mask_img <= threshold] = 0\n        mask_img[mask_img > threshold] = 1\n        mask_img = mask_img.astype(np.bool)\n        mask_img = skimage.morphology.remove_small_holes(mask_img, 85)\n        mask_img = skimage.morphology.remove_small_objects(mask_img, 1).astype(np.uint8)\n        markers = ndi.label(img_copy, output=np.uint32)[0]\n        labeled_array = skimage.segmentation.watershed(\n            mask_img, markers, mask=mask_img, watershed_line=True\n        )\n        return labeled_array\n\n    def label_cell(self, nuclei_pred, cell_pred):\n        nuclei_label = self.__wsh(\n            nuclei_pred[..., 2] / 255.0,\n            0.4,\n            1 - (nuclei_pred[..., 1] + cell_pred[..., 1]) / 255.0 > 0.05,\n            nuclei_pred[..., 2] / 255,\n            threshold_adjustment=-0.25,\n            small_object_size_cutoff=45,\n        )\n\n        # for hpa_image, to remove the small pseduo nuclei\n        nuclei_label = skimage.measure.label(nuclei_label)\n        # this is to remove the cell borders' signal from cell mask.\n        # could use np.logical_and with some revision, to replace this func.\n        # Tuned for segmentation hpa images\n        threshold_value = max(0.22, skimage.filters.threshold_otsu(cell_pred[..., 2] / 255) * 0.5)\n        # exclude the green area first\n        cell_region = np.multiply(\n            cell_pred[..., 2] / 255 > threshold_value,\n            np.invert(np.asarray(cell_pred[..., 1] / 255 > 0.05, dtype=np.int8)),\n        )\n        sk = np.asarray(cell_region, dtype=np.int8)\n        distance = np.clip(cell_pred[..., 2], 255 * threshold_value, cell_pred[..., 2])\n        cell_label = skimage.segmentation.watershed(-distance, nuclei_label, mask=sk)\n        selem = skimage.morphology.disk(6)\n        cell_label = skimage.morphology.closing(cell_label, selem)\n        cell_label = self.__fill_holes(cell_label)\n        # this part is to use green channel, and extend cell label to green channel\n        # benefit is to exclude cells clear on border but without nucleus\n        sk = np.asarray(\n            np.add(\n                np.asarray(cell_label > 0, dtype=np.int8),\n                np.asarray(cell_pred[..., 1] / 255 > 0.05, dtype=np.int8),\n            )\n            > 0,\n            dtype=np.int8,\n        )\n        cell_label = skimage.segmentation.watershed(-distance, cell_label, mask=sk)\n        cell_label = self.__fill_holes(cell_label)\n        cell_label = np.asarray(cell_label > 0, dtype=np.uint8)\n        cell_label = skimage.measure.label(cell_label)\n        cell_label = skimage.measure.label(cell_label)\n        cell_label = np.asarray(cell_label, dtype=np.uint16)\n        nuclei_label = np.multiply(cell_label > 0, nuclei_label) > 0\n        nuclei_label = skimage.measure.label(nuclei_label)\n        nuclei_label = np.multiply(cell_label, nuclei_label > 0)\n        return nuclei_label, cell_label\n\n    def __call__(self, images=None, image_root=None, image_ids=None, precombined=True):\n        if not precombined:\n            image_paths = self.build_image_names(image_root, image_ids)\n            # For nuclei\n            nuclei_segmentations = self.pred_nuclei(image_paths[2])\n            # For full cells\n            cell_segmentations = self.pred_cells(image_paths)\n        else:\n             # For nuclei\n            nuclei_segmentations = self.pred_nuclei(images[..., 2])\n            # For full cells\n            cell_segmentations = self.pred_cells(images[..., [0, 3, 2]], precombined=precombined)\n        nuclei_masks, cell_masks = [], []\n        for nuclei, cell in zip(nuclei_segmentations, cell_segmentations):\n            nuclei_mask, cell_mask = self.label_cell(nuclei, cell)\n            nuclei_masks.append(nuclei_mask)\n            cell_masks.append(cell_mask)\n        return nuclei_masks, cell_masks","metadata":{"papermill":{"duration":0.111628,"end_time":"2021-03-25T15:30:41.134068","exception":false,"start_time":"2021-03-25T15:30:41.02244","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\nsegmentator = CellSegment(device, cache_root=kernel_data_root)\n# For nuclei\nnuclei_mask,  cell_mask = segmentator(images=image.transpose(1, 2, 0)[None])\nnuclei_mask, cell_mask = nuclei_mask[0], cell_mask[0]\n# Plot prms grid with original image\n_, axs = plt.subplots(1, 3, figsize=(15, 5))\nvisualize(image, labels, axs[0], \"Original image\")\nvisualize(nuclei_mask, labels, axs[1], \"Nuclei Mask\")\nvisualize(cell_mask, labels, axs[2], \"Cell Mask\")","metadata":{"papermill":{"duration":0.088437,"end_time":"2021-03-25T15:30:41.291535","exception":false,"start_time":"2021-03-25T15:30:41.203098","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# predict instance masks via proposal retrieval\nproposals = [(nuclei_mask == i) for i in np.unique(nuclei_mask) if i != 0]\nretrieval_cfg=dict(proposals=proposals, param=(0.95, 1e-5, 0.8))\nmodels[fold].inference()\ninstance_list = models[fold](torch.tensor(image[:3]).unsqueeze(0).to(device), retrieval_cfg=retrieval_cfg)\n# visualization\nif instance_list is None or not len(instance_list):\n    print('No object detected')\nelse:\n    # peak response maps are merged if they select similar proposals\n    vis = prm.prm_visualize(instance_list, class_names=list(map(str, range(19))))\n    # peak response maps are merged if they select similar proposals\n    # Plot prms grid with original image\n    _, axs = plt.subplots(1, 3, figsize=(15, 5))\n    visualize(image, labels, axs[0], \"Original image\")\n    visualize(vis[0].transpose(2, 0, 1), labels, axs[1], \"Prediction\")\n    visualize(vis[1].transpose(2, 0, 1), labels, axs[2], \"Peak Response Maps\")","metadata":{"papermill":{"duration":0.154663,"end_time":"2021-03-25T15:30:42.214197","exception":false,"start_time":"2021-03-25T15:30:42.059534","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TestDataset:\n    def __init__(self, competition_root, kernel_data_root, resize=(256, 256)):\n        segmentor = CellSegment(device=device, cache_root=kernel_data_root)\n\n        self.data = []\n        colors = ['red', 'green', 'blue', 'yellow']\n        regx = \"(.+)(?=_(red|green|blue|yellow).png)\"\n        image_root = os.path.join(competition_root, \"test\")\n        image_ids = set(map(lambda name: re.match(regx, name)[0], os.listdir(image_root)))\n\n        for image_id in tqdm(image_ids, total=len(image_ids)):\n            if not os.path.exists('test'): os.makedirs('test')\n            image_path = os.path.join(competition_root, \"test\", f\"{image_id}_{{}}.png\")\n            image = [cv2.imread(image_path.format(c), cv2.IMREAD_GRAYSCALE) for c in colors]\n            image_height, image_width = image[0].shape\n            image = np.stack([cv2.resize(c, resize) for c in image], axis=-1)\n            np.save(os.path.join(\"test\", f\"{image_id}.npy\"), image)\n            nuclei_masks, cell_masks = segmentor(images=image[None], precombined=True)\n            np.save(os.path.join(\"test\", f\"{image_id}_nuclei_mask.npy\"), nuclei_masks[0])\n            np.save(os.path.join(\"test\", f\"{image_id}_cell_mask.npy\"), cell_masks[0])\n\n            self.data += [{\n                \"image_id\": image_id,\n                \"image_width\": image_width,\n                \"image_height\": image_height,\n                \"image_path\": os.path.join(\"test\", f\"{image_id}.npy\"),\n                \"nuclei_mask_path\": os.path.join(\"test\", f\"{image_id}_nuclei_mask.npy\"),\n                \"cell_mask_path\": os.path.join(\"test\", f\"{image_id}_cell_mask.npy\"),\n            }]\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        image_path = self.data[idx][\"image_path\"]\n        image = np.load(image_path, allow_pickle=True)#[..., :3]\n        image = image.transpose(2, 0, 1).astype(np.float32)\n        image = np.divide(image, 255.0)\n\n        nuclei_mask_path = self.data[idx][\"nuclei_mask_path\"]\n        nuclei_mask = np.load(nuclei_mask_path, allow_pickle=True)\n        nuclei_proposals = np.stack([(nuclei_mask == i) for i in np.unique(nuclei_mask) if i != 0])\n\n        cell_mask_path = self.data[idx][\"cell_mask_path\"]\n        cell_mask = np.load(cell_mask_path, allow_pickle=True)\n        cell_proposals = np.stack([(cell_mask == i) for i in np.unique(cell_mask) if i != 0])\n        return image, nuclei_proposals, cell_proposals","metadata":{"papermill":{"duration":0.100231,"end_time":"2021-03-25T15:30:43.074572","exception":false,"start_time":"2021-03-25T15:30:42.974341","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntest_dataset = TestDataset(competition_root, kernel_data_root)\ntest_loader = DataLoader(test_dataset, batch_size=1)","metadata":{"papermill":{"duration":0.109369,"end_time":"2021-03-25T15:30:43.270853","exception":false,"start_time":"2021-03-25T15:30:43.161484","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx = np.random.randint(len(test_dataset))\nfold = np.random.randint(num_folds)\nimage, nuclei_proposals, cell_proposals = test_dataset[idx]\nmodels[fold].eval()\nlogits = models[fold](torch.tensor(image[:3]).unsqueeze(0).to(device))\nlabels = torch.sigmoid(logits[0]).gt(0.5).nonzero(as_tuple=True)[0].cpu().numpy()\nvisualize(image, labels)","metadata":{"papermill":{"duration":0.101991,"end_time":"2021-03-25T15:30:43.453916","exception":false,"start_time":"2021-03-25T15:30:43.351925","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models[fold].inference()\nretrieval_cfg = dict(proposals=cell_proposals, param=(0.95, 1e-5, 0.8))\ninstance_list = models[fold](torch.tensor(image[:3]).unsqueeze(0).to(device), retrieval_cfg=retrieval_cfg)\n# visualization\nif instance_list is None or not len(instance_list):\n    print('No object detected')\nelse:\n    # peak response maps are merged if they select similar proposals\n    vis = prm.prm_visualize(instance_list, class_names=list(map(str, range(19))))\n    # peak response maps are merged if they select similar proposals\n    # Plot prms grid with original image\n    _, axs = plt.subplots(1, 3, figsize=(15, 5))\n    visualize(image, labels, axs[0], \"Original image\")\n    visualize(vis[0].transpose(2, 0, 1), labels, axs[1], \"Prediction\")\n    visualize(vis[1].transpose(2, 0, 1), labels, axs[2], \"Peak Response Maps\")","metadata":{"papermill":{"duration":0.10593,"end_time":"2021-03-25T15:30:43.640713","exception":false,"start_time":"2021-03-25T15:30:43.534783","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(fold):\n    masks, scores = [], []\n    models[fold].inference()\n    loader = tqdm(test_loader, desc = f\"Predicting with fold {fold+1}\")\n    for batch_idx, (image, nuclei_proposals, cell_proposals) in enumerate(loader, start=1):\n        image = image[:, :3].to(device)\n        retrieval_cfg = dict(proposals=nuclei_proposals[0].numpy(), param=(0.95, 1e-5, 0.8))\n        instance_list = models[fold](image, retrieval_cfg=retrieval_cfg)\n        score = np.zeros((cell_proposals[0].size(0), 19))\n        score[:, 18] = 1.0\n        for instance in instance_list:\n            intersections = np.logical_and(nuclei_proposals[0].numpy(), instance[\"mask\"]).sum(axis=(1, 2))\n            unions = np.logical_or(nuclei_proposals[0].numpy(), instance[\"mask\"]).sum(axis=(1, 2))\n            ious = intersections / (unions + 1e-10)\n            nuclei_proposal = nuclei_proposals[0].numpy()[np.argmax(ious)]\n            intersections = np.logical_and(cell_proposals[0].numpy(), nuclei_proposal).sum(axis=(1, 2))\n            score[np.argmax(intersections), [instance[\"category\"], 18]] = [instance[\"score\"], 0.0]\n        masks.append(cell_proposals[0].numpy()); scores.append(score)            \n        del image; del nuclei_proposals; del cell_proposals\n        del retrieval_cfg; del instance_list; del score\n        torch.cuda.empty_cache()\n    loader.write(\"Done!\")\n    return masks, scores","metadata":{"papermill":{"duration":0.100785,"end_time":"2021-03-25T15:30:44.353186","exception":false,"start_time":"2021-03-25T15:30:44.252401","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"masks, scores = predict(0)","metadata":{"papermill":{"duration":0.11669,"end_time":"2021-03-25T15:30:44.556423","exception":false,"start_time":"2021-03-25T15:30:44.439733","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx = np.random.randint(len(test_dataset))*0\nimage, _, _ = test_dataset[idx]\nmask, score = masks[idx], scores[idx]\nlabels = np.argmax(score, axis=1)\nmask = np.sum(mask * labels[:, None, None], axis=0)0+\nlabels = score.sum(axis=0)\n_, axs = plt.subplots(1, 2, figsize=(15, 5))\nvisualize(image, labels, axs[0], \"Test image\")\nvisualize(mask, labels, axs[1], \"Predicted Mask\")","metadata":{"papermill":{"duration":0.108261,"end_time":"2021-03-25T15:30:44.750776","exception":false,"start_time":"2021-03-25T15:30:44.642515","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def encode_binary_mask(mask):\n    \"\"\"Converts a binary mask into OID challenge encoding ascii text.\"\"\"\n\n    # check input mask --\n    if mask.dtype != np.bool:\n        raise ValueError(\"encode_binary_mask expects a binary mask, received dtype == %s\" % mask.dtype)\n\n    mask = np.squeeze(mask)\n    if len(mask.shape) != 2:\n        raise ValueError(\"encode_binary_mask expects a 2d mask, received shape == %s\" % mask.shape)\n\n    # convert input mask to expected COCO API input --\n    mask_to_encode = mask.reshape(mask.shape[0], mask.shape[1], 1)\n    mask_to_encode = mask_to_encode.astype(np.uint8)\n    mask_to_encode = np.asfortranarray(mask_to_encode)\n\n    # RLE encode mask --\n    encoded_mask = coco_mask.encode(mask_to_encode)[0][\"counts\"]\n\n    # compress and base64 encoding --\n    binary_str = zlib.compress(encoded_mask, zlib.Z_BEST_COMPRESSION)\n    base64_str = base64.b64encode(binary_str)\n    return base64_str.decode()","metadata":{"papermill":{"duration":0.096215,"end_time":"2021-03-25T15:30:44.935843","exception":false,"start_time":"2021-03-25T15:30:44.839628","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = {\"ID\": [], \"ImageWidth\": [], \"ImageHeight\": [], \"PredictionString\": []}\nfor image_masks, image_scores, data in zip(masks, scores, test_dataset.data):\n    dsize = (data[\"image_height\"], data[\"image_width\"])\n    for mask, score in zip(image_masks, image_scores):\n        mask = cv2.resize(src=mask.astype(int), dsize=dsize, interpolation=cv2.INTER_NEAREST)\n        label = np.argmax(score)\n        enc_str = encode_binary_mask(mask.astype(bool))\n        submission[\"ID\"].append(data[\"image_id\"])\n        submission[\"ImageWidth\"].append(data[\"image_width\"])\n        submission[\"ImageHeight\"].append(data[\"image_height\"])\n        submission[\"PredictionString\"].append(\"%.0f %.4f %s\"%(label, score[label], enc_str))\nsubmission = pd.DataFrame(submission)\nsubmission = submission.groupby([\"ID\", \"ImageWidth\", \"ImageHeight\"])[\"PredictionString\"].apply(\" \".join).reset_index()","metadata":{"papermill":{"duration":0.111049,"end_time":"2021-03-25T15:30:45.133399","exception":false,"start_time":"2021-03-25T15:30:45.02235","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv(\"submission.csv\", index=False)\nsubmission.head()","metadata":{"papermill":{"duration":0.106738,"end_time":"2021-03-25T15:30:45.327821","exception":false,"start_time":"2021-03-25T15:30:45.221083","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for item in os.listdir():\n    if os.path.isdir(item):\n        shutil.rmtree(item, ignore_errors=False)","metadata":{"papermill":{"duration":0.08743,"end_time":"2021-03-25T15:30:45.502316","exception":false,"start_time":"2021-03-25T15:30:45.414886","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}