{"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":1807973,"sourceType":"datasetVersion","datasetId":1074109}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Import","metadata":{}},{"cell_type":"code","source":"#certificate verification (needed for segmentatin_models_pytorch library to work)\nimport ssl\nssl._create_default_https_context = ssl._create_unverified_context","metadata":{"execution":{"iopub.status.busy":"2024-04-07T21:10:06.576810Z","iopub.execute_input":"2024-04-07T21:10:06.577701Z","iopub.status.idle":"2024-04-07T21:10:06.587157Z","shell.execute_reply.started":"2024-04-07T21:10:06.577666Z","shell.execute_reply":"2024-04-07T21:10:06.586267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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/\nimport torch as torch \nimport torch.nn as nn  \nimport numpy as np\nfrom tqdm import tqdm\nimport os,sys,cv2\nfrom torch.cuda.amp import autocast\nimport matplotlib.pyplot as plt\nimport albumentations as aug\n!pip install segmentation_models_pytorch\nimport segmentation_models_pytorch as smp\nfrom albumentations.pytorch import ToTensorV2\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.parallel import DataParallel\nfrom glob import glob","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-04-07T21:10:07.022738Z","iopub.execute_input":"2024-04-07T21:10:07.023062Z","iopub.status.idle":"2024-04-07T21:10:32.299124Z","shell.execute_reply.started":"2024-04-07T21:10:07.023037Z","shell.execute_reply":"2024-04-07T21:10:32.298303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"#percentage of augmentation\np_augment = 0.05 \nclass config:\n    #use the SE-ResNeXt backbone \n    backbone = 'se_resnext50_32x4d'\n    # prediction target size\n    target_size = 1\n    in_chans = 1  \n    #img and input sizes\n    image_size = 1024 \n    input_size = 1024 \n\n    train_batch_size = 1 \n    valid_batch_size = 2\n\n    epochs = 40\n    #learning rate\n    lr = 8e-5\n    # fold\n    valid_id = 1\n    # training augmentations\n    train_aug_list = [\n        aug.Rotate(limit=270, p= 0.5),\n        aug.RandomScale(scale_limit=(0.8,1.25),interpolation=cv2.INTER_CUBIC,p=p_augment),\n        aug.RandomCrop(input_size, input_size,p=1),\n        aug.RandomGamma(p=p_augment*2/3),\n        aug.RandomBrightnessContrast(p=p_augment,),\n        aug.GaussianBlur(p=p_augment),\n        aug.MotionBlur(p=p_augment),\n        aug.GridDistortion(num_steps=5, distort_limit=0.3, p=p_augment),\n        ToTensorV2(transpose_mask=True),\n    ]\n    train_aug = aug.Compose(train_aug_list)\n    valid_aug_list = [\n        ToTensorV2(transpose_mask=True),\n    ]\n    valid_aug = aug.Compose(valid_aug_list)","metadata":{"execution":{"iopub.status.busy":"2024-04-07T21:10:32.300713Z","iopub.execute_input":"2024-04-07T21:10:32.301060Z","iopub.status.idle":"2024-04-07T21:10:32.309806Z","shell.execute_reply.started":"2024-04-07T21:10:32.301030Z","shell.execute_reply":"2024-04-07T21:10:32.308874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"#define our model, FPN \nclass CustomModel(nn.Module):\n    def __init__(self, config, weight=None):\n        super().__init__()\n        self.model = smp.FPN(encoder_name='se_resnext50_32x4d',in_channels=config.in_chans,classes=config.target_size,activation=None,)\n    def forward(self, image):\n        output = self.model(image)\n        # output = output.squeeze(-1)\n        return output[:,0]#.sigmoid()\n\ndef build_model(weight=\"imagenet\"):\n    #load_dotenv()\n    #print('model_name', config.model_name)\n    print('backbone', config.backbone)\n    model = CustomModel(config, weight)\n    return model.cuda()","metadata":{"execution":{"iopub.status.busy":"2024-04-07T21:10:32.311140Z","iopub.execute_input":"2024-04-07T21:10:32.311493Z","iopub.status.idle":"2024-04-07T21:10:32.323427Z","shell.execute_reply.started":"2024-04-07T21:10:32.311463Z","shell.execute_reply":"2024-04-07T21:10:32.322705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Resizing","metadata":{}},{"cell_type":"code","source":"def resize(img, image_size=1024):\n    height, width = img.shape[:2]\n    \n    # if the image dimensions already equal/larger than the target size, return the original image\n    if height >= image_size and width >= image_size:\n        return img\n    \n    #padding sizes\n    pad_height = (image_size - height) // 2 if height < image_size else 0\n    pad_width = (image_size - width) // 2 if width < image_size else 0\n    \n    # add padding to the image to match target size\n    img_result = np.pad(img, ((pad_height, pad_height), (pad_width, pad_width)), 'constant', constant_values=0)\n    \n    # for odd dimensions trim extra pixel from the bottom/right side \n    img_result = img_result[:image_size, :image_size]\n    \n    return img_result","metadata":{"execution":{"iopub.status.busy":"2024-04-07T21:10:32.326162Z","iopub.execute_input":"2024-04-07T21:10:32.326454Z","iopub.status.idle":"2024-04-07T21:10:32.338674Z","shell.execute_reply.started":"2024-04-07T21:10:32.326428Z","shell.execute_reply":"2024-04-07T21:10:32.337828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions","metadata":{}},{"cell_type":"code","source":"def min_max_normalization(x):\n    \"\"\"Normalize a tensor to [0, 1] across the last dimension.\n\n    Args:\n        x (torch.Tensor): Input tensor with shape (batch, f1, ...).\n\n    Returns:\n        torch.Tensor: Normalized tensor.\n    \"\"\"\n    original_shape = x.shape  #original shape for reshaping later\n    if x.ndim > 2:\n        x = x.reshape(x.shape[0], -1)  # flatten the tensor while keeping the batch dimension\n    \n    minval = x.min(dim=-1, keepdim=True).values\n    maxval = x.max(dim=-1, keepdim=True).values\n\n    #prevent division by zero via small epsilon and normalize\n    x = (x - minval) / (maxval - minval + 1e-9)\n    return x.reshape(original_shape)\n\ndef norm_and_clip(x, smooth = 1e-5):\n    \"\"\"Standardize a tensor to have mean=0 and std=1, with clipping.\n\n    Args:\n        x (torch.Tensor): Input tensor.\n        smooth (float): Smoothing term to avoid division by zero.\n\n    Returns:\n        torch.Tensor: Standardized tensor.\n    \"\"\"\n    dim = list(range(1, x.ndim))  # dimensions over which to compute the mean and std\n    mean = x.mean(dim=dim, keepdim=True)\n    std = x.std(dim=dim, keepdim=True)\n\n    return (x - mean) / (std + smooth)\n\n#define data loader class\nclass Data_loader(Dataset):\n     \n    def __init__(self,paths,is_label):\n        self.paths=paths\n        self.paths.sort()\n        self.is_label=is_label\n    \n    def __len__(self):\n        return len(self.paths)\n    \n    def __getitem__(self,index):\n         \n        img = cv2.imread(self.paths[index],cv2.IMREAD_GRAYSCALE)\n        \n        img = resize(img , image_size = config.image_size )\n\n        img = torch.from_numpy(img.copy())\n        if self.is_label:\n            img=(img!=0).to(torch.uint8)*255\n        else:\n            img=img.to(torch.uint8)\n        return img\n\ndef load_data(paths,is_label=False):\n    data_loader=Data_loader(paths,is_label)\n    data_loader=DataLoader(data_loader, batch_size=16, num_workers=2)  \n    data=[]\n    for x in tqdm(data_loader):\n        data.append(x)\n    x=torch.cat(data,dim=0)\n    del data\n    if not is_label:\n        x=(min_max_normalization(x.to(torch.float16)[None])[0]*255).to(torch.uint8)\n    return x\n\ndef dice_coef(y_pred, y_true, thr=0.5, dim=(-1,-2), epsilon=0.001):\n    y_pred=y_pred.sigmoid()\n    y_true = y_true.to(torch.float32)\n    y_pred = (y_pred>thr).to(torch.float32)\n    inter = (y_true*y_pred).sum(dim=dim)\n    den = y_true.sum(dim=dim) + y_pred.sum(dim=dim)\n    dice = ((2*inter+epsilon)/(den+epsilon)).mean()\n    return dice\n\nclass DiceLoss(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(DiceLoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        \n        #comment out if your model contains a sigmoid or equivalent activation layer\n        inputs = inputs.sigmoid()   \n        \n        #flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()                            \n        dice = (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n        \n        return 1 - dice\n    \nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=0.25, gamma=2.0):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n\n    def forward(self, inputs, targets):\n        BCE_loss = nn.functional.binary_cross_entropy_with_logits(inputs, targets, reduction='none')\n        pt = torch.exp(-BCE_loss)  # pt is the probability of being classified as the target class\n        F_loss = self.alpha * (1-pt)**self.gamma * BCE_loss\n        return F_loss.mean()\n\nclass Custom_Dataset(Dataset):\n    def __init__(self,x,y,arg=False):\n        super(Dataset,self).__init__()\n        self.x=x\n        self.y=y\n        self.image_size=config.image_size\n        self.in_chans=config.in_chans\n        self.arg=arg\n        if arg:\n            self.transform=config.train_aug\n        else: \n            self.transform=config.valid_aug\n\n    def __len__(self) -> int:\n        # calculate the total number of samples\n        return sum([y.shape[0] - self.in_chans for y in self.y])\n    \n    def __getitem__(self,index):\n        # Determine which set of x,y the current index falls into\n        i = 0\n        for x in self.x:\n            if index > x.shape[0] - self.in_chans:\n                index -= x.shape[0] - self.in_chans\n                i += 1\n            else:\n                break\n        x = self.x[i]\n        y = self.y[i]\n        \n        # central crop calculation \n        x_index = (x.shape[1] - self.image_size) // 2\n        y_index = (x.shape[2] - self.image_size) // 2\n        \n        # taking the relevant crop from the image and the mask\n        x = x[index:index + self.in_chans, x_index:x_index + self.image_size, y_index:y_index + self.image_size]\n        y = y[index + self.in_chans // 2, x_index:x_index + self.image_size, y_index:y_index + self.image_size]\n\n        # apply transfs\n        data = self.transform(image=x.numpy().transpose(1,2,0), mask=y.numpy())\n        x = data['image']\n        y = data['mask'] >= 127  # binarize mask based on threshold\n        \n        # apply augmentation\n        if self.arg:\n            i = np.random.randint(4)\n            x = x.rot90(i, dims=(1,2))\n            y = y.rot90(i, dims=(0,1))\n            for i in range(3):\n                if np.random.randint(2):\n                    x = x.flip(dims=(i,))\n                    if i >= 1:\n                        y = y.flip(dims=(i-1,))\n        return x, y","metadata":{"execution":{"iopub.status.busy":"2024-04-07T21:10:32.340157Z","iopub.execute_input":"2024-04-07T21:10:32.340480Z","iopub.status.idle":"2024-04-07T21:10:32.370811Z","shell.execute_reply.started":"2024-04-07T21:10:32.340450Z","shell.execute_reply":"2024-04-07T21:10:32.369895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load data ","metadata":{}},{"cell_type":"code","source":"train_x=[]\ntrain_y=[]\n\nroot_path=\"/kaggle/input/blood-vessel-segmentation/\"\npaths=[\"/kaggle/input/blood-vessel-segmentation/train/kidney_1_dense\"]\nfor i,path in enumerate(paths):\n    if path==\"/kaggle/input/blood-vessel-segmentation/train/kidney_3_dense\":\n        continue\n    x=load_data(glob(f\"{path}/images/*\"),is_label=False)\n    print(x.shape)\n    y=load_data(glob(f\"{path}/labels/*\"),is_label=True)\n    print(y.shape)\n    train_x.append(x)\n    train_y.append(y)\n\n    #(C,H,W)\n    #aug\n    train_x.append(x.permute(1,2,0))\n    train_y.append(y.permute(1,2,0))\n    train_x.append(x.permute(2,0,1))\n    train_y.append(y.permute(2,0,1))\npath2=\"/kaggle/input/blood-vessel-segmentation/train/kidney_3_dense\"\npaths_y=glob(f\"{path2}/labels/*\")\npaths_x=[x.replace(\"labels\",\"images\").replace(\"dense\",\"sparse\") for x in paths_y]\n\nval_x=load_data(paths_x,is_label=False)\nprint(val_x.shape)\nval_y=load_data(paths_y,is_label=True)\nprint(val_y.shape)","metadata":{"execution":{"iopub.status.busy":"2024-04-07T21:10:32.372061Z","iopub.execute_input":"2024-04-07T21:10:32.372333Z","iopub.status.idle":"2024-04-07T21:13:18.345857Z","shell.execute_reply.started":"2024-04-07T21:10:32.372307Z","shell.execute_reply":"2024-04-07T21:13:18.344741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"torch.backends.cudnn.enabled = True\ntorch.backends.cudnn.benchmark = True\n\n# initialize datasets and data loaders\ntrain_dataset = Custom_Dataset(train_x, train_y, arg=True)\ntrain_loader = DataLoader(train_dataset, batch_size=config.train_batch_size, num_workers=2, shuffle=True, pin_memory=True)\nval_dataset = Custom_Dataset([val_x], [val_y])\nval_loader = DataLoader(val_dataset, batch_size=config.valid_batch_size, num_workers=2, shuffle=False, pin_memory=True)\n\n# model, loss fn, optimizer, and scheduler setup\nmodel = DataParallel(build_model()).cuda()\nloss_function = DiceLoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=config.lr)\nscaler = torch.cuda.amp.GradScaler()\nscheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=config.lr, steps_per_epoch=len(train_loader), epochs=config.epochs + 1, pct_start=0.1)\n\n#training routine for an epoch\ndef train_epoch(loader, model, loss_function, optimizer, scaler, scheduler, is_train=True):\n    if is_train:\n        model.train()\n    else:\n        model.eval()\n\n    losses = []\n    scores = []\n\n    progress_bar = tqdm(loader, desc=\"processing\", leave=True)\n    for x, y in progress_bar:\n        x, y = x.cuda().float(), y.cuda().float()\n        x = norm_and_clip(x.reshape(-1, *x.shape[2:])).reshape(x.shape)\n\n        with autocast():\n            pred = model(x)\n            loss = loss_function(pred, y)\n\n        if is_train:\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            scheduler.step()\n\n        score = dice_coef(pred.detach(), y)\n        losses.append(loss.item())\n        scores.append(score.item())\n\n        progress_bar.set_description(f\"Loss: {np.mean(losses):.4f}, Score: {np.mean(scores):.4f}\")\n    \n    return np.mean(losses), np.mean(scores)\n#train for all epochs\ndef train1():\n    for epoch in range(2):\n        train_loss, train_score = train_epoch(train_loader, model, loss_function, optimizer, scaler, scheduler, is_train=True)\n        val_loss, val_score = train_epoch(val_loader, model, loss_function, optimizer, scaler, scheduler, is_train=False)\n        \n        print(f\"Epoch {epoch}: Train Loss {train_loss:.4f}, Train Score {train_score:.4f}, Val Loss {val_loss:.4f}, Val Score {val_score:.4f}\")\n\n    torch.save(model.module.state_dict(), f\"./{config.backbone}_{epoch}_train_loss{train_loss:.2f}_train_score{train_score:.2f}_val_loss{val_loss:.2f}_val_score{val_score:.2f}.pt\")","metadata":{"execution":{"iopub.status.busy":"2024-04-07T21:26:36.239461Z","iopub.execute_input":"2024-04-07T21:26:36.240257Z","iopub.status.idle":"2024-04-07T21:26:36.687811Z","shell.execute_reply.started":"2024-04-07T21:26:36.240225Z","shell.execute_reply":"2024-04-07T21:26:36.687053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.backends.cudnn.enabled = True\ntorch.backends.cudnn.benchmark = True\n    \ntrain_dataset=Custom_Dataset(train_x,train_y,arg=True)\ntrain_dataset = DataLoader(train_dataset, batch_size=config.train_batch_size ,num_workers=2, shuffle=True, pin_memory=True)\nval_dataset=Custom_Dataset([val_x],[val_y])\nval_dataset = DataLoader(val_dataset, batch_size=config.valid_batch_size, num_workers=2, shuffle=False, pin_memory=True)\n\nmodel=build_model()\nmodel=DataParallel(model)\n\nloss_fc=DiceLoss()\n#use AdamW optimizer which is extension of Adam optimizer with weight decay, better for generalization \noptimizer=torch.optim.AdamW(model.parameters(),lr=config.lr)\nscaler=torch.cuda.amp.GradScaler()\n#one cycle learning rate schedule. Cyclical lr helps model improve accuracy in fewer steps\nscheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=config.lr,\n                                                steps_per_epoch=len(train_dataset), epochs=config.epochs+1,\n                                                pct_start=0.1,)\ndef train():\n    for epoch in range(config.epochs):\n        model.train()\n        time=tqdm(range(len(train_dataset)))\n        #init losses and scores\n        losses=0\n        scores=0\n        #iterate through train set \n        for i,(x,y) in enumerate(train_dataset):\n            x=x.cuda().to(torch.float32)\n            y=y.cuda().to(torch.float32)\n            #norm and clip (prevent NaN gradients)\n            x=norm_and_clip(x.reshape(-1,*x.shape[2:])).reshape(x.shape)\n            with autocast():\n                pred=model(x)\n                loss=loss_fc(pred,y)\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            scheduler.step()\n            score=dice_coef(pred.detach(),y)\n            losses=(losses*i+loss.item())/(i+1)\n            scores=(scores*i+score)/(i+1)\n            time.set_description(f\"epoch number:{epoch},loss:{losses:.4f},score:{scores:.4f},lr{optimizer.param_groups[0]['lr']:.4e}\")\n            time.update()\n            del loss,pred\n        time.close()\n\n        model.eval()\n        time=tqdm(range(len(val_dataset)))\n        val_losses=0\n        val_scores=0\n        for i,(x,y) in enumerate(val_dataset):\n            x=x.cuda().to(torch.float32)\n            y=y.cuda().to(torch.float32)\n            x=norm_and_clip(x.reshape(-1,*x.shape[2:])).reshape(x.shape)\n\n            with autocast():\n                with torch.no_grad():\n                    pred=model(x)\n                    loss=loss_fc(pred,y)\n            score=dice_coef(pred.detach(),y)\n            val_losses=(val_losses*i+loss.item())/(i+1)\n            val_scores=(val_scores*i+score)/(i+1)\n            time.set_description(f\"val-->loss:{val_losses:.4f},score:{val_scores:.4f}\")\n            time.update()\n\n        time.close()\n    torch.save(model.module.state_dict(), f\"./FPN_epochs{config.epochs}_val_acc{val_scores:.2f}.pt\")\n\n    time.close()","metadata":{"execution":{"iopub.status.busy":"2024-04-07T21:26:38.674154Z","iopub.execute_input":"2024-04-07T21:26:38.674927Z","iopub.status.idle":"2024-04-07T21:26:39.040881Z","shell.execute_reply.started":"2024-04-07T21:26:38.674887Z","shell.execute_reply":"2024-04-07T21:26:39.040132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#visualize\nimport matplotlib.pyplot as plt\n\n# Access the first element from your DataLoader\nsample_image, sample_label = next(iter(train_dataset))\n\n# Plot the image and its corresponding label\nplt.figure(figsize=(10, 5))\n\n# Plot the image\nplt.subplot(1, 2, 1)\nplt.imshow(sample_image.squeeze(), cmap='gray')  # Assuming the image is grayscale\nplt.title('Sample Image')\nplt.axis('off')\n\n# Plot the label\nplt.subplot(1, 2, 2)\nplt.imshow(sample_label.squeeze(), cmap='gray')  # Assuming the label is grayscale\nplt.title('Sample Label')\nplt.axis('off')\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-07T21:18:31.278750Z","iopub.execute_input":"2024-04-07T21:18:31.279143Z","iopub.status.idle":"2024-04-07T21:18:32.237162Z","shell.execute_reply.started":"2024-04-07T21:18:31.279117Z","shell.execute_reply":"2024-04-07T21:18:32.236141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train1()","metadata":{"execution":{"iopub.status.busy":"2024-04-07T21:26:47.035263Z","iopub.execute_input":"2024-04-07T21:26:47.036208Z","iopub.status.idle":"2024-04-07T21:39:46.925175Z","shell.execute_reply.started":"2024-04-07T21:26:47.036173Z","shell.execute_reply":"2024-04-07T21:39:46.923985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}