{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":7122490,"sourceType":"datasetVersion","datasetId":4108361},{"sourceId":7141413,"sourceType":"datasetVersion","datasetId":4121860},{"sourceId":7149637,"sourceType":"datasetVersion","datasetId":4127856},{"sourceId":7150171,"sourceType":"datasetVersion","datasetId":4128229},{"sourceId":7150740,"sourceType":"datasetVersion","datasetId":4128617},{"sourceId":150248402,"sourceType":"kernelVersion"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Prepare","metadata":{}},{"cell_type":"code","source":"!python -m pip install --no-index --find-links=/kaggle/input/pip-download-for-segmentation-models-pytorch segmentation-models-pytorch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-12-09T04:18:04.253875Z","iopub.execute_input":"2023-12-09T04:18:04.254324Z","iopub.status.idle":"2023-12-09T04:18:23.725512Z","shell.execute_reply.started":"2023-12-09T04:18:04.254286Z","shell.execute_reply":"2023-12-09T04:18:23.724342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\nfrom torch.cuda import amp\nimport torch.optim as optim\nimport albumentations as A\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2023-12-09T04:18:23.727479Z","iopub.execute_input":"2023-12-09T04:18:23.727786Z","iopub.status.idle":"2023-12-09T04:18:32.262124Z","shell.execute_reply.started":"2023-12-09T04:18:23.727757Z","shell.execute_reply":"2023-12-09T04:18:32.261322Z"},"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-09T04:18:32.263232Z","iopub.execute_input":"2023-12-09T04:18:32.263528Z","iopub.status.idle":"2023-12-09T04:18:32.270067Z","shell.execute_reply.started":"2023-12-09T04:18:32.263502Z","shell.execute_reply":"2023-12-09T04:18:32.269064Z"},"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-09T04:18:32.272184Z","iopub.execute_input":"2023-12-09T04:18:32.272580Z","iopub.status.idle":"2023-12-09T04:18:32.309604Z","shell.execute_reply.started":"2023-12-09T04:18:32.272544Z","shell.execute_reply":"2023-12-09T04:18:32.308618Z"},"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)}\")\n\n# ===============================================================================\n# /kaggle/input/blood-vessel-segmentation/train/kidney_1_dense\n# /kaggle/input/blood-vessel-segmentation/train/kidney_2\n# sanity check\nbase_path = '/kaggle/input/blood-vessel-segmentation/train'  \ntrain_base_path = '/kaggle/input/patched-sennet-kidney-1-data'\nval_img = \"kidney_2\"#'kidney_3_sparse',kidney_1_dense\nval_mask = \"kidney_2\"#'kidney_3_sparse',kidney_1_dense\n\nimages_val_path = os.path.join(base_path, val_img, 'images')\nlabels_val_path = os.path.join(base_path, val_mask, 'labels')\nimage_val_files = sorted([os.path.join(images_val_path, f) for f in os.listdir(images_val_path) if f.endswith('.tif')])\nlabel_val_files = sorted([os.path.join(labels_val_path, f) for f in os.listdir(labels_val_path) if f.endswith('.tif')])","metadata":{"execution":{"iopub.status.busy":"2023-12-09T04:18:32.310935Z","iopub.execute_input":"2023-12-09T04:18:32.311344Z","iopub.status.idle":"2023-12-09T04:18:32.614844Z","shell.execute_reply.started":"2023-12-09T04:18:32.311309Z","shell.execute_reply":"2023-12-09T04:18:32.614134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = BuildDataset(ls_images, [], transforms=None)\nsanity_dataset = BuildDataset(image_val_files,label_val_files,transforms=None )\nsanity_loader = DataLoader(sanity_dataset, batch_size = 1, num_workers=0, shuffle=False, pin_memory=True)\ntest_loader = DataLoader(test_dataset, batch_size = 1, num_workers=0, shuffle=False, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T04:18:32.616007Z","iopub.execute_input":"2023-12-09T04:18:32.616366Z","iopub.status.idle":"2023-12-09T04:18:32.623911Z","shell.execute_reply.started":"2023-12-09T04:18:32.616332Z","shell.execute_reply":"2023-12-09T04:18:32.622655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# resnext50_32x4d\n# resnet50\n\n#/kaggle/input/resnext-800/best_epoch800.bin\n#/kaggle/input/renext-patch-50d/renext_patch.bin\nclass CFG:\n    backbone = \"resnext50_32x4d\"\n    bin_path = '/kaggle/input/unet-attetion-patch/best_epoch (11).bin'\n    test_batchsize = 1\n    img_size = [800,800]\n    remove_area = 30\n    threshold = 0.9\n    over_lap = 0.1\n    num_classes   = 1\n    patch_size = 800\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-09T04:18:32.625664Z","iopub.execute_input":"2023-12-09T04:18:32.626595Z","iopub.status.idle":"2023-12-09T04:18:32.712597Z","shell.execute_reply.started":"2023-12-09T04:18:32.626562Z","shell.execute_reply":"2023-12-09T04:18:32.711413Z"},"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":"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    model.load_state_dict(torch.load(path))\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-12-09T04:18:32.714222Z","iopub.execute_input":"2023-12-09T04:18:32.714605Z","iopub.status.idle":"2023-12-09T04:18:32.743781Z","shell.execute_reply.started":"2023-12-09T04:18:32.714574Z","shell.execute_reply":"2023-12-09T04:18:32.742758Z"},"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                   CFG.device, \n                   CFG.bin_path)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T04:18:32.745146Z","iopub.execute_input":"2023-12-09T04:18:32.746094Z","iopub.status.idle":"2023-12-09T04:18:36.580821Z","shell.execute_reply.started":"2023-12-09T04:18:32.746061Z","shell.execute_reply":"2023-12-09T04:18:36.579814Z"},"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-09T04:18:36.584693Z","iopub.execute_input":"2023-12-09T04:18:36.585040Z","iopub.status.idle":"2023-12-09T04:18:36.591253Z","shell.execute_reply.started":"2023-12-09T04:18:36.585012Z","shell.execute_reply":"2023-12-09T04:18:36.590404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef 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-09T04:18:36.592372Z","iopub.execute_input":"2023-12-09T04:18:36.592692Z","iopub.status.idle":"2023-12-09T04:18:36.602001Z","shell.execute_reply.started":"2023-12-09T04:18:36.592667Z","shell.execute_reply":"2023-12-09T04:18:36.601191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def patch_image(img, patch_size, model = None, over_lap=0.2):\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    \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    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(type(patch_img),\" and shape is \",patch_img.shape)\n#                 print(f'inside the patch function: size of patched image: {np.shape(patch_img)}')\n                patches.append(patch_img)\n    patches = torch.cat(patches,dim = 0)\n    pred = model(patches)\n    return pred\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-09T04:18:36.603236Z","iopub.execute_input":"2023-12-09T04:18:36.603733Z","iopub.status.idle":"2023-12-09T04:18:36.616993Z","shell.execute_reply.started":"2023-12-09T04:18:36.603701Z","shell.execute_reply":"2023-12-09T04:18:36.616094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# inference for resize model","metadata":{}},{"cell_type":"code","source":"# rles = []\n# pbar = tqdm(enumerate(test_loader), total=len(test_loader), desc='Inference ')\n# for step, (images, shapes) in pbar:\n# #     print(np.shape(images))\n#     shapes = shapes.numpy()\n#     images = images.to(device, dtype=torch.float)\n#     with torch.no_grad():\n#         preds = model(images)\n#         print(f'shape of raw prediction {np.shape(preds)}')\n# #         preds = (nn.Sigmoid()(preds)>0.5).double()\n#         preds = (preds>CFG.threshold).float()\n#         print(f'shape of prediction after sigmoid {np.shape(preds)}')\n#     preds = preds.cpu().numpy().astype(np.uint8)\n\n#     for pred, shape in zip(preds, shapes):\n#         pred = cv2.resize(pred[0], (shape[1], shape[0]), cv2.INTER_NEAREST)\n#         rle = rle_encode_submission(remove_small_objects(pred,30))\n#         rles.append(rle)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T04:18:36.618134Z","iopub.execute_input":"2023-12-09T04:18:36.618404Z","iopub.status.idle":"2023-12-09T04:18:36.632206Z","shell.execute_reply.started":"2023-12-09T04:18:36.618381Z","shell.execute_reply":"2023-12-09T04:18:36.631529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference for patch model","metadata":{}},{"cell_type":"code","source":"rles = []\npbar = tqdm(enumerate(test_loader), total=len(test_loader), desc='Inference ')\nfor step, (images, shapes) in pbar:\n#     print(np.shape(images))\n    shapes = shapes.numpy()\n    images = images.to(device, dtype=torch.float)\n    with torch.no_grad():\n        ori_shape = images.shape\n#         print(ori_shape)\n#         print(np.shape(images))\n        patches = patch_image(images,patch_size = CFG.patch_size,model = model,over_lap = CFG.over_lap)\n        pred  = combine_patches_torch(patches,ori_shape,patch_size = CFG.patch_size,over_lap = CFG.over_lap )\n        pred = pred.cpu().numpy()\n        pred = ((pred>CFG.threshold)*255).astype(np.uint8)\n    pred = pred.squeeze()\n#     print(f'shape of raw prediction {np.shape(pred)}')\n    rle = rle_encode(remove_small_objects(pred,CFG.remove_area))\n    rles.append(rle)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T04:18:36.633316Z","iopub.execute_input":"2023-12-09T04:18:36.633657Z","iopub.status.idle":"2023-12-09T04:18:47.386579Z","shell.execute_reply.started":"2023-12-09T04:18:36.633623Z","shell.execute_reply":"2023-12-09T04:18:47.385637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Sanity check","metadata":{}},{"cell_type":"code","source":"# sample_ids = [random.randint(0, len(sanity_dataset)) for _ in range(20)]\n# for id in sample_ids:\n#     images, mask = sanity_dataset[id]\n#     images = images.to(device, dtype=torch.float)\n#     with torch.no_grad():\n        \n#         image = images.unsqueeze(0)\n#         ori_shape = image.shape\n#         patches = patch_image(image,patch_size = CFG.patch_size,model = model,over_lap = CFG.over_lap)\n#         pred  = combine_patches_torch(patches,ori_shape,patch_size = CFG.patch_size,over_lap = CFG.over_lap )\n#         pred = (nn.Sigmoid()(pred)>.5).float()\n#         pred = pred.cpu().numpy().astype(np.uint8)\n#         plt.figure(figsize=(9, 4))\n#         plt.subplot(1,3,1)\n#         image = image.cpu().detach().numpy()\n#         image = image.squeeze(0)\n#         image = np.transpose(image,(1,2,0))\n#         plt.imshow(image)\n#         plt.subplot(1,3,2)\n#         plt.imshow(mask)\n#         plt.subplot(1,3,3)\n#         pred = pred.squeeze()\n#         plt.imshow(pred)\n#         plt.show()\n    \n\n# for images, shapes in test_loader:\n# #     print(np.shape(images))\n#     shapes = shapes.numpy()\n#     images = images.to(device, dtype=torch.float)\n#     with torch.no_grad():\n#         ori_shape = images.shape\n# #         print(ori_shape)\n#         image = images[0]\n#         patches = patch_image(images,512,model = model)\n#         pred  = combine_patches_torch(patches,ori_shape,512)\n#         pred = (nn.Sigmoid()(pred)>.9).float()\n#         pred = pred.cpu().numpy().astype(np.uint8)\n#         plt.figure(figsize=(9, 4))\n#         plt.subplot(1,2,1)\n#         image = image.cpu().detach().numpy()\n#         image = np.transpose(image,(1,2,0))\n#         plt.imshow(image)\n#         plt.subplot(1,2,2)\n#         pred = pred.squeeze()\n#         plt.imshow(pred)","metadata":{"execution":{"iopub.status.busy":"2023-12-09T04:18:47.388017Z","iopub.execute_input":"2023-12-09T04:18:47.388450Z","iopub.status.idle":"2023-12-09T04:18:47.395247Z","shell.execute_reply.started":"2023-12-09T04:18:47.388411Z","shell.execute_reply":"2023-12-09T04:18:47.394208Z"},"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-09T04:18:47.396345Z","iopub.execute_input":"2023-12-09T04:18:47.396638Z","iopub.status.idle":"2023-12-09T04:18:47.412033Z","shell.execute_reply.started":"2023-12-09T04:18:47.396613Z","shell.execute_reply":"2023-12-09T04:18:47.411164Z"},"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-09T04:18:47.413168Z","iopub.execute_input":"2023-12-09T04:18:47.413508Z","iopub.status.idle":"2023-12-09T04:18:47.433591Z","shell.execute_reply.started":"2023-12-09T04:18:47.413477Z","shell.execute_reply":"2023-12-09T04:18:47.432746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2023-12-09T04:18:47.434626Z","iopub.execute_input":"2023-12-09T04:18:47.434969Z","iopub.status.idle":"2023-12-09T04:18:47.452454Z","shell.execute_reply.started":"2023-12-09T04:18:47.434935Z","shell.execute_reply":"2023-12-09T04:18:47.451393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}