{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":7309407,"sourceType":"datasetVersion","datasetId":4171518},{"sourceId":150248402,"sourceType":"kernelVersion"}],"dockerImageVersionId":30588,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys, os\nimport cv2\nimport pandas as pd\nfrom glob import glob\nimport numpy as np\n\nfrom timeit import default_timer as timer\n!python -m pip install --no-index --find-links=/kaggle/input/pip-download-for-segmentation-models-pytorch segmentation-models-pytorch\nimport segmentation_models_pytorch as smp\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast\nimport matplotlib\nimport matplotlib.pyplot as plt\n\nprint('IMPORT OK  !!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-29T17:31:39.299629Z","iopub.execute_input":"2023-12-29T17:31:39.299915Z","iopub.status.idle":"2023-12-29T17:32:03.675789Z","shell.execute_reply.started":"2023-12-29T17:31:39.299889Z","shell.execute_reply":"2023-12-29T17:32:03.674749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Optional, Union, List\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport segmentation_models_pytorch as smp\nfrom segmentation_models_pytorch.base import (\n    SegmentationModel,\n    SegmentationHead,\n    ClassificationHead,\n)\nfrom segmentation_models_pytorch.decoders.unet.decoder import UnetDecoder\n\nclass SenUNet(nn.Module):\n    def __init__(self, encoder_name=\"resnest26d\",output_stride=32,\n                 encoder_depth=5 , n_time=8, pickup_index=4,  \n        decoder_use_batchnorm: bool = True,\n        decoder_channels: List[int] = (256, 128, 64, 32, 16),\n        decoder_attention_type: Optional[str] = None, classes=1, activation=None):\n        super(SenUNet, self).__init__()\n        kwargs = dict(\n            in_chans=1,\n            features_only=True,\n#             output_stride=output_stride,\n            pretrained=False,\n            out_indices=tuple(range(encoder_depth)),\n        )\n        self.encoder = timm.create_model(encoder_name, **kwargs)\n        self._out_channels = [\n            3,\n        ] + self.encoder.feature_info.channels()\n\n        self.decoder = UnetDecoder(\n            encoder_channels=self._out_channels,\n            decoder_channels=decoder_channels,\n            n_blocks=encoder_depth,\n            use_batchnorm=decoder_use_batchnorm,\n            center=True if encoder_name.startswith(\"vgg\") else False,\n            attention_type=decoder_attention_type,\n        )\n\n        self.segmentation_head = SegmentationHead(\n            in_channels=decoder_channels[-1],\n            out_channels=classes,\n            activation=activation,\n            kernel_size=3,\n        )\n\n\n        # not all models support output stride argument, drop it by default\n#         if output_stride == 32:\n#             kwargs.pop(\"output_stride\")\n\n        #self.conv4_3d_1 = Residual3DBlock(512)\n        self.n_time = n_time\n        self.pickup_index = pickup_index\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        h = (H//32)*32\n        w = (W//32)*32\n        x = x[:,:,:h,:w]\n        features = self.encoder(x)        \n        features = [\n            x,\n        ] + features\n\n        decoder_output = self.decoder(*features)\n\n        masks = self.segmentation_head(decoder_output)\n        masks = F.pad(masks,[0,W-w,0,H-h,0,0,0,0], mode='constant', value=0)\n        \n        return masks","metadata":{"execution":{"iopub.status.busy":"2023-12-29T17:32:03.677683Z","iopub.execute_input":"2023-12-29T17:32:03.678002Z","iopub.status.idle":"2023-12-29T17:32:03.690989Z","shell.execute_reply.started":"2023-12-29T17:32:03.677977Z","shell.execute_reply":"2023-12-29T17:32:03.690035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Conv2dReLU(nn.Sequential):\n    def __init__(\n        self,\n        in_channels,\n        out_channels,\n        kernel_size,\n        padding=0,\n        stride=1,\n        use_layernorm=True,\n    ):\n\n        if use_layernorm == \"inplace\" and InPlaceABN is None:\n            raise RuntimeError(\n                \"In order to use `use_batchnorm='inplace'` inplace_abn package must be installed. \"\n                + \"To install see: https://github.com/mapillary/inplace_abn\"\n            )\n\n        conv = nn.Conv2d(\n            in_channels,\n            out_channels,\n            kernel_size,\n            stride=stride,\n            padding=padding,\n            bias=not (use_layernorm),\n        )\n        relu = nn.GELU()\n\n        if use_layernorm == \"inplace\":\n            bn = InPlaceABN(out_channels, activation=\"leaky_relu\", activation_param=0.0)\n            relu = nn.Identity()\n\n        elif use_layernorm and use_layernorm != \"inplace\":\n            bn = LayerNorm2d(out_channels)\n\n        else:\n            bn = nn.Identity()\n\n        super(Conv2dReLU, self).__init__(conv, bn, relu)\n\nclass SenUNetStem(nn.Module):\n    def __init__(self, encoder_name=\"resnest26d\",output_stride=32,\n                 encoder_depth=5 , n_time=8, pickup_index=4,  \n        decoder_use_batchnorm: bool = True,\n        decoder_channels: List[int] = (256, 128, 64, 32, 16),\n        decoder_attention_type: Optional[str] = None, classes=1, activation=None):\n        super(SenUNetStem, self).__init__()\n        kwargs = dict(\n            in_chans=1,\n            features_only=True,\n            # output_stride=output_stride,\n            pretrained=False,\n            out_indices=tuple(range(encoder_depth)),\n        )\n        self.conv_stem = Conv2dReLU(1, 16, 3, use_layernorm=False)\n        self.encoder = timm.create_model(encoder_name, **kwargs)\n        self._out_channels = [\n            32,\n        ] + self.encoder.feature_info.channels()\n\n        self.decoder = UnetDecoder(\n            encoder_channels=self._out_channels,\n            decoder_channels=decoder_channels,\n            n_blocks=encoder_depth,\n            use_batchnorm=decoder_use_batchnorm,\n            center=True if encoder_name.startswith(\"vgg\") else False,\n            attention_type=decoder_attention_type,\n        )\n\n        self.segmentation_head = SegmentationHead(\n            in_channels=decoder_channels[-1],\n            out_channels=classes,\n            activation=activation,\n            kernel_size=3,\n        )\n\n\n        # not all models support output stride argument, drop it by default\n        # if output_stride == 32:\n        #     kwargs.pop(\"output_stride\")\n\n        #self.conv4_3d_1 = Residual3DBlock(512)\n        self.n_time = n_time\n        self.pickup_index = pickup_index\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        h = (H//32)*32\n        w = (W//32)*32\n        x = x[:,:,:h,:w]\n        stem = self.conv_stem(x)\n        features = self.encoder(x)        \n        features = [\n            stem,\n        ] + features\n\n        decoder_output = self.decoder(*features)\n\n        masks = self.segmentation_head(decoder_output)\n        masks = F.pad(masks,[0,W-w,0,H-h,0,0,0,0], mode='constant', value=0)\n        \n        return masks","metadata":{"execution":{"iopub.status.busy":"2023-12-29T17:35:41.455395Z","iopub.execute_input":"2023-12-29T17:35:41.455896Z","iopub.status.idle":"2023-12-29T17:35:41.473909Z","shell.execute_reply.started":"2023-12-29T17:35:41.455859Z","shell.execute_reply":"2023-12-29T17:35:41.472937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#https://www.kaggle.com/competitions/blood-vessel-segmentation/discussion/456033\ndef remove_small_objects(mask, min_size):\n    # Find all connected components (labels)\n    num_label, label, stats, centroid = cv2.connectedComponentsWithStats(mask, connectivity=8)\n\n    # create a mask where small objects are removed\n    processed = np.zeros_like(mask)\n    for l in range(1, num_label):\n        if stats[l, cv2.CC_STAT_AREA] >= min_size:\n            processed[label == l] = 255\n\n    return processed\n\ndef rle_encode(mask):\n    pixel = mask.flatten()\n    pixel = np.concatenate([[0], pixel, [0]])\n    run = np.where(pixel[1:] != pixel[:-1])[0] + 1\n    run[1::2] -= run[::2]\n    rle = ' '.join(str(r) for r in run)\n    if rle == '':\n        rle = '1 0'\n    return rle\n\n#-------------------------------\n\ncheckpoint_file = \"\"\n\nnet = SenUNetStem(encoder_name=\"maxvit_tiny_tf_512.in1k\")\n#run_check_net()\nstate_dict = torch.load(\"/kaggle/input/sennet-models/baseline58_ema.pth\", map_location=lambda storage, loc: storage)\nprint(net.load_state_dict(state_dict, strict=False))  # True\n\nnet = net.eval()\nnet = net.cuda()\n#net = torch.compile(net)\n\npredict_mode = \"tile\"","metadata":{"execution":{"iopub.status.busy":"2023-12-29T17:35:44.245309Z","iopub.execute_input":"2023-12-29T17:35:44.246072Z","iopub.status.idle":"2023-12-29T17:35:48.770743Z","shell.execute_reply.started":"2023-12-29T17:35:44.246038Z","shell.execute_reply":"2023-12-29T17:35:48.769948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2","metadata":{"execution":{"iopub.status.busy":"2023-12-25T10:16:27.225004Z","iopub.execute_input":"2023-12-25T10:16:27.225291Z","iopub.status.idle":"2023-12-25T10:16:28.612215Z","shell.execute_reply.started":"2023-12-25T10:16:27.225267Z","shell.execute_reply":"2023-12-25T10:16:28.611435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_aug = A.Compose([\n        ToTensorV2(transpose_mask=True)\n    ])","metadata":{"execution":{"iopub.status.busy":"2023-12-25T10:16:28.613263Z","iopub.execute_input":"2023-12-25T10:16:28.613637Z","iopub.status.idle":"2023-12-25T10:16:28.618381Z","shell.execute_reply.started":"2023-12-25T10:16:28.613613Z","shell.execute_reply":"2023-12-25T10:16:28.617273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"paths = glob(\"/kaggle/input/blood-vessel-segmentation/test/*\")","metadata":{"execution":{"iopub.status.busy":"2023-12-25T10:16:28.620535Z","iopub.execute_input":"2023-12-25T10:16:28.620935Z","iopub.status.idle":"2023-12-25T10:16:28.633862Z","shell.execute_reply.started":"2023-12-25T10:16:28.620905Z","shell.execute_reply":"2023-12-25T10:16:28.63299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_img(path):\n    img = cv2.imread(path, cv2.IMREAD_GRAYSCALE).astype(np.float32)\n    return img\n\ndef normalize(img):\n    img = (img - img.min())/(img.max() - img.min() +0.001)\n    return img\n\ndef norm_by_percentile(volume, low=10, high=99.8, alpha=0.01):\n    xmin = np.percentile(volume,low)\n    xmax = np.percentile(volume,high)\n    x = (volume-xmin)/(xmax-xmin)\n    if 1:\n        x[x>1]=(x[x>1]-1)*alpha +1\n        x[x<0]=(x[x<0])*alpha\n    #x = np.clip(x,0,1)\n    return x","metadata":{"execution":{"iopub.status.busy":"2023-12-25T10:16:28.634813Z","iopub.execute_input":"2023-12-25T10:16:28.63513Z","iopub.status.idle":"2023-12-25T10:16:28.642636Z","shell.execute_reply.started":"2023-12-25T10:16:28.635098Z","shell.execute_reply":"2023-12-25T10:16:28.641719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_full_image(net, img):\n    mask = F.sigmoid(net(img)).float().data.cpu()\n                    \n    mask += torch.flip(F.sigmoid(net(torch.flip(img, dims=[2,]))), \n                                      dims=[2,]).float().data.cpu()\n    mask += torch.flip(F.sigmoid(net(torch.flip(img, dims=[3,]))), \n                                  dims=[3,]).float().data.cpu()\n    mask = mask / 3.0\n    return mask[0 , 0]\n\ndef predict_tile_image(net, img, sliding_window, tile_img_size):\n    _, _, H, W = img.size()\n    full_mask = torch.zeros(H, W)\n    count_mask = torch.zeros(H, W)\n    \n    for start_h_idx in range(0, H, sliding_window):\n        for start_w_idx in range(0, W, sliding_window):\n            tile_img = img[:, :, start_h_idx: start_h_idx + tile_img_size, start_w_idx: start_w_idx + tile_img_size]\n            _, _, part_H, part_W = tile_img.size()\n            pad_tile_img = F.pad(tile_img, (0, tile_img_size-part_W, 0, tile_img_size-part_H), \"constant\", 0)\n            mask = F.sigmoid(net(pad_tile_img)).float().data.cpu()[:, :, :part_H, :part_W]     \n            mask += torch.flip(F.sigmoid(net(torch.flip(pad_tile_img, dims=[2,]))), \n                                      dims=[2,]).float().data.cpu()[:, :, :part_H, :part_W]\n            mask += torch.flip(F.sigmoid(net(torch.flip(pad_tile_img, dims=[3,]))), \n                                  dims=[3,]).float().data.cpu()[:, :, :part_H, :part_W]\n            mask /= 3.0\n            full_mask[start_h_idx: start_h_idx + tile_img_size,\n                      start_w_idx: start_w_idx + tile_img_size] += mask[0, 0]\n            count_mask[start_h_idx: start_h_idx + tile_img_size, \n                       start_w_idx: start_w_idx + tile_img_size] += 1.0\n    full_mask = full_mask / count_mask\n    return full_mask","metadata":{"execution":{"iopub.status.busy":"2023-12-25T10:16:28.643773Z","iopub.execute_input":"2023-12-25T10:16:28.644075Z","iopub.status.idle":"2023-12-25T10:16:28.657362Z","shell.execute_reply.started":"2023-12-25T10:16:28.644051Z","shell.execute_reply":"2023-12-25T10:16:28.65659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_data = []\n\nfor image_folder in paths:\n    image_paths = sorted(glob(f'{image_folder}/images/*.tif'))\n    transformed_images = []\n    image_ids = []\n    for image_path in image_paths:\n        image_id = image_path.split(\"/\")[-3] + \"_\" + image_path.split(\"/\")[-1].replace(\".tif\", \"\")\n        img = load_img(image_path)\n        transformed = test_aug(image=img)\n        transformed_image = transformed['image']\n        transformed_images.append(transformed_image)\n        image_ids.append(image_id)\n    volume = torch.cat(transformed_images)\n    prob_volume = torch.zeros_like(volume)\n    volume = volume.unsqueeze(dim=0)\n    counter = 0\n    for axis_volume in [\"xy\", \"yz\", \"xz\"]:\n        with torch.no_grad():\n            if axis_volume == \"xy\":\n                for z in range(volume.size()[1]):\n                    normalize_img = norm_by_percentile(volume[:, z, :, :]).unsqueeze(dim=0).cuda()\n                    if predict_mode == \"full\":\n                        prob_volume[z, :, :] += predict_full_image(net, normalize_img)\n                    else:\n                        prob_volume[z, :, :] += predict_tile_image(net, normalize_img, 256, 512)\n                counter += 1\n            elif axis_volume == \"yz\" and len(image_paths) != 3:\n                for x in range(volume.size()[3]):\n                    normalize_img = norm_by_percentile(volume[:, :, :, x]).unsqueeze(dim=0).cuda()\n                    if predict_mode == \"full\":\n                        prob_volume[:, :, x] += predict_full_image(net, normalize_img)\n                    else:\n                        prob_volume[:, :, x] += predict_tile_image(net, normalize_img, 256, 512)\n                counter += 1\n            elif axis_volume == \"xz\" and len(image_paths) != 3:\n                for y in range(volume.size()[2]):\n                    normalize_img = norm_by_percentile(volume[:, :, y, :]).unsqueeze(dim=0).cuda()  \n                    if predict_mode == \"full\":\n                        prob_volume[:, y, :] += predict_full_image(net, normalize_img)\n                    else:\n                        prob_volume[:, y, :] += predict_tile_image(net, normalize_img, 256, 512)\n                counter += 1\n    prob_volume /= counter\n    prob_volume = prob_volume.numpy()\n    for i, image_id in enumerate(image_ids):\n        p = ((prob_volume[i, :, :]>0.3)*255).astype(np.uint8)\n        \n\n        #---post processing ---\n        #remove small\n        #https://www.kaggle.com/competitions/blood-vessel-segmentation/discussion/456033\n        p = remove_small_objects(p, min_size=1)\n\n        #----------------------\n        rle = rle_encode(p)\n            \n        df_data.append({\n            'id': image_id,\n            'rle': rle, #'1 0', \n        })\n","metadata":{"execution":{"iopub.status.busy":"2023-12-25T10:16:28.658645Z","iopub.execute_input":"2023-12-25T10:16:28.658972Z","iopub.status.idle":"2023-12-25T10:16:41.575116Z","shell.execute_reply.started":"2023-12-25T10:16:28.658947Z","shell.execute_reply":"2023-12-25T10:16:41.574351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission = pd.DataFrame(df_data)\ndf_submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-12-25T10:16:41.576139Z","iopub.execute_input":"2023-12-25T10:16:41.576397Z","iopub.status.idle":"2023-12-25T10:16:41.588608Z","shell.execute_reply.started":"2023-12-25T10:16:41.576375Z","shell.execute_reply":"2023-12-25T10:16:41.587683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}