{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","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":1807973,"sourceType":"datasetVersion","datasetId":1074109},{"sourceId":7553845,"sourceType":"datasetVersion","datasetId":4363835},{"sourceId":150248402,"sourceType":"kernelVersion"}],"dockerImageVersionId":30648,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints/\n!cp /kaggle/input/se-net-pretrained-imagenet-weights/* /root/.cache/torch/hub/checkpoints/\n!python -m pip install --no-index -q --find-links=/kaggle/input/pip-download-for-segmentation-models-pytorch segmentation-models-pytorch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-01T16:58:36.451703Z","iopub.execute_input":"2024-02-01T16:58:36.452208Z","iopub.status.idle":"2024-02-01T16:59:12.796196Z","shell.execute_reply.started":"2024-02-01T16:58:36.452181Z","shell.execute_reply":"2024-02-01T16:59:12.794918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nfrom albumentations.pytorch import ToTensorV2\nfrom torch.nn.parallel import DataParallel\nimport segmentation_models_pytorch as smp\nfrom torch.cuda.amp import autocast\nimport torch.nn.functional as fn\nimport matplotlib.pyplot as plt\nimport albumentations as alb\nfrom torch import optim\nimport os, sys, cv2, gc\nfrom tqdm import tqdm\nfrom glob import glob\nimport torch.nn as nn\nimport pandas as pd\nimport numpy as np\nimport datetime\nimport ctypes\nimport torch\nimport gc\n\nDEBUG = False\nlibc = ctypes.CDLL(\"libc.so.6\")\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"execution":{"iopub.status.busy":"2024-02-01T16:59:12.798466Z","iopub.execute_input":"2024-02-01T16:59:12.798816Z","iopub.status.idle":"2024-02-01T16:59:28.082863Z","shell.execute_reply.started":"2024-02-01T16:59:12.798780Z","shell.execute_reply":"2024-02-01T16:59:28.081884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    \n    model_name = \"Unet\"\n    backbone = \"se_resnext50_32x4d\"\n    path_to_model = \"/kaggle/input/sennet-models/se_resnext50_32x4d_19epoch_1024image_size_2024_02_04_17_00.pt\"\n\n    target_size = 1\n    in_chans = 1\n    image_size = 1024\n    input_size = 1024\n    tile_size = image_size\n    stride = tile_size // 4\n    drop_egde_pixel = 16\n    \n    batch_size = 1\n    eval_batch_size = batch_size\n\n    epochs = 20\n    lr = 5e-5\n    chopping_percentile = 1e-3\n    th_percentile = 0.00123\n\n    train_aug = alb.Compose(\n        [\n            alb.Rotate(limit=45, p=0.5),\n            alb.RandomScale(scale_limit=(0.8,1.25),interpolation=cv2.INTER_CUBIC,p=0.5),\n            alb.RandomCrop(input_size, input_size,p=1),\n            alb.RandomGamma(p=0.75),\n            alb.RandomBrightnessContrast(p=0.5,),\n            alb.GaussianBlur(p=0.5),\n            alb.MotionBlur(p=0.5),\n            alb.GridDistortion(num_steps=5, distort_limit=0.3, p=0.5),\n            alb.pytorch.ToTensorV2(transpose_mask=True),\n        ]\n    )\n    eval_aug = alb.Compose(\n        [\n            alb.pytorch.ToTensorV2(transpose_mask=True),\n        ]\n    )","metadata":{"execution":{"iopub.status.busy":"2024-02-01T16:59:28.088923Z","iopub.execute_input":"2024-02-01T16:59:28.089545Z","iopub.status.idle":"2024-02-01T16:59:28.099921Z","shell.execute_reply.started":"2024-02-01T16:59:28.089504Z","shell.execute_reply":"2024-02-01T16:59:28.099100Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def min_max_normalization(x: torch.Tensor) -> torch.Tensor:    \n    \n    x_flatten = x.view(x.size(0), -1)\n    min_in_batch = x_flatten.min(dim=-1, keepdim=True).values.view(x.size(0), 1, 1).to(torch.float16)\n    max_in_batch = x_flatten.max(dim=-1, keepdim=True).values.view(x.size(0), 1, 1).to(torch.float16)\n    \n    return (\n        (x.to(torch.float16) - min_in_batch)\n        / (max_in_batch - min_in_batch + 1e-4)\n    )\n\ndef norm_with_clip(x: torch.Tensor) -> torch.Tensor: \n    \n    dims = list(range(1, x.ndim))\n    mean = x.mean(dim=dims, keepdim=True)\n    std = x.std(dim=dims, keepdim=True)\n    x_r = (\n        (x - mean)\n        / (std + 1e-4)\n    )\n    \n    mask = (x_r>5)\n    x_r[mask]=(x_r[mask]-5) * 1e-3 + 5\n    mask = (x_r<-5)\n    x_r[mask]=(x_r[mask]+5) * 1e-3 - 5\n    \n    return x_r\n\ndef add_noise(x: torch.Tensor, expectation, variance) -> torch.Tensor:\n    \n    return x + torch.normal(expectation, variance, x.shape)","metadata":{"execution":{"iopub.status.busy":"2024-02-01T16:59:28.100983Z","iopub.execute_input":"2024-02-01T16:59:28.101308Z","iopub.status.idle":"2024-02-01T16:59:28.131273Z","shell.execute_reply.started":"2024-02-01T16:59:28.101277Z","shell.execute_reply":"2024-02-01T16:59:28.130446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ShotDataset(Dataset):\n    \n    def __init__(self, paths, is_label):\n        \n        self.paths = sorted(paths)\n        self.is_label = is_label\n    \n    def __len__(self):\n        \n        return len(self.paths)\n    \n    def __getitem__(self, index):\n        \n        img = cv2.imread(self.paths[index], cv2.IMREAD_GRAYSCALE)\n        img = torch.from_numpy(img).to(torch.uint8)\n        \n        if self.is_label:\n            img = (img!=0) * 255\n        \n        return img.to(torch.uint8)\n\n\ndef load_data(paths, is_label):\n    \n    dataset = ShotDataset(paths, is_label)\n    x = [dataset[i] for i in tqdm(range(len(dataset)))]\n    x = torch.stack(x, dim=0)\n    \n    if not is_label:\n        \n        x_flattened = x.view(-1).numpy()\n        index = int(len(x_flattened) * CFG.chopping_percentile)\n        max_value = np.partition(x_flattened, -index)[-index]\n        min_value = np.partition(x_flattened, index)[index]\n        \n        x = torch.clip(x, min_value, max_value)\n        x = min_max_normalization(x)\n        x = (x*255).to(torch.uint8)\n        \n    return x","metadata":{"execution":{"iopub.status.busy":"2024-02-01T16:59:28.132206Z","iopub.execute_input":"2024-02-01T16:59:28.132430Z","iopub.status.idle":"2024-02-01T16:59:28.150571Z","shell.execute_reply.started":"2024-02-01T16:59:28.132410Z","shell.execute_reply":"2024-02-01T16:59:28.149732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class InferenceModel(nn.Module):\n    \n    def __init__(self, CFG, weight=None):\n        super().__init__()\n        \n        self.model = smp.Unet(\n            encoder_name=CFG.backbone, \n            encoder_weights=weight,\n            in_channels=CFG.in_chans,\n            classes=CFG.target_size,\n            activation=None\n        )\n\n    def _forward(self, image):\n        output = self.model(image)\n        return output[:, 0]\n    \n    def forward(self, images, batch_size=CFG.batch_size):\n        \n        images = images.to(torch.float32)\n        images = norm_with_clip(images)\n        \n        if CFG.input_size != CFG.image_size:\n            images = nn.functional.interpolate(\n                images,\n                size=(CFG.input_size,CFG.input_size),\n                mode='bilinear',\n                align_corners=True\n            )\n        \n        original_shape = [images.size(0), images.size(-2), images.size(-1)]\n#         images = [torch.rot90(images, k=i, dims=(-2,-1)) for i in range(4)]\n        images = [torch.rot90(images, k=i, dims=(-2,-1)) for i in [0, 2]]\n        images = torch.cat(images)\n        \n        with autocast():\n            with torch.no_grad():\n                images = [\n                    self._forward(images[i*batch_size:(i+1)*batch_size])\n                    for i in range(images.size(0)//batch_size + 1)\n                ]\n                images = torch.cat(images)\n                \n        images = fn.sigmoid(images)\n        images = images.view(2, *original_shape)\n#         images = [torch.rot90(image, k=-i, dims=(-2,-1)) for i, image in enumerate(images)]\n        images = [torch.rot90(image, k=-i, dims=(-2,-1)) for i, image in zip([0, 2], images)]\n        images = torch.stack(images).mean(dim=0)\n        \n        if CFG.input_size != CFG.image_size:\n            images = nn.functional.interpolate(\n                images,\n                size=(CFG.image_size, CFG.image_size),\n                mode='bilinear',\n                align_corners=True\n            )[0]\n        \n        return images\n\n\ndef build_inference_model(weight=\"imagenet\"):\n    print(f\"Model name - {CFG.model_name}\")\n    print(f\"Backbone   - {CFG.backbone}\")\n\n    model = InferenceModel(CFG, weight)\n    return model.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-02-01T16:59:28.151826Z","iopub.execute_input":"2024-02-01T16:59:28.152543Z","iopub.status.idle":"2024-02-01T16:59:28.167525Z","shell.execute_reply.started":"2024-02-01T16:59:28.152474Z","shell.execute_reply":"2024-02-01T16:59:28.166628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PipelineDataset(Dataset):\n    \n    def __init__(self, data):\n        \n        self.len = data.size(0)\n        self.data = torch.cat(\n            (\n                torch.zeros(CFG.in_chans // 2, *data.shape[1:]),\n                data,\n                torch.zeros(CFG.in_chans // 2, *data.shape[1:]),\n#                 torch.zeros(CFG.in_chans % 2, *data.shape[1:])\n            )\n        )\n        \n    def __len__(self):\n        \n        return self.len\n    \n    def __getitem__(self, index):\n        \n        return self.data[index:index+CFG.in_chans]\n    \n    \ndef get_images_names(location):\n    \n    image_paths = sorted(glob(f\"{location}/images/*\"))\n    \n    def get_image_name(image_path):\n        \n        path = image_path.split(\"/\")[-3:]\n        path.pop(1)\n        name = \"_\".join(path)\n        \n        return name[:-4]\n    \n    labels = []\n    for image_path in image_paths:\n        labels.append(get_image_name(image_path))\n\n    return labels\n\n    \ndef add_border(image, border_size):\n\n    image_mean = int(image.float().mean())\n    image_b = image.clone()\n    \n    image_b = torch.cat(\n        (\n            image_b,\n            torch.full((image_b.size(0), border_size, image_b.size(2)), image_mean, device=image.device)\n        ),\n        dim=1\n    )\n    image_b = torch.cat(\n        (\n            torch.full((image_b.size(0), border_size, image_b.size(2)), image_mean, device=image.device),\n            image_b\n        ),\n        dim=1\n    )\n    image_b = torch.cat(\n        (\n            torch.full((image_b.size(0), image_b.size(1), border_size), image_mean, device=image.device),\n            image_b\n        ),\n        dim=2\n    )\n    image_b = torch.cat(\n        (\n            image_b,\n            torch.full((image_b.size(0), image_b.size(1), border_size), image_mean, device=image.device)\n        ),\n        dim=2\n    )\n    \n    return image_b.to(torch.uint8)\n\n\ndef rle_encode(signal):\n    \n    flattened = signal.flatten()\n    flattened = np.concatenate([[0], flattened, [0]])\n    positions = np.where(flattened[1:] != flattened[:-1])[0] + 1\n    positions[1::2] -= positions[::2]\n    encoded = \"1 0\" if len(positions) == 0 else \" \".join(str(pos) for pos in positions)\n    \n    return encoded","metadata":{"execution":{"iopub.status.busy":"2024-02-01T16:59:28.168718Z","iopub.execute_input":"2024-02-01T16:59:28.169067Z","iopub.status.idle":"2024-02-01T16:59:28.187255Z","shell.execute_reply.started":"2024-02-01T16:59:28.169028Z","shell.execute_reply":"2024-02-01T16:59:28.186310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = build_inference_model()\nmodel.load_state_dict(torch.load(\n    CFG.path_to_model,\n    device\n))\nmodel.eval();\nmodel = DataParallel(model)","metadata":{"execution":{"iopub.status.busy":"2024-02-01T16:59:28.188489Z","iopub.execute_input":"2024-02-01T16:59:28.188820Z","iopub.status.idle":"2024-02-01T16:59:30.554405Z","shell.execute_reply.started":"2024-02-01T16:59:28.188790Z","shell.execute_reply":"2024-02-01T16:59:30.553590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    torch.save(\n        load_data(glob(\"/kaggle/input/blood-vessel-segmentation/train/kidney_2/images/*\"), False),\n        \"data.pt\"\n    )\n    paths = [\"/kaggle/input/blood-vessel-segmentation/train/kidney_2\"]\nelse:\n    paths = glob(\"/kaggle/input/blood-vessel-segmentation/test/*\")\n    \nborder_size = CFG.tile_size // 2\noutput = {\n    \"maps\": [],\n    \"labels\": []\n}","metadata":{"execution":{"iopub.status.busy":"2024-01-31T18:08:13.967737Z","iopub.execute_input":"2024-01-31T18:08:13.968450Z","iopub.status.idle":"2024-01-31T18:11:00.516072Z","shell.execute_reply.started":"2024-01-31T18:08:13.968418Z","shell.execute_reply":"2024-01-31T18:11:00.515031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for path in paths:\n    \n    if DEBUG:\n        data = torch.load('data.pt')\n    else:\n        data = load_data(glob(f\"{path}/images/*\"), False)\n    predicted_maps = torch.zeros_like(data, dtype=torch.uint8)\n    \n    for permutation in range(3):\n        \n        dataset = PipelineDataset(data)\n        x_points = range(0, data.size(-2)+border_size+1, CFG.stride)\n        y_points = range(0, data.size(-1)+border_size+1, CFG.stride)\n\n        for index in tqdm(range(len(dataset)), total=len(dataset)):\n\n            images = dataset[index].clone().to(device)\n            images = torch.unsqueeze(\n                add_border(images, border_size),\n                dim=0\n            ).to(torch.float32)\n\n            points = []\n            activation_maps = []\n            activation_map_cum_sum = torch.zeros(images.size(-2), images.size(-1)).to(device)\n            activation_map_counter = activation_map_cum_sum.clone()\n\n            for x1 in x_points:\n                for y1 in y_points:\n                    x2 = x1 + CFG.tile_size\n                    y2 = y1 + CFG.tile_size\n\n                    if x2 > images.size(-2) or y2 > images.size(-1):\n                        continue\n\n                    points.append(\n                        (\n                            x1+CFG.drop_egde_pixel, x2-CFG.drop_egde_pixel,\n                            y1+CFG.drop_egde_pixel, y2-CFG.drop_egde_pixel\n                        )\n                    )\n                    activation_maps.append(images[..., x1:x2, y1:y2])\n\n            activation_maps = torch.cat(activation_maps)\n            activation_maps = model.forward(activation_maps)\n            activation_maps = activation_maps[\n                ...,\n                CFG.drop_egde_pixel:-CFG.drop_egde_pixel,\n                CFG.drop_egde_pixel:-CFG.drop_egde_pixel\n            ]        \n\n            for i, (x1, x2, y1, y2) in enumerate(points):\n                activation_map_cum_sum[x1:x2, y1:y2] += activation_maps[i]\n                activation_map_counter[x1:x2, y1:y2] += 1.\n\n            activation_map = activation_map_cum_sum / activation_map_counter\n            activation_map = activation_map[border_size:-border_size, border_size:-border_size]\n            predicted_maps[index] += (activation_map * 255 / 3).to(torch.uint8).cpu()\n\n        predicted_maps = predicted_maps.permute(1, 2, 0)\n        data = data.permute(1, 2, 0)\n        \n        del dataset\n        _ = gc.collect()\n    \n    output[\"maps\"].append(predicted_maps)\n    output[\"labels\"].extend(\n        get_images_names(path)\n    )","metadata":{"execution":{"iopub.status.busy":"2024-01-31T18:11:05.427993Z","iopub.execute_input":"2024-01-31T18:11:05.429105Z","iopub.status.idle":"2024-01-31T18:12:25.820760Z","shell.execute_reply.started":"2024-01-31T18:11:05.429067Z","shell.execute_reply":"2024-01-31T18:12:25.817315Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"treshold = [x.flatten().numpy() for x in output[\"maps\"]]\ntreshold = np.concatenate(treshold)\nindex = -int(len(treshold) * CFG.th_percentile)\ntreshold = np.partition(treshold, index)[index]\nsubmission_df = []\n\nfor index in range(len(output[\"labels\"])):\n    labels = output[\"labels\"][index]\n    i = 0\n    \n    while i < len(output[\"maps\"]) and index >= len(output[\"maps\"][i]):\n        index -= len(output[\"maps\"][i])\n        i += 1\n            \n    mask_pred = (output[\"maps\"][i][index] > treshold).numpy()\n    rle = rle_encode(mask_pred)\n    \n    submission_df.append(\n        pd.DataFrame(\n            data={\n                'id': labels,\n                'rle': rle,\n            },\n            index=[0]\n        )\n    )\n\nsubmission_df = pd.concat(submission_df)\nsubmission_df.to_csv('submission.csv', index=False)","metadata":{},"execution_count":null,"outputs":[]}]}