{"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":"import sys\nsys.path.append('/kaggle/input/pretrainedmodels/pretrainedmodels-0.7.4')\nsys.path.append('/kaggle/input/efficientnet-pytorch/EfficientNet-PyTorch-master')\nsys.path.append('/kaggle/input/pytorch-image-models/pytorch-image-models')\nsys.path.append('/kaggle/input/segmentation-models-pytorch/segmentation_models_pytorch')\nsys.path.append('/kaggle/input/addict')\nimport numpy as np\nimport cv2\nimport pandas as pd\nfrom glob import glob\nimport torch.nn as nn\nimport json\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchmetrics import AveragePrecision\nimport torch\nimport torchvision\nfrom torchvision.transforms import transforms\nimport albumentations as A \nfrom albumentations.pytorch.transforms import ToTensorV2\nimport os\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm\nimport torch.nn.functional as F\nimport time\nimport base64\nimport typing as t\nimport zlib\nimport segmentation_models_pytorch as smp\n\nclass CFG:\n    data_path = '/kaggle/input/hubmap-hacking-the-human-vasculature/'\n    batch_size = 1\n    device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\n    th = 0.55\n    chepoint_dir = '/kaggle/input/hubmap-checpoint/'\n    model_types = ['UnetPlusPlus']\n    encoder_name_list = ['se_resnext50_32x4d']\n    is_tta = False\n    size = 512\n    org_size = 512\n    encoder_depth = 4\n    decoder_channels = [512, 256, 128, 64]\n    \n    test_aug = [\n        A.Resize(size, size),\n        A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n        ToTensorV2()\n    ]","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-24T11:27:31.822331Z","iopub.execute_input":"2023-06-24T11:27:31.822922Z","iopub.status.idle":"2023-06-24T11:27:50.330680Z","shell.execute_reply.started":"2023-06-24T11:27:31.822890Z","shell.execute_reply":"2023-06-24T11:27:50.329604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\n!mkdir /kaggle/working/packages\n!cp -r /kaggle/input/pycocotools/* /kaggle/working/packages\nos.chdir(\"/kaggle/working/packages/pycocotools-2.0.6/\")\n!python setup.py install\n!pip install . --no-index --find-links /kaggle/working/packages/\nos.chdir(\"/kaggle/working\")\nfrom pycocotools import _mask as coco_mask","metadata":{"execution":{"iopub.status.busy":"2023-06-24T11:27:50.334922Z","iopub.execute_input":"2023-06-24T11:27:50.335493Z","iopub.status.idle":"2023-06-24T11:28:40.050224Z","shell.execute_reply.started":"2023-06-24T11:27:50.335466Z","shell.execute_reply":"2023-06-24T11:28:40.048956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HUBMAPDataset(Dataset):\n    def __init__(self, image_dir, transform=None):\n        self.image_dir = image_dir\n        self.img_list = os.listdir(image_dir)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.img_list)\n\n    def __getitem__(self, idx):\n        # Load image\n        image_path = os.path.join(self.image_dir, self.img_list[idx])\n        image = cv2.imread(image_path)  \n\n        if self.transform:\n            data = self.transform(image=image)\n            image = data['image']\n        return self.img_list[idx][:-4],image","metadata":{"execution":{"iopub.status.busy":"2023-06-24T11:28:40.052102Z","iopub.execute_input":"2023-06-24T11:28:40.052463Z","iopub.status.idle":"2023-06-24T11:28:40.065607Z","shell.execute_reply.started":"2023-06-24T11:28:40.052432Z","shell.execute_reply":"2023-06-24T11:28:40.064679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EnsembleModel:\n    def __init__(self):\n        self.models = []\n\n    def __call__(self, x):\n        outputs = [model(x) for model in self.models]\n        outputs = torch.stack(outputs, dim=0)\n        avg_preds = torch.mean(outputs, dim=0)\n        return avg_preds\n\n    def add_model(self, model):\n        self.models.append(model)\n\ndef build_ensemble_model():\n    model = EnsembleModel()\n    model_types = CFG.model_types\n    encoder_name_list = CFG.encoder_name_list\n    for _type, encoder_name in zip(model_types, encoder_name_list):\n        model_dir = CFG.chepoint_dir + _type + '/' + encoder_name + '/'\n        model_list = os.listdir(model_dir)\n        for i in range(len(model_list)):\n            if _type == 'Unet':\n                _model = smp.Unet(encoder_name=encoder_name, activation='sigmoid', encoder_depth=CFG.encoder_depth, decoder_channels=CFG.decoder_channels, encoder_weights=None)\n            elif _type == 'PSP':\n                _model = smp.PSPNet(encoder_name=encoder_name, activation='sigmoid', encoder_depth=CFG.encoder_depth, decoder_channels=CFG.decoder_channels, encoder_weights=None)\n            elif _type == 'FPN':\n                _model = smp.FPN(encoder_name=encoder_name, activation='sigmoid', encoder_depth=CFG.encoder_depth, decoder_channels=CFG.decoder_channels, encoder_weights=None)\n            elif _type == 'PAN':\n                _model = smp.PAN(encoder_name=encoder_name, activation='sigmoid', encoder_depth=CFG.encoder_depth, decoder_channels=CFG.decoder_channels, encoder_weights=None)\n            elif _type == 'UnetPlusPlus':\n                _model = smp.UnetPlusPlus(encoder_name=encoder_name, activation='sigmoid', encoder_depth=CFG.encoder_depth, decoder_channels=CFG.decoder_channels, encoder_weights=None)\n            _model.to(CFG.device)\n            model_path = model_dir + model_list[i]\n            print(model_path)\n            state = torch.load(model_path, map_location=CFG.device)\n#             _model.load_state_dict(state['model_state_dict'])\n            _model.load_state_dict(state)\n            _model.eval()\n            model.add_model(_model)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-06-24T11:28:40.069337Z","iopub.execute_input":"2023-06-24T11:28:40.069692Z","iopub.status.idle":"2023-06-24T11:28:40.084630Z","shell.execute_reply.started":"2023-06-24T11:28:40.069658Z","shell.execute_reply":"2023-06-24T11:28:40.082914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def TTA(x: torch.Tensor, model: nn.Module):\n    # x.shape=(batch,c,h,w)\n    shape = x.shape\n    x = [x, *[torch.rot90(x, k=i, dims=(-2, -1)) for i in range(1, 4)]]\n    x = torch.cat(x, dim=0)\n    x = model(x)\n    x = x.reshape(4, shape[0], 1, *shape[-2:])\n    x = [torch.rot90(x[i], k=-i, dims=(-2, -1)) for i in range(4)]\n    x = torch.stack(x, dim=0)\n    return x.mean(0)","metadata":{"execution":{"iopub.status.busy":"2023-06-24T11:28:40.086212Z","iopub.execute_input":"2023-06-24T11:28:40.086577Z","iopub.status.idle":"2023-06-24T11:28:40.097314Z","shell.execute_reply.started":"2023-06-24T11:28:40.086547Z","shell.execute_reply":"2023-06-24T11:28:40.096304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = build_ensemble_model()","metadata":{"execution":{"iopub.status.busy":"2023-06-24T11:28:40.099678Z","iopub.execute_input":"2023-06-24T11:28:40.100421Z","iopub.status.idle":"2023-06-24T11:28:55.705411Z","shell.execute_reply.started":"2023-06-24T11:28:40.100391Z","shell.execute_reply":"2023-06-24T11:28:55.704435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Test:\n    \n#     def encode_binary_mask(self,mask):\n#         if mask.dtype != bool:\n#             raise ValueError(\n#                 \"encode_binary_mask expects a binary mask, received dtype == %s\" %\n#                 mask.dtype)\n\n#         mask = np.squeeze(mask)\n#         if len(mask.shape) != 2:\n#             raise ValueError(\n#                 \"encode_binary_mask expects a 2d mask, received shape == %s\" %\n#                 mask.shape)\n\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#         encoded_mask = coco_mask.encode(mask_to_encode)[0][\"counts\"]\n#         binary_str = zlib.compress(encoded_mask, zlib.Z_BEST_COMPRESSION)\n#         base64_str = base64.b64encode(binary_str)\n#         return base64_str\n    \n    def encode_binary_mask(self, mask: np.ndarray) -> t.Text:\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(\n            \"encode_binary_mask expects a binary mask, received dtype == %s\" %\n            mask.dtype)\n\n      mask = np.squeeze(mask)\n      if len(mask.shape) != 2:\n        raise ValueError(\n            \"encode_binary_mask expects a 2d mask, received shape == %s\" %\n            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\n    \n    \n    def encode_output(self,outputs,idx):\n        blood_vessel = torch.argmax(outputs, 1) \n        blood_vessel = blood_vessel == 1\n        blood_vessel = blood_vessel * 1\n    \n        blood_vessel = blood_vessel.cpu().numpy()\n        all_encode = {} \n        for i in range(blood_vessel.shape[0]):\n            list_encode = []\n            sliceImage = blood_vessel[i,:,:]\n            binarized = sliceImage > 0\n            coded_len = self.encode_binary_mask(binarized)\n            list_encode.append(coded_len)\n            all_encode[idx[i]] =list_encode\n        return all_encode\n\n    \n   \n    def get_test_transforms(self):\n        return A.Compose(CFG.test_aug)\n    \n    def test_dataloader(self,image_folder):\n        dataset = HUBMAPDataset(image_dir=image_folder, \n                                transform=self.get_test_transforms())\n        return DataLoader(dataset, batch_size=CFG.batch_size,shuffle=False, num_workers=4)\n    \n    \n\n    def evaluate(self,model):\n#         predictions = []\n#         outputdict ={}\n        ids = []\n        heights = []\n        widths = []\n        prediction_strings = []\n        sample = None\n        with torch.no_grad():\n            test_dataloader = self.test_dataloader(CFG.data_path + 'test/')\n            bar = tqdm(enumerate(test_dataloader), total=len(test_dataloader))\n            \n            for step, (idx, images) in bar:\n                images = images.to(CFG.device)\n                if CFG.is_tta:\n                    pred = TTA(images, model)\n                else:\n                    pred = model(images)\n                pred = F.interpolate(pred, size=[CFG.org_size, CFG.org_size], mode='bilinear', align_corners=False)\n                if sample is None: sample=pred\n                pred_string = ''\n#               pred_scores = []\n#                 for m in range(pred.shape[0]):\n#                     score = pred[m][0].cpu().numpy()\n#                     score = np.mean(score[np.nonzero(score)])\n#                     pred_scores.append(score)\n                pred = (pred > CFG.th).float().cpu().numpy()\n                \n                for m in range(len(pred)):\n                    # 先膨胀，然后进行连通组件分析\n                    kernel = np.ones(shape=(3, 3), dtype=np.uint8)\n                    binary_mask = cv2.dilate((pred[m][0] * 255), kernel, 3)\n                    num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(binary_mask.astype(np.uint8))\n                    # 将每个连通组件保存为一个单独的 mask 图像\n                    for i in range(1, num_labels):\n                        # 创建一个空白图像，与输入 mask 图像大小和类型相同\n                        mask_i = np.zeros_like(binary_mask)\n                        # 选择标签为 i 的像素，将其赋值为白色\n                        mask_i[labels == i] = 1\n                        mask = mask_i[:, :, np.newaxis].astype(np.bool)\n                        score = 1.0\n                        encoded = self.encode_binary_mask(mask)\n                        if i==0:\n                            pred_string += f\"0 {score} {encoded.decode('utf-8')}\"\n\n                        else:\n                            pred_string += f\" 0 {score} {encoded.decode('utf-8')}\"\n                b, c, h, w = images.shape\n                ids.append(idx[0])\n                heights.append(h)\n                widths.append(w)\n                prediction_strings.append(pred_string)\n#             for step, (idx, images) in bar:\n#                 images = images.to(CFG.device)\n#                 if CFG.is_tta:\n#                     outputs = TTA(images, model)\n#                 else:\n#                     outputs = model(images)\n#                 outputs = (outputs > CFG.th).float()\n#                 encoded = self.encode_output(outputs, idx)\n#                 for key in encoded: \n#                     outputdict[key] = \" \".join([f\"0 1.0 {x.decode('utf-8')}\" for x in  encoded[key]])\n#                 print(outputdict)    \n        return ids, heights, widths, prediction_strings, sample\n        \ntest = Test()\nids, heights, widths, prediction_strings, sample=test.evaluate(model)","metadata":{"execution":{"iopub.status.busy":"2023-06-24T11:29:21.922628Z","iopub.execute_input":"2023-06-24T11:29:21.923127Z","iopub.status.idle":"2023-06-24T11:29:23.303279Z","shell.execute_reply.started":"2023-06-24T11:29:21.923079Z","shell.execute_reply":"2023-06-24T11:29:23.290701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"top10 = [sample[i].detach().permute(1,2,0).cpu().numpy() for i in range(min(10,len(sample)))]\nimg = 0\nfor i in top10:\n    img += i\n    img = np.clip(img, 0, 1)\nplt.imshow(img)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-06-24T11:29:27.734424Z","iopub.execute_input":"2023-06-24T11:29:27.734820Z","iopub.status.idle":"2023-06-24T11:29:28.029872Z","shell.execute_reply.started":"2023-06-24T11:29:27.734785Z","shell.execute_reply":"2023-06-24T11:29:28.028954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame()\nsubmission['id'] = ids\nsubmission['height'] = heights\nsubmission['width'] = widths\nsubmission['prediction_string'] = prediction_strings\nsubmission = submission.set_index('id')\nsubmission.to_csv(\"submission.csv\")\n!cat submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-06-24T11:29:33.147539Z","iopub.execute_input":"2023-06-24T11:29:33.147971Z","iopub.status.idle":"2023-06-24T11:29:34.522472Z","shell.execute_reply.started":"2023-06-24T11:29:33.147933Z","shell.execute_reply":"2023-06-24T11:29:34.520871Z"},"trusted":true},"execution_count":null,"outputs":[]}]}