{"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":"none","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":1807973,"sourceType":"datasetVersion","datasetId":1074109},{"sourceId":7532567,"sourceType":"datasetVersion","datasetId":4363835},{"sourceId":150248402,"sourceType":"kernelVersion"}],"dockerImageVersionId":30648,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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-02T16:41:54.056213Z","iopub.execute_input":"2024-02-02T16:41:54.056679Z","iopub.status.idle":"2024-02-02T16:42:28.172449Z","shell.execute_reply.started":"2024-02-02T16:41:54.056641Z","shell.execute_reply":"2024-02-02T16:42:28.170991Z"},"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 numpy as np\nimport datetime\nimport ctypes\nimport torch\nimport gc\n\nlibc = ctypes.CDLL(\"libc.so.6\")\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"execution":{"iopub.status.busy":"2024-02-02T16:42:28.175286Z","iopub.execute_input":"2024-02-02T16:42:28.176090Z","iopub.status.idle":"2024-02-02T16:42:40.029905Z","shell.execute_reply.started":"2024-02-02T16:42:28.176044Z","shell.execute_reply":"2024-02-02T16:42:40.028555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    \n    model_name = 'Unet'\n    backbone = 'se_resnext50_32x4d'\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 = 16\n    eval_batch_size = batch_size * 2\n\n    epochs = 50\n    lr = 5e-5\n    chopping_percentile = 1e-3\n\n    train_aug = alb.Compose(\n        [\n            alb.HorizontalFlip(p=0.5),\n#             alb.ChannelDropout(channel_drop_range=(1, in_chans // 2), p=0.5),\n            alb.GaussNoise(var_limit=10, p=0.9),\n            alb.Rotate(limit=15, p=0.5),\n            alb.RandomScale(scale_limit=(0.8, 1.2), interpolation=cv2.INTER_CUBIC, p=0.4),\n            alb.RandomCrop(input_size, input_size, p=1),\n            alb.RandomBrightnessContrast(p=0.6),\n            alb.GaussianBlur(p=0.5),\n#             alb.MotionBlur(p=0.5),\n            alb.GridDistortion(num_steps=4, distort_limit=0.2, p=0.4),\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-02T16:42:40.031584Z","iopub.execute_input":"2024-02-02T16:42:40.032100Z","iopub.status.idle":"2024-02-02T16:42:40.044373Z","shell.execute_reply.started":"2024-02-02T16:42:40.032059Z","shell.execute_reply":"2024-02-02T16:42:40.042635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainModel(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\ndef build_train_model(weight=\"imagenet\"):\n    print(f\"Model name - {CFG.model_name}\")\n    print(f\"Backbone   - {CFG.backbone}\")\n\n    model = TrainModel(CFG, weight)\n    return model.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-02-02T16:42:40.047533Z","iopub.execute_input":"2024-02-02T16:42:40.048068Z","iopub.status.idle":"2024-02-02T16:42:40.062364Z","shell.execute_reply.started":"2024-02-02T16:42:40.048024Z","shell.execute_reply":"2024-02-02T16:42:40.061195Z"},"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\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\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-02T16:42:40.063943Z","iopub.execute_input":"2024-02-02T16:42:40.064366Z","iopub.status.idle":"2024-02-02T16:42:40.077300Z","shell.execute_reply.started":"2024-02-02T16:42:40.064335Z","shell.execute_reply":"2024-02-02T16:42:40.075934Z"},"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)","metadata":{"execution":{"iopub.status.busy":"2024-02-02T16:42:40.079116Z","iopub.execute_input":"2024-02-02T16:42:40.079649Z","iopub.status.idle":"2024-02-02T16:42:40.093610Z","shell.execute_reply.started":"2024-02-02T16:42:40.079619Z","shell.execute_reply":"2024-02-02T16:42:40.092261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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\n\n\ndef dice_coef(target, prediction, th=0.5, dim=(-2, -1)):\n    \n    prediction_r = fn.sigmoid(prediction) > th\n    target_r = target.to(torch.float32)\n    prediction_r = prediction_r.to(torch.float32)\n    \n    inter = (target_r * prediction_r).sum(dim=dim)\n    union = target_r.sum(dim=dim) + prediction_r.sum(dim=dim)\n    coef = ((2*inter + 1e-4) / (union + 1e-4)).mean()\n    \n    return coef\n\n\ndef add_border(image, border_size, zeros=False, axis=(0, 1)):\n\n    image_mean = int(image.float().mean()) if not zeros else 0.\n    image_b = image.clone()\n    \n    if 0 in axis:\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, dtype=image.dtype)\n            ),\n            dim=1\n        )\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, dtype=image.dtype),\n                image_b\n            ),\n            dim=1\n        )\n        \n    if 1 in axis:\n        image_b = torch.cat(\n            (\n                torch.full((image_b.size(0), image_b.size(1), border_size), image_mean, device=image.device, dtype=image.dtype),\n                image_b\n            ),\n            dim=2\n        )\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, dtype=image.dtype)\n            ),\n            dim=2\n        )\n    \n    return image_b.to(torch.uint8)","metadata":{"execution":{"iopub.status.busy":"2024-02-02T16:42:40.095031Z","iopub.execute_input":"2024-02-02T16:42:40.095518Z","iopub.status.idle":"2024-02-02T16:42:40.115447Z","shell.execute_reply.started":"2024-02-02T16:42:40.095476Z","shell.execute_reply":"2024-02-02T16:42:40.114248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LOAD = False","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not LOAD:\n    x_train = []\n    y_train = []\n\n    train_paths = [\"/kaggle/input/blood-vessel-segmentation/train/kidney_1_dense\"]\n    for path in train_paths:\n\n        border_size = CFG.image_size // 16\n        x = load_data(glob(f\"{path}/images/*\"), is_label=False)\n        x_train.extend(\n            (\n                x,\n    #             x.permute(0, 2, 1),\n    #             torch.flip(x, dims=(-1,)),\n                x.permute(1, 2, 0),\n    #             x.permute(1, 0, 2),\n    #             torch.flip(x.permute(1, 2, 0), dims=(-1,)),\n                x.permute(2, 0, 1),\n    #             x.permute(2, 1, 0),\n    #             torch.flip(x.permute(2, 0, 1), dims=(-1,))\n            )\n        )\n\n        y = load_data(glob(f\"{path}/labels/*\"), is_label=True)\n    #     y = add_border(y, border_size, zeros=True, axis=(1,))\n        y_train.extend(\n            (\n                y,\n    #             y.permute(0, 2, 1),\n    #             torch.flip(y, dims=(-1,)),\n                y.permute(1, 2, 0),\n    #             y.permute(1, 0, 2),\n    #             torch.flip(y.permute(1, 2, 0), dims=(-1,)),\n                y.permute(2, 0, 1),\n    #             y.permute(2, 1, 0),\n    #             torch.flip(y.permute(2, 0, 1), dims=(-1,))\n            )\n        )\n\n    torch.save(x_train[0], \"kidney_1x_012.pt\")\n    torch.save(x_train[1], \"kidney_1x_120.pt\")\n    torch.save(x_train[2], \"kidney_1x_201.pt\")\n\n    torch.save(y_train[0], \"kidney_1y_012.pt\")\n    torch.save(y_train[1], \"kidney_1y_120.pt\")\n    torch.save(y_train[2], \"kidney_1y_201.pt\")\nelse:\n    x_train, y_train = [None for _ in range(3)], [None for _ in range(3)]\n\n    x_train[0] = torch.load(\"kidney_1x_012.pt\")\n    x_train[1] = torch.load(\"kidney_1x_120.pt\")\n    x_train[2] = torch.load(\"kidney_1x_201.pt\")\n\n    y_train[0] = torch.load(\"kidney_1y_012.pt\")\n    y_train[1] = torch.load(\"kidney_1y_120.pt\")\n    y_train[2] = torch.load(\"kidney_1y_201.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-02-02T16:42:40.117422Z","iopub.execute_input":"2024-02-02T16:42:40.118009Z","iopub.status.idle":"2024-02-02T16:47:09.229179Z","shell.execute_reply.started":"2024-02-02T16:42:40.117978Z","shell.execute_reply":"2024-02-02T16:47:09.227470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# x_train[0] = add_border(x_train[0], CFG.image_size // 8, axis=(1,))\n# x_train[1] = add_border(x_train[1], CFG.image_size // 8, axis=(0,))\n# y_train[0] = add_border(y_train[0], CFG.image_size // 8, axis=(1,))\n# y_train[1] = add_border(y_train[1], CFG.image_size // 8, axis=(0,))\n\n# gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not LOAD:\n    eval_path = \"/kaggle/input/blood-vessel-segmentation/train/kidney_3_dense\"\n    eval_paths_y = glob(f\"{eval_path}/labels/*\")\n    eval_paths_x = [x.replace(\"labels\", \"images\").replace(\"dense\", \"sparse\") for x in eval_paths_y]\n    x_eval, y_eval = [], []\n\n    x = load_data(eval_paths_x, is_label=False)\n    y = load_data(eval_paths_y, is_label=True)\n\n    torch.save(x, \"kidney_3x_012.pt\")\n    torch.save(y, \"kidney_3y_012.pt\")\nelse:\n    x_eval, y_eval = [None for _ in range(1)], [None for _ in range(1)]\n\n    x_eval[0] = torch.load(\"kidney_3x_012.pt\")\n    y_eval[0] = torch.load(\"kidney_3y_012.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-02-02T16:47:09.233540Z","iopub.execute_input":"2024-02-02T16:47:09.235075Z","iopub.status.idle":"2024-02-02T16:48:37.042069Z","shell.execute_reply.started":"2024-02-02T16:47:09.235030Z","shell.execute_reply":"2024-02-02T16:48:37.040842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# x_eval[0] = add_border(x_eval[0], CFG.image_size // 8, axis=(1,))\n# y_eval[0] = add_border(y_eval[0], CFG.image_size // 8, axis=(1,))\n\n# gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DiceLoss(nn.Module):\n    \n    def __init__(self):\n        super(DiceLoss, self).__init__()\n\n    def forward(self, targets, inputs, smooth=1e-3):\n        \n        inputs = fn.sigmoid(inputs)   \n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()                            \n        union = inputs.sum() + targets.sum()\n        dice = (2*intersection + smooth) / (union + smooth)  \n        \n        return 1 - dice\n    \n\nclass Kaggld_Dataset(Dataset):\n    \n    def __init__(self, x, y, train=False):\n        super(Dataset,self).__init__()\n        \n        self.x = x #list[(C,H,W),...]\n        self.y = y #list[(C,H,W),...]\n        self.image_size = CFG.image_size\n        self.in_chans = CFG.in_chans\n        self.train = train\n        self.transform = CFG.train_aug if train else CFG.eval_aug\n        \n    def __len__(self):\n        \n        return sum([len(y) for y in self.y])\n    \n    def __getitem__(self, index):\n        \n        for x, y in zip(self.x, self.y):\n            if index >= len(x):\n                index -= len(x)\n            else:\n                break\n        \n        part_size = CFG.in_chans // 2\n        x_index = index\n        y_index = index\n        x_shift = np.random.randint(0, x.size(1)-self.image_size)\n        y_shift = np.random.randint(0, x.size(2)-self.image_size)\n        \n        x = x[..., x_shift:x_shift+self.image_size, y_shift:y_shift+self.image_size]\n        y = y[..., x_shift:x_shift+self.image_size, y_shift:y_shift+self.image_size]\n        \n        if index < part_size:\n            x = torch.cat(\n                (\n                    torch.zeros(part_size, *x.shape[1:]),\n                    x\n                )\n            )\n            x_index += part_size\n            \n        if index >= len(y) - 1 - part_size:\n            x = torch.cat(\n                (\n                    x,\n                    torch.zeros(part_size, *x.shape[1:])\n                )\n            )\n            \n        x = x[x_index]\n        y = y[y_index]\n        \n        data = self.transform(image=x.numpy().transpose(1, 2, 0), mask=y.numpy())\n        x = data['image'].to(torch.uint8)\n        y = (data['mask'] >= 127).to(torch.uint8)\n                        \n        return x, y","metadata":{"execution":{"iopub.status.busy":"2024-02-02T12:17:45.978266Z","iopub.status.idle":"2024-02-02T12:17:45.978616Z","shell.execute_reply.started":"2024-02-02T12:17:45.978430Z","shell.execute_reply":"2024-02-02T12:17:45.978444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torch.backends.cudnn.enabled = True\n# torch.backends.cudnn.benchmark = True\n    \n# train_dataset = Kaggld_Dataset(x_train, y_train, train=True)\n# eval_dataset = Kaggld_Dataset(x_eval, y_eval)\n\n# train_dataloader = DataLoader(train_dataset, batch_size=CFG.batch_size, num_workers=2, shuffle=True, pin_memory=True)\n# eval_dataloader = DataLoader(eval_dataset, batch_size=CFG.eval_batch_size, num_workers=2, shuffle=False, pin_memory=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = build_train_model()\n# model = DataParallel(model)","metadata":{"execution":{"iopub.status.busy":"2024-02-02T11:22:46.395618Z","iopub.execute_input":"2024-02-02T11:22:46.395947Z","iopub.status.idle":"2024-02-02T11:22:47.723172Z","shell.execute_reply.started":"2024-02-02T11:22:46.395922Z","shell.execute_reply":"2024-02-02T11:22:47.722208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# criterion = DiceLoss()\n# optimizer = optim.AdamW(\n#     model.parameters(),\n#     lr=CFG.lr\n# )\n# scaler = torch.cuda.amp.GradScaler()\n# scheduler = optim.lr_scheduler.OneCycleLR(\n#     optimizer,\n#     max_lr=CFG.lr,\n#     steps_per_epoch=len(train_dataloader),\n#     epochs=CFG.epochs+1,\n#     pct_start=0.1\n# )\n\n# for epoch in range(CFG.epochs):\n    \n#     model.train()\n#     log = tqdm(range(len(train_dataloader)))\n#     mean_loss_on_train = 0\n#     mean_dice_on_train = 0\n    \n#     for i, (train, target) in enumerate(train_dataloader):\n#         optimizer.zero_grad()\n        \n#         train = train.to(torch.float32)\n#         target = target.to(device).to(torch.float32)\n#         train = norm_with_clip(train)\n#         train = train.to(device)\n        \n#         with autocast():\n#             prediction = model.forward(train)\n#             loss = criterion(target, prediction)\n            \n#         scaler.scale(loss).backward()\n#         scaler.step(optimizer)\n#         scaler.update()\n#         scheduler.step()\n        \n#         score = dice_coef(target, prediction.detach())\n#         mean_loss_on_train = (mean_loss_on_train*i + loss.item()) / (i+1)\n#         mean_dice_on_train = (mean_dice_on_train*i + score) / (i+1)\n#         log.set_description(f\"train | epoch:{epoch}; loss:{loss.item():.3f}, dice:{score:.3f}; mean_loss:{mean_loss_on_train:.3f}, mean_dice:{mean_dice_on_train:.3f}\")\n#         log.update()\n    \n#     log.close()\n#     model.eval()\n#     log = tqdm(range(len(eval_dataloader)))\n#     mean_loss_on_eval = 0\n#     mean_dice_on_eval = 0\n    \n#     for i, (eval_, target) in enumerate(eval_dataloader):\n        \n#         eval_ = eval_.to(device).to(torch.float32)\n#         target = target.to(device).to(torch.float32)\n#         eval_ = norm_with_clip(eval_)\n        \n#         with autocast():\n#             with torch.no_grad():\n#                 prediction = model(eval_)\n#                 loss = criterion(target, prediction)\n                \n#         score = dice_coef(target, prediction.detach())\n#         mean_loss_on_train = (mean_loss_on_train*i + loss.item()) / (i+1)\n#         mean_dice_on_train = (mean_dice_on_train*i + score) / (i+1)\n#         log.set_description(f\"eval  | mean_loss:{mean_loss_on_train:.3f}, mean_dice:{mean_dice_on_train:.3f}\")\n#         log.update()\n        \n#     torch.save(model.module.state_dict(), f\"./{CFG.backbone}_{epoch}epoch_{CFG.image_size}image_size_{datetime.datetime.now():%Y_%m_%d_%H_%M}.pt\")\n#     log.close()","metadata":{"execution":{"iopub.status.busy":"2024-02-02T11:22:48.671866Z","iopub.execute_input":"2024-02-02T11:22:48.672388Z","iopub.status.idle":"2024-02-02T11:22:49.591529Z","shell.execute_reply.started":"2024-02-02T11:22:48.672354Z","shell.execute_reply":"2024-02-02T11:22:49.589885Z"},"trusted":true},"execution_count":null,"outputs":[]}]}