{"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":"none","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":7079885,"sourceType":"datasetVersion","datasetId":4078343},{"sourceId":7080192,"sourceType":"datasetVersion","datasetId":4078562},{"sourceId":7081124,"sourceType":"datasetVersion","datasetId":4079220},{"sourceId":7081204,"sourceType":"datasetVersion","datasetId":4079283},{"sourceId":7116413,"sourceType":"datasetVersion","datasetId":4104090},{"sourceId":7124585,"sourceType":"datasetVersion","datasetId":4109820},{"sourceId":7148971,"sourceType":"datasetVersion","datasetId":4127365},{"sourceId":150248402,"sourceType":"kernelVersion"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Prepare","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nfrom tqdm import tqdm\nimport pandas as pd\nimport numpy as np\nfrom glob import glob\nimport gc\nimport time\nfrom collections import defaultdict\nimport  matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\nimport copy\nimport cv2\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import lr_scheduler\nimport torch.nn.functional as F\nfrom torch.cuda import amp\nimport torch.optim as optim\nimport albumentations as A","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:00.379000Z","iopub.execute_input":"2023-12-07T21:03:00.379882Z","iopub.status.idle":"2023-12-07T21:03:07.545594Z","shell.execute_reply.started":"2023-12-07T21:03:00.379840Z","shell.execute_reply":"2023-12-07T21:03:07.544434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_img(path):\n    img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n    img = np.tile(img[...,None], [1, 1, 3]) # gray to rgb\n    img = img.astype('float32') # original is uint16\n    mx = np.max(img)\n    if mx:\n        img/=mx # scale image to [0, 1]\n    return img\n\ndef load_msk(path):\n    msk = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n    msk = msk.astype('float32')\n    msk/=255.0\n    return msk","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:07.547537Z","iopub.execute_input":"2023-12-07T21:03:07.548078Z","iopub.status.idle":"2023-12-07T21:03:07.556450Z","shell.execute_reply.started":"2023-12-07T21:03:07.548044Z","shell.execute_reply":"2023-12-07T21:03:07.555103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BuildDataset(torch.utils.data.Dataset):\n    def __init__(self, img_paths, msk_paths=[], transforms=None):\n        self.img_paths  = img_paths\n        self.msk_paths  = msk_paths\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.img_paths)\n    \n    def __getitem__(self, index):\n        img_path  = self.img_paths[index]\n        img = load_img(img_path)\n        \n        if len(self.msk_paths)>0:\n            msk_path = self.msk_paths[index]\n            msk = load_msk(msk_path)\n            if self.transforms:\n                data = self.transforms(image=img, mask=msk)\n                img  = data['image']\n                msk  = data['mask']\n            img = np.transpose(img, (2, 0, 1))\n            return torch.tensor(img), torch.tensor(msk)\n        else:\n            orig_size = img.shape\n            if self.transforms:\n                data = self.transforms(image=img)\n                img  = data['image']\n            img = np.transpose(img, (2, 0, 1))\n            return torch.tensor(img), torch.tensor(np.array([orig_size[0], orig_size[1]]))","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:07.558809Z","iopub.execute_input":"2023-12-07T21:03:07.559288Z","iopub.status.idle":"2023-12-07T21:03:07.575571Z","shell.execute_reply.started":"2023-12-07T21:03:07.559242Z","shell.execute_reply":"2023-12-07T21:03:07.574418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATASET_FOLDER = \"/kaggle/input/blood-vessel-segmentation\"\nls_images = glob(os.path.join(DATASET_FOLDER, \"test\", \"*\", \"images\", \"*.tif\"))\nprint(f\"found images: {len(ls_images)}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:07.578294Z","iopub.execute_input":"2023-12-07T21:03:07.578872Z","iopub.status.idle":"2023-12-07T21:03:07.602244Z","shell.execute_reply.started":"2023-12-07T21:03:07.578838Z","shell.execute_reply":"2023-12-07T21:03:07.600947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    backbone = \"efficientnet-b1\"\n    train_bs = 12\n    valid_bs = 24\n    img_size = [512,512]\n    epochs = 10\n    lr = 2e-3\n    threshold = 0.998\n    bin_path = \"/kaggle/input/unet-attention-adam/best_epoch (10).bin\"\n\n    num_classes   = 1\n    device        = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    data_transforms = {\n        \"train\": A.Compose([\n            A.Resize(*img_size, interpolation=cv2.INTER_NEAREST),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n        ], p=1.0),\n        \n        \"valid\": A.Compose([\n            A.Resize(*img_size, interpolation=cv2.INTER_NEAREST),\n        ], p=1.0)\n    }\n    ","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:07.604177Z","iopub.execute_input":"2023-12-07T21:03:07.604954Z","iopub.status.idle":"2023-12-07T21:03:07.615433Z","shell.execute_reply.started":"2023-12-07T21:03:07.604912Z","shell.execute_reply":"2023-12-07T21:03:07.613930Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# for new way remember to make batch_size to 1","metadata":{}},{"cell_type":"code","source":"test_dataset = BuildDataset(ls_images, [], transforms=CFG.data_transforms[\"valid\"])\ntest_loader = DataLoader(test_dataset, batch_size = 1, num_workers=0, shuffle=False, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:07.617548Z","iopub.execute_input":"2023-12-07T21:03:07.618295Z","iopub.status.idle":"2023-12-07T21:03:07.629675Z","shell.execute_reply.started":"2023-12-07T21:03:07.618261Z","shell.execute_reply":"2023-12-07T21:03:07.628166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class attention_block(nn.Module):\n    def __init__(self, F_g, F_l, n_coefficients):\n        super(attention_block, self).__init__()\n\n        self.bypass = nn.Sequential(\n            nn.Conv2d(F_g, n_coefficients, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(n_coefficients)\n        )\n\n        self.upsample = nn.Sequential(\n            nn.Conv2d(F_l, n_coefficients, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(n_coefficients)\n        )\n\n        self.psi = nn.Sequential(\n            nn.Conv2d(n_coefficients, 1, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(1),\n            nn.Sigmoid()\n        )\n\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, gate, skip_connection):\n        \"\"\"\n        :param gate: gating signal from previous layer\n        :param skip_connection: activation from corresponding encoder layer\n        :return: output activations\n        \"\"\"\n        g1 = self.upsample(gate)\n        x1 = self.bypass(skip_connection)\n        psi = self.relu(g1 + x1)\n        psi = self.psi(psi)\n        out = skip_connection * psi\n        return out\n\nclass MyUNet(nn.Module):\n    def contracting_block(self, in_channels, out_channels, kernel_size=3):\n        block = torch.nn.Sequential(\n                    torch.nn.Conv2d(kernel_size=kernel_size, in_channels=in_channels, out_channels=out_channels, padding=1),\n                    torch.nn.ReLU(),\n                    torch.nn.BatchNorm2d(out_channels),\n                    torch.nn.Conv2d(kernel_size=kernel_size, in_channels=out_channels, out_channels=out_channels, padding=1),\n                    torch.nn.ReLU(),\n                )\n        return block\n    def expansive_block(self, in_channels, mid_channel, out_channels, kernel_size=3):\n            block = torch.nn.Sequential(\n                    torch.nn.Conv2d(kernel_size=kernel_size, in_channels=in_channels, out_channels=mid_channel, padding=1),\n                    torch.nn.ReLU(),\n                    torch.nn.BatchNorm2d(mid_channel),\n                    torch.nn.Conv2d(kernel_size=kernel_size, in_channels=mid_channel, out_channels=mid_channel, padding=1),\n                    torch.nn.ReLU(),\n                    torch.nn.BatchNorm2d(mid_channel),\n                    torch.nn.ConvTranspose2d(in_channels=mid_channel, out_channels=out_channels, kernel_size=3, stride=2, padding=1, output_padding=1)\n            )\n            return  block\n\n    def final_block(self, in_channels, mid_channel, out_channels, kernel_size=1):\n            block = torch.nn.Sequential(\n                    \n                    torch.nn.Conv2d(kernel_size=3, in_channels=in_channels, out_channels=mid_channel, padding=1),\n                    torch.nn.ReLU(),\n                    torch.nn.BatchNorm2d(mid_channel),\n                    torch.nn.Conv2d(kernel_size=kernel_size, in_channels=mid_channel, out_channels=out_channels, padding=0),\n                    torch.nn.ReLU(),\n                    torch.nn.BatchNorm2d(out_channels),\n            )\n            return  block\n\n    def __init__(self, in_channel, out_channel):\n        super(MyUNet, self).__init__()\n        #Encode\n        self.conv_encode1 = self.contracting_block(in_channels=in_channel, out_channels=64)\n        self.conv_maxpool1 = torch.nn.MaxPool2d(kernel_size=2)\n        self.conv_encode2 = self.contracting_block(64, 128)\n        self.conv_maxpool2 = torch.nn.MaxPool2d(kernel_size=2)\n        self.conv_encode3 = self.contracting_block(128, 256)\n        self.conv_maxpool3 = torch.nn.MaxPool2d(kernel_size=2)\n        # Bottleneck\n        self.bottleneck = torch.nn.Sequential(\n                            torch.nn.Conv2d(kernel_size=3, in_channels=256, out_channels=512, padding=1),\n                            torch.nn.ReLU(),\n                            torch.nn.BatchNorm2d(512),\n                            torch.nn.Conv2d(kernel_size=3, in_channels=512, out_channels=512, padding=1),\n                            torch.nn.ReLU(),\n                            torch.nn.BatchNorm2d(512),\n                            torch.nn.ConvTranspose2d(in_channels=512, out_channels=256, kernel_size=3, stride=2, padding=1, output_padding=1)\n                            )\n\n        # Decode\n        self.attention3 = attention_block(256, 256, 128)\n        self.conv_decode3 = self.expansive_block(512, 256, 128)\n        self.attention2 = attention_block(128, 128, 64)\n        self.conv_decode2 = self.expansive_block(256, 128, 64)\n        self.attention1 = attention_block(64, 64, 32)\n        self.final_layer = self.final_block(128, 64, out_channel)\n    \n    def forward(self, x):\n        # Encode\n        encode_block1 = self.conv_encode1(x)\n        encode_pool1 = self.conv_maxpool1(encode_block1)\n        encode_block2 = self.conv_encode2(encode_pool1)\n        encode_pool2 = self.conv_maxpool2(encode_block2)\n        encode_block3 = self.conv_encode3(encode_pool2)\n        encode_pool3 = self.conv_maxpool3(encode_block3)\n        # Bottleneck\n        bottle_neck1 = self.bottleneck(encode_pool3)\n        # Decode\n        att_3 = self.attention3(bottle_neck1, encode_block3) \n        decode_block1 = torch.cat((bottle_neck1, att_3), 1)\n        cat_layer2 = self.conv_decode3(decode_block1)\n        \n        att_2 = self.attention2(cat_layer2, encode_block2) \n        decode_block2 = torch.cat((cat_layer2, att_2), 1)\n        cat_layer1 = self.conv_decode2(decode_block2)\n        \n        \n        att_1 = self.attention1(cat_layer1, encode_block1) \n        decode_block3 = torch.cat((cat_layer1, att_1), 1)\n        final_layer = self.final_layer(decode_block3)\n        return final_layer\ndef build_model(backbone, num_classes, device):\n    model = MyUNet(in_channel=3, out_channel=num_classes)\n    model.to(device)\n    return model\n\ndef load_model(backbone, num_classes, device, path):\n    model = build_model(backbone, num_classes, device)\n    try:\n        model.load_state_dict(torch.load(path))\n    except:\n        model.load_state_dict(torch.load(path, map_location=torch.device('cpu')))\n    model.eval()\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:07.631550Z","iopub.execute_input":"2023-12-07T21:03:07.632175Z","iopub.status.idle":"2023-12-07T21:03:07.667875Z","shell.execute_reply.started":"2023-12-07T21:03:07.632141Z","shell.execute_reply":"2023-12-07T21:03:07.666697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nmodel = load_model(CFG.backbone, \n                   1, \n                   device, \n                   CFG.bin_path)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:07.669209Z","iopub.execute_input":"2023-12-07T21:03:07.670865Z","iopub.status.idle":"2023-12-07T21:03:08.516179Z","shell.execute_reply.started":"2023-12-07T21:03:07.670824Z","shell.execute_reply":"2023-12-07T21:03:08.515220Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    rle = ' '.join(str(x) for x in runs)\n    if rle=='':\n        rle = '1 0'\n    return rle","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:08.517513Z","iopub.execute_input":"2023-12-07T21:03:08.518167Z","iopub.status.idle":"2023-12-07T21:03:08.526394Z","shell.execute_reply.started":"2023-12-07T21:03:08.518131Z","shell.execute_reply":"2023-12-07T21:03:08.525070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def remove_small_objects(img, min_size):\n    # Find all connected components (labels)\n    num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(img, connectivity=8)\n\n    # Create a mask where small objects are removed\n    new_img = np.zeros_like(img)\n    for label in range(1, num_labels):\n        if stats[label, cv2.CC_STAT_AREA] >= min_size:\n            new_img[labels == label] = 255\n\n    return new_img","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:08.531304Z","iopub.execute_input":"2023-12-07T21:03:08.532377Z","iopub.status.idle":"2023-12-07T21:03:08.545596Z","shell.execute_reply.started":"2023-12-07T21:03:08.532325Z","shell.execute_reply":"2023-12-07T21:03:08.543671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def patch_image(img, patch_size, model = None, over_lap=0.1):\n    \"\"\"\n    Splits the image into patches with overlap.\n\n    \"\"\"\n    shape = img.shape\n\n    height, width = shape[2],shape[3]\n\n    stride = patch_size * (1 - over_lap)\n    num_patches = np.ceil(np.array([height, width]) / stride).astype(np.int64)\n    starts = [np.int64(np.linspace(0, width - patch_size, num_patches[1])),\n              np.int64(np.linspace(0, height - patch_size, num_patches[0]))]\n    patches = []\n\n    for y in starts[1]:\n        for x in starts[0]:\n            if model != None: \n                patch_img = img[:,:,y:y + patch_size, x:x + patch_size]\n#                 print(f'inside the patch function: size of patched image: {np.shape(patch_img)}')\n                pred = model(patch_img)\n                patches.append(pred)\n           \n\n    return patches\n\n\ndef combine_patches_torch(patches, original_shape, patch_size, over_lap=0.1):\n    height, width = original_shape[2],original_shape[3]\n    stride = int(patch_size * (1 - over_lap))\n    combined = np.zeros((height, width), dtype=np.float32)\n    weight = np.zeros((height, width), dtype=np.float32)\n\n    num_patches_y = np.ceil(height / stride).astype(np.int64)\n    num_patches_x = np.ceil(width / stride).astype(np.int64)\n\n    starts_y = np.linspace(0, height - patch_size, num_patches_y).astype(np.int64)\n    starts_x = np.linspace(0, width - patch_size, num_patches_x).astype(np.int64)\n\n    patch_idx = 0\n    for y in starts_y:\n        for x in starts_x:\n            \n            patch = patches[patch_idx].detach().cpu()\n            patch = patch.numpy().astype(np.float32)\n#             print(f'inside the combine function: type of combine = {type(combined)}, shape of patches = {np.shape(patch)}')\n            # with torch, I cannot add different sized tensor together \n            combined[y:y + patch_size, x:x + patch_size] += patch.squeeze()\n            weight[y:y + patch_size, x:x + patch_size] += 1.0\n            patch_idx += 1\n\n    # Avoid division by zero\n    weight[weight == 0] = 1.0\n    combined = combined / weight\n    combined = torch.from_numpy(combined)\n    combined = combined.unsqueeze(0).unsqueeze(0)\n    \n    return combined","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:08.548621Z","iopub.execute_input":"2023-12-07T21:03:08.550055Z","iopub.status.idle":"2023-12-07T21:03:08.572853Z","shell.execute_reply.started":"2023-12-07T21:03:08.550005Z","shell.execute_reply":"2023-12-07T21:03:08.571471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"new way","metadata":{}},{"cell_type":"code","source":"rles = []\npbar = tqdm(enumerate(test_loader), total=len(test_loader), desc='Inference ')\nfor step, (images, shapes) in pbar:\n    shapes = shapes.numpy()\n    images = images.to(CFG.device, dtype=torch.float)\n    with torch.no_grad():\n        preds = model(images)\n        preds = (nn.Sigmoid()(preds)>0.998).double()\n    preds = preds.cpu().numpy().astype(np.uint8)\n    for pred, shape in zip(preds, shapes):\n        pred = cv2.resize(pred[0], (shape[1], shape[0]), cv2.INTER_NEAREST)\n        pred = remove_small_objects(pred, 10)\n        rle = rle_encode(pred)\n        rles.append(rle)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:08.574577Z","iopub.execute_input":"2023-12-07T21:03:08.576036Z","iopub.status.idle":"2023-12-07T21:03:34.072033Z","shell.execute_reply.started":"2023-12-07T21:03:08.575978Z","shell.execute_reply":"2023-12-07T21:03:34.070648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ids = []\nfor p_img in tqdm(ls_images):\n    path_ = p_img.split(os.path.sep)\n    # parse the submission ID\n    dataset = path_[-3]\n    slice_id, _ = os.path.splitext(path_[-1])\n    ids.append(f\"{dataset}_{slice_id}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:34.073566Z","iopub.execute_input":"2023-12-07T21:03:34.074021Z","iopub.status.idle":"2023-12-07T21:03:34.085902Z","shell.execute_reply.started":"2023-12-07T21:03:34.073989Z","shell.execute_reply":"2023-12-07T21:03:34.084695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame.from_dict({\n    \"id\": ids,\n    \"rle\": rles\n})\nsubmission.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:34.087431Z","iopub.execute_input":"2023-12-07T21:03:34.088238Z","iopub.status.idle":"2023-12-07T21:03:34.113827Z","shell.execute_reply.started":"2023-12-07T21:03:34.088192Z","shell.execute_reply":"2023-12-07T21:03:34.112857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2023-12-07T21:03:34.115169Z","iopub.execute_input":"2023-12-07T21:03:34.116040Z","iopub.status.idle":"2023-12-07T21:03:34.142253Z","shell.execute_reply.started":"2023-12-07T21:03:34.115986Z","shell.execute_reply":"2023-12-07T21:03:34.140971Z"},"trusted":true},"execution_count":null,"outputs":[]}]}