{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### How to set thresholds for each organ\ncustomize mmseg/models/segmentors/encorder_decoder.py to get sigmoid raw outputs.  \nIf you want to use sigmoid raw outputs of each class, add sigmoid=True in test_cfg.   \nArgmax outputs as usual if simoid=False (default).  \n  \nmodel = dict(\n    ...,\n    train_cfg=dict(),\n    test_cfg=dict(mode='whole', sigmoid=True)\n    )","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"%%writefile /path/to/mmsegmentation/mmseg/models/segmentors/encoder_decoder.py\n\n# Copyright (c) OpenMMLab. All rights reserved.\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom mmseg.core import add_prefix\nfrom mmseg.ops import resize\nfrom .. import builder\nfrom ..builder import SEGMENTORS\nfrom .base import BaseSegmentor\n\n\n@SEGMENTORS.register_module()\nclass EncoderDecoder(BaseSegmentor):\n    \"\"\"Encoder Decoder segmentors.\n\n    EncoderDecoder typically consists of backbone, decode_head, auxiliary_head.\n    Note that auxiliary_head is only used for deep supervision during training,\n    which could be dumped during inference.\n    \"\"\"\n\n    def __init__(self,\n                 backbone,\n                 decode_head,\n                 neck=None,\n                 auxiliary_head=None,\n                 train_cfg=None,\n                 test_cfg=None,\n                 pretrained=None,\n                 init_cfg=None):\n        super(EncoderDecoder, self).__init__(init_cfg)\n        if pretrained is not None:\n            assert backbone.get('pretrained') is None, \\\n                'both backbone and segmentor set pretrained weight'\n            backbone.pretrained = pretrained\n        self.backbone = builder.build_backbone(backbone)\n        if neck is not None:\n            self.neck = builder.build_neck(neck)\n        self._init_decode_head(decode_head)\n        self._init_auxiliary_head(auxiliary_head)\n\n        self.train_cfg = train_cfg\n        self.test_cfg = test_cfg\n\n        assert self.with_decode_head\n\n    def _init_decode_head(self, decode_head):\n        \"\"\"Initialize ``decode_head``\"\"\"\n        self.decode_head = builder.build_head(decode_head)\n        self.align_corners = self.decode_head.align_corners\n        self.num_classes = self.decode_head.num_classes\n\n    def _init_auxiliary_head(self, auxiliary_head):\n        \"\"\"Initialize ``auxiliary_head``\"\"\"\n        if auxiliary_head is not None:\n            if isinstance(auxiliary_head, list):\n                self.auxiliary_head = nn.ModuleList()\n                for head_cfg in auxiliary_head:\n                    self.auxiliary_head.append(builder.build_head(head_cfg))\n            else:\n                self.auxiliary_head = builder.build_head(auxiliary_head)\n\n    def extract_feat(self, img):\n        \"\"\"Extract features from images.\"\"\"\n        x = self.backbone(img)\n        if self.with_neck:\n            x = self.neck(x)\n        return x\n\n    def encode_decode(self, img, img_metas):\n        \"\"\"Encode images with backbone and decode into a semantic segmentation\n        map of the same size as input.\"\"\"\n        x = self.extract_feat(img)\n        out = self._decode_head_forward_test(x, img_metas)\n        out = resize(\n            input=out,\n            size=img.shape[2:],\n            mode='bilinear',\n            align_corners=self.align_corners)\n        return out\n\n    def _decode_head_forward_train(self, x, img_metas, gt_semantic_seg):\n        \"\"\"Run forward function and calculate loss for decode head in\n        training.\"\"\"\n        losses = dict()\n        loss_decode = self.decode_head.forward_train(x, img_metas,\n                                                     gt_semantic_seg,\n                                                     self.train_cfg)\n\n        losses.update(add_prefix(loss_decode, 'decode'))\n        return losses\n\n    def _decode_head_forward_test(self, x, img_metas):\n        \"\"\"Run forward function and calculate loss for decode head in\n        inference.\"\"\"\n        seg_logits = self.decode_head.forward_test(x, img_metas, self.test_cfg)\n        return seg_logits\n\n    def _auxiliary_head_forward_train(self, x, img_metas, gt_semantic_seg):\n        \"\"\"Run forward function and calculate loss for auxiliary head in\n        training.\"\"\"\n        losses = dict()\n        if isinstance(self.auxiliary_head, nn.ModuleList):\n            for idx, aux_head in enumerate(self.auxiliary_head):\n                loss_aux = aux_head.forward_train(x, img_metas,\n                                                  gt_semantic_seg,\n                                                  self.train_cfg)\n                losses.update(add_prefix(loss_aux, f'aux_{idx}'))\n        else:\n            loss_aux = self.auxiliary_head.forward_train(\n                x, img_metas, gt_semantic_seg, self.train_cfg)\n            losses.update(add_prefix(loss_aux, 'aux'))\n\n        return losses\n\n    def forward_dummy(self, img):\n        \"\"\"Dummy forward function.\"\"\"\n        seg_logit = self.encode_decode(img, None)\n\n        return seg_logit\n\n    def forward_train(self, img, img_metas, gt_semantic_seg):\n        \"\"\"Forward function for training.\n\n        Args:\n            img (Tensor): Input images.\n            img_metas (list[dict]): List of image info dict where each dict\n                has: 'img_shape', 'scale_factor', 'flip', and may also contain\n                'filename', 'ori_shape', 'pad_shape', and 'img_norm_cfg'.\n                For details on the values of these keys see\n                `mmseg/datasets/pipelines/formatting.py:Collect`.\n            gt_semantic_seg (Tensor): Semantic segmentation masks\n                used if the architecture supports semantic segmentation task.\n\n        Returns:\n            dict[str, Tensor]: a dictionary of loss components\n        \"\"\"\n\n        x = self.extract_feat(img)\n\n        losses = dict()\n\n        loss_decode = self._decode_head_forward_train(x, img_metas,\n                                                      gt_semantic_seg)\n        losses.update(loss_decode)\n\n        if self.with_auxiliary_head:\n            loss_aux = self._auxiliary_head_forward_train(\n                x, img_metas, gt_semantic_seg)\n            losses.update(loss_aux)\n\n        return losses\n\n    # TODO refactor\n    def slide_inference(self, img, img_meta, rescale):\n        \"\"\"Inference by sliding-window with overlap.\n\n        If h_crop > h_img or w_crop > w_img, the small patch will be used to\n        decode without padding.\n        \"\"\"\n\n        h_stride, w_stride = self.test_cfg.stride\n        h_crop, w_crop = self.test_cfg.crop_size\n        batch_size, _, h_img, w_img = img.size()\n        num_classes = self.num_classes\n        h_grids = max(h_img - h_crop + h_stride - 1, 0) // h_stride + 1\n        w_grids = max(w_img - w_crop + w_stride - 1, 0) // w_stride + 1\n        preds = img.new_zeros((batch_size, num_classes, h_img, w_img))\n        count_mat = img.new_zeros((batch_size, 1, h_img, w_img))\n        for h_idx in range(h_grids):\n            for w_idx in range(w_grids):\n                y1 = h_idx * h_stride\n                x1 = w_idx * w_stride\n                y2 = min(y1 + h_crop, h_img)\n                x2 = min(x1 + w_crop, w_img)\n                y1 = max(y2 - h_crop, 0)\n                x1 = max(x2 - w_crop, 0)\n                crop_img = img[:, :, y1:y2, x1:x2]\n                crop_seg_logit = self.encode_decode(crop_img, img_meta)\n                preds += F.pad(crop_seg_logit,\n                               (int(x1), int(preds.shape[3] - x2), int(y1),\n                                int(preds.shape[2] - y2)))\n\n                count_mat[:, :, y1:y2, x1:x2] += 1\n        assert (count_mat == 0).sum() == 0\n        if torch.onnx.is_in_onnx_export():\n            # cast count_mat to constant while exporting to ONNX\n            count_mat = torch.from_numpy(\n                count_mat.cpu().detach().numpy()).to(device=img.device)\n        preds = preds / count_mat\n        if rescale:\n            # remove padding area\n            resize_shape = img_meta[0]['img_shape'][:2]\n            preds = preds[:, :, :resize_shape[0], :resize_shape[1]]\n            preds = resize(\n                preds,\n                size=img_meta[0]['ori_shape'][:2],\n                mode='bilinear',\n                align_corners=self.align_corners,\n                warning=False)\n        return preds\n\n    def whole_inference(self, img, img_meta, rescale):\n        \"\"\"Inference with full image.\"\"\"\n\n        seg_logit = self.encode_decode(img, img_meta)\n        if rescale:\n            # support dynamic shape for onnx\n            if torch.onnx.is_in_onnx_export():\n                size = img.shape[2:]\n            else:\n                # remove padding area\n                resize_shape = img_meta[0]['img_shape'][:2]\n                seg_logit = seg_logit[:, :, :resize_shape[0], :resize_shape[1]]\n                size = img_meta[0]['ori_shape'][:2]\n            seg_logit = resize(\n                seg_logit,\n                size=size,\n                mode='bilinear',\n                align_corners=self.align_corners,\n                warning=False)\n\n        return seg_logit\n\n    def inference(self, img, img_meta, rescale):\n        \"\"\"Inference with slide/whole style.\n\n        Args:\n            img (Tensor): The input image of shape (N, 3, H, W).\n            img_meta (dict): Image info dict where each dict has: 'img_shape',\n                'scale_factor', 'flip', and may also contain\n                'filename', 'ori_shape', 'pad_shape', and 'img_norm_cfg'.\n                For details on the values of these keys see\n                `mmseg/datasets/pipelines/formatting.py:Collect`.\n            rescale (bool): Whether rescale back to original shape.\n\n        Returns:\n            Tensor: The output segmentation map.\n        \"\"\"\n\n        assert self.test_cfg.mode in ['slide', 'whole']\n        ori_shape = img_meta[0]['ori_shape']\n        assert all(_['ori_shape'] == ori_shape for _ in img_meta)\n        if self.test_cfg.mode == 'slide':\n            seg_logit = self.slide_inference(img, img_meta, rescale)\n        else:\n            seg_logit = self.whole_inference(img, img_meta, rescale)\n\n        if self.test_cfg.get('sigmoid', False):\n            output = F.sigmoid(seg_logit)\n        else:\n            output = F.softmax(seg_logit, dim=1)\n        flip = img_meta[0]['flip']\n        if flip:\n            flip_direction = img_meta[0]['flip_direction']\n            assert flip_direction in ['horizontal', 'vertical']\n            if flip_direction == 'horizontal':\n                output = output.flip(dims=(3, ))\n            elif flip_direction == 'vertical':\n                output = output.flip(dims=(2, ))\n\n        return output\n\n    def simple_test(self, img, img_meta, rescale=True):\n        \"\"\"Simple test with single image.\"\"\"\n        seg_logit = self.inference(img, img_meta, rescale)\n        if self.test_cfg.get('sigmoid', False):\n            seg_pred = seg_logit\n        else:\n            seg_pred = seg_logit.argmax(dim=1)\n        if torch.onnx.is_in_onnx_export():\n            # our inference backend only support 4D output\n            seg_pred = seg_pred.unsqueeze(0)\n            return seg_pred\n        seg_pred = seg_pred.cpu().numpy()\n        # unravel batch dim\n        seg_pred = list(seg_pred)\n        return seg_pred\n\n    def aug_test(self, imgs, img_metas, rescale=True):\n        \"\"\"Test with augmentations.\n\n        Only rescale=True is supported.\n        \"\"\"\n        # aug_test rescale all imgs back to ori_shape for now\n        assert rescale\n        # to save memory, we get augmented seg logit inplace\n        seg_logit = self.inference(imgs[0], img_metas[0], rescale)\n        for i in range(1, len(imgs)):\n            cur_seg_logit = self.inference(imgs[i], img_metas[i], rescale)\n            seg_logit += cur_seg_logit\n        seg_logit /= len(imgs)\n        if self.test_cfg.get('sigmoid', False):\n            seg_pred = seg_logit\n        else:\n            seg_pred = seg_logit.argmax(dim=1)\n        seg_pred = seg_pred.cpu().numpy()\n        # unravel batch dim\n        seg_pred = list(seg_pred)\n        return seg_pred\n","metadata":{},"execution_count":null,"outputs":[]}]}