{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-07-27T10:14:19.519377Z","iopub.execute_input":"2023-07-27T10:14:19.519755Z","iopub.status.idle":"2023-07-27T10:14:19.538520Z","shell.execute_reply.started":"2023-07-27T10:14:19.519725Z","shell.execute_reply":"2023-07-27T10:14:19.537542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import Libaries","metadata":{}},{"cell_type":"code","source":"!pip install segmentation-models-pytorch\nimport segmentation_models_pytorch as smp\n","metadata":{"scrolled":true,"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-07-27T10:14:20.199858Z","iopub.execute_input":"2023-07-27T10:14:20.200886Z","iopub.status.idle":"2023-07-27T10:14:42.868043Z","shell.execute_reply.started":"2023-07-27T10:14:20.200843Z","shell.execute_reply":"2023-07-27T10:14:42.867036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install adabound\nfrom adabound import AdaBound","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:14:42.870377Z","iopub.execute_input":"2023-07-27T10:14:42.870827Z","iopub.status.idle":"2023-07-27T10:14:54.693750Z","shell.execute_reply.started":"2023-07-27T10:14:42.870789Z","shell.execute_reply":"2023-07-27T10:14:54.692602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport random\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision\nimport torchvision.utils as vutils\nfrom torchvision import models\nimport warnings\nwarnings.filterwarnings('ignore')\n\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau, CosineAnnealingLR, StepLR, MultiStepLR, CyclicLR\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms as T, datasets as dset\nfrom sklearn.model_selection import train_test_split\nfrom matplotlib import pyplot as plt\n\nfrom zipfile import ZipFile\nfrom tqdm import tqdm\nfrom glob import glob\nfrom PIL import Image","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:16:27.241599Z","iopub.execute_input":"2023-07-27T10:16:27.241984Z","iopub.status.idle":"2023-07-27T10:16:27.250062Z","shell.execute_reply.started":"2023-07-27T10:16:27.241952Z","shell.execute_reply":"2023-07-27T10:16:27.248952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport cv2\nimport zipfile\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torchvision\nimport torch.nn as nn\nimport albumentations as A\nimport torch.optim as optim\nfrom torchvision import models\nimport torch.nn.functional as F\nfrom torch.optim import lr_scheduler\nimport torchvision.datasets as datasets\nimport torchvision.transforms as transforms\nfrom torchvision.datasets import ImageFolder\nfrom albumentations.pytorch import ToTensorV2 \nfrom torch.utils.data import DataLoader, Dataset\nimport torchvision.utils as vutils","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:14:55.304706Z","iopub.execute_input":"2023-07-27T10:14:55.305091Z","iopub.status.idle":"2023-07-27T10:14:56.535081Z","shell.execute_reply.started":"2023-07-27T10:14:55.305058Z","shell.execute_reply":"2023-07-27T10:14:56.534109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"batch_size = 4\nn_iters = 100\nepochs = 50\nlearning_rate = 0.0001\nn_workers = 2\n\n\nwidth = 256\nheight = 256\nchannels = 3\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nseed = 44\nrandom.seed(seed)\ntorch.manual_seed(seed)","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:14:56.536297Z","iopub.execute_input":"2023-07-27T10:14:56.537208Z","iopub.status.idle":"2023-07-27T10:14:56.577335Z","shell.execute_reply.started":"2023-07-27T10:14:56.537180Z","shell.execute_reply":"2023-07-27T10:14:56.576345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    split_pct = 0.2\n    split_pct1 = 0.1\n    learning_rate = 0.0002\n    batch_size = 4\n    epochs = 50","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:14:56.578756Z","iopub.execute_input":"2023-07-27T10:14:56.579093Z","iopub.status.idle":"2023-07-27T10:14:56.584931Z","shell.execute_reply.started":"2023-07-27T10:14:56.579062Z","shell.execute_reply":"2023-07-27T10:14:56.583984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Data","metadata":{}},{"cell_type":"code","source":"# extract train data\nwith ZipFile('../input/carvana-image-masking-challenge/train.zip', 'r') as zf:\n    zf.extractall('../working')\n     \nwith ZipFile('../input/carvana-image-masking-challenge/train_masks.zip', 'r') as zf:\n    zf.extractall('../working')","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:14:56.586342Z","iopub.execute_input":"2023-07-27T10:14:56.586674Z","iopub.status.idle":"2023-07-27T10:15:04.853776Z","shell.execute_reply.started":"2023-07-27T10:14:56.586638Z","shell.execute_reply":"2023-07-27T10:15:04.852798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" class MyDataset(Dataset):\n    def __init__(self, root_dir: str, train=True, transforms=None):\n        super(MyDataset, self).__init__()\n        self.train = train\n        self.transforms = transforms\n        \n        file_path = root_dir + 'train/*.*'\n        file_mask_path = root_dir + 'train_masks/*.*'\n        \n        self.images = sorted(glob(file_path))\n        self.image_mask = sorted(glob(file_mask_path))\n        \n        # manually split the train/valid data\n        split_ratio = int(len(self.images) * 0.8)\n        if train:\n            self.images = self.images[:split_ratio]\n            self.image_mask = self.image_mask[:split_ratio]\n        else:\n            self.images = self.images[split_ratio:]\n            self.image_mask = self.image_mask[split_ratio:]\n        \n        \n    def __getitem__(self, index: int):\n        image = Image.open(self.images[index]).convert('RGB')\n        image_mask = Image.open(self.image_mask[index]).convert('L')\n\n        \n        if self.transforms:\n            image = self.transforms(image)\n            image_mask = self.transforms(image_mask)\n                        \n        return {'img': image, 'mask': image_mask}\n    \n    def __len__(self):\n        return len(self.images)\n    \n\ntransforms = T.Compose([\n    T.Resize((width, height)),\n    T.ToTensor(),\n#     T.Normalize(mean=[0.485, 0.456, 0.406],\n#                 std=[0.229, 0.224, 0.225]),\n#     T.RandomHorizontalFlip()\n])\ntrain_dataset = MyDataset(root_dir='../working/',\n                                 train=True,\n                                 transforms=transforms)\nval_dataset = MyDataset(root_dir='../working/',\n                                train=False,\n                                transforms=transforms)\n\n\"\"\"train_dataset_loader = DataLoader(dataset=train_dataset,\n                                 batch_size=batch_size,\n                                 shuffle=True,\n                                 num_workers=n_workers)\nval_dataset_loader = DataLoader(dataset=val_dataset,\n                                batch_size=batch_size,\n                                shuffle=False,\n                                num_workers=n_workers)\"\"\"\n\n        ","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:15:04.855061Z","iopub.execute_input":"2023-07-27T10:15:04.855439Z","iopub.status.idle":"2023-07-27T10:15:04.969935Z","shell.execute_reply.started":"2023-07-27T10:15:04.855405Z","shell.execute_reply":"2023-07-27T10:15:04.968926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\n\n# Assuming the original validation dataset contains 20% of the total data\nval_length = len(val_dataset)\ntest_length = val_length // 2\nval_length = val_length - test_length\n\n# Split the validation dataset into validation and test datasets\nval_dataset, test_dataset = torch.utils.data.random_split(val_dataset, [val_length, test_length])\n\n# Now you have three datasets: train_dataset, val_dataset, and test_dataset\n# The train_dataset contains 80% of the data, the val_dataset contains 10%, and the test_dataset contains 10%.\n\n# You can then create DataLoader objects for each dataset as needed.\ntrain_dataset_loader = DataLoader(dataset=train_dataset, batch_size=CFG.batch_size, shuffle=True, num_workers=n_workers)\nval_dataset_loader = DataLoader(dataset=val_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=n_workers)\ntest_dataset_loader = DataLoader(dataset=test_dataset, batch_size=4, shuffle=False, num_workers=n_workers)","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:15:04.971272Z","iopub.execute_input":"2023-07-27T10:15:04.971662Z","iopub.status.idle":"2023-07-27T10:15:04.986604Z","shell.execute_reply.started":"2023-07-27T10:15:04.971634Z","shell.execute_reply":"2023-07-27T10:15:04.985495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Sample Data","metadata":{}},{"cell_type":"code","source":"samples = next(iter(train_dataset_loader))\n\nfig, (ax1, ax2) = plt.subplots(nrows=2, ncols=1, figsize=(12, 4))\nfig.tight_layout()\n\n\nax1.axis('off')\nax1.set_title('input image')\nax1.imshow(np.transpose(vutils.make_grid(samples['img'], padding=2).numpy(),\n                       (1, 2, 0)))\n\nax2.axis('off')\nax2.set_title('input mask')\nax2.imshow(np.transpose(vutils.make_grid(samples['mask'], padding=2).numpy(),\n                       (1, 2, 0)), cmap='gray')\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:15:04.991210Z","iopub.execute_input":"2023-07-27T10:15:04.991867Z","iopub.status.idle":"2023-07-27T10:15:06.248380Z","shell.execute_reply.started":"2023-07-27T10:15:04.991841Z","shell.execute_reply":"2023-07-27T10:15:06.247256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ARCHITECTURE","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels, mid_channels=None, dilation=1):\n        super().__init__()\n        if not mid_channels:\n            mid_channels = out_channels\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=dilation, dilation=dilation),\n            nn.BatchNorm2d(mid_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, dilation=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n\n\nclass Down(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool2d(2),\n            DoubleConv(in_channels, out_channels)\n        )\n\n    def forward(self, x):\n        return self.maxpool_conv(x)\n\n\nclass Up(nn.Module):\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super().__init__()\n\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n            self.conv = DoubleConv(in_channels, out_channels, in_channels // 2)\n        else:\n            self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)\n            self.conv = DoubleConv(in_channels, out_channels)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n        x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2])\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\n\nclass OutConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(OutConv, self).__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        return self.conv(x)\n\n\nclass TransUNet(nn.Module):\n    def __init__(self, n_channels, n_classes, bilinear=True):\n        super(TransUNet, self).__init__()\n        self.n_channels = n_channels\n        self.n_classes = n_classes\n        self.bilinear = bilinear\n\n        self.inc = DoubleConv(n_channels, 64, dilation=2)  # Set dilation=2 for the first convolutional layer\n\n        self.down1 = Down(64, 128)\n        self.down2 = Down(128, 256)\n        self.down3 = Down(256, 512)\n        factor = 2 if bilinear else 1\n        self.down4 = Down(512, 1024 // factor)\n        self.up1 = Up(1024, 512 // factor, bilinear)\n        self.up2 = Up(512, 256 // factor, bilinear)\n        self.up3 = Up(256, 128 // factor, bilinear)\n        self.up4 = Up(128, 64, bilinear)\n        self.outc = OutConv(64, n_classes)\n\n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n\n        x6 = self.up1(x5, x4)\n        x7 = self.up2(x6, x3)\n        x8 = self.up3(x7, x2)\n        x9 = self.up4(x8, x1)\n\n        logits = self.outc(x9)\n        logits = torch.sigmoid(logits)\n        return logits\n","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:15:06.250246Z","iopub.execute_input":"2023-07-27T10:15:06.250866Z","iopub.status.idle":"2023-07-27T10:15:06.274667Z","shell.execute_reply.started":"2023-07-27T10:15:06.250823Z","shell.execute_reply":"2023-07-27T10:15:06.273347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metrics","metadata":{}},{"cell_type":"code","source":"def dice_score(pred: torch.Tensor, mask: torch.Tensor):\n    dice = (2 * (pred * mask).sum()) / (pred + mask).sum()\n    return (dice.cpu().numpy())\n\nfrom torchmetrics import JaccardIndex\njaccard = JaccardIndex(task='binary', num_classes=2).to(CFG.device)\n\nfrom torchmetrics.classification import BinaryAccuracy\ntrain_accuracy = BinaryAccuracy().to(CFG.device)\nvalid_accuracy = BinaryAccuracy().to(CFG.device)\ntest_accuracy = BinaryAccuracy().to(CFG.device)\n\nfrom torchmetrics.classification import BinaryRecall\nmetric = BinaryRecall().to(CFG.device)\n\nfrom torchmetrics.classification import BinaryPrecision\nmetric1 = BinaryPrecision().to(CFG.device)\n\nfrom torchmetrics.classification import BinarySpecificity\nmetric2 = BinarySpecificity().to(CFG.device)","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:15:06.276059Z","iopub.execute_input":"2023-07-27T10:15:06.276726Z","iopub.status.idle":"2023-07-27T10:15:19.622298Z","shell.execute_reply.started":"2023-07-27T10:15:06.276683Z","shell.execute_reply":"2023-07-27T10:15:19.621233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BestMetric:\n    def __init__(self):\n        self.best_value = float('-inf')\n    \n    def update(self, value):\n        if value > self.best_value:\n            self.best_value = value\n        return self.best_value\n    \n    def reset(self):\n        self.best_value = float('-inf')\n    \n    def get_best(self):\n        return self.best_value","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:15:19.623875Z","iopub.execute_input":"2023-07-27T10:15:19.624711Z","iopub.status.idle":"2023-07-27T10:15:19.631506Z","shell.execute_reply.started":"2023-07-27T10:15:19.624682Z","shell.execute_reply":"2023-07-27T10:15:19.630434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Plotting","metadata":{}},{"cell_type":"code","source":"def plot_pred_img(samples, pred):\n    fig, (ax1, ax2, ax3) = plt.subplots(nrows=3, ncols=1, figsize=(12, 6))\n    fig.tight_layout()\n\n\n    ax1.axis('off')\n    ax1.set_title('input image')\n    ax1.imshow(np.transpose(vutils.make_grid(samples['img'], padding=2).numpy(),\n                           (1, 2, 0)))\n\n    ax2.axis('off')\n    ax2.set_title('input mask')\n    ax2.imshow(np.transpose(vutils.make_grid(samples['mask'], padding=2).numpy(),\n                           (1, 2, 0)), cmap='gray')\n    \n    ax3.axis('off')\n    ax3.set_title('predicted mask')\n    ax3.imshow(np.transpose(vutils.make_grid(pred, padding=2).cpu().numpy(),\n                           (1, 2, 0)), cmap='gray')\n\n    plt.show()\n    \n    \ndef plot_train_progress_A(model):\n#     model.eval()\n\n#     with torch.no_grad():\n    samples = next(iter(val_dataset_loader))\n    val_img = samples['img'].to(device)\n    val_mask = samples['mask'].to(device)\n\n    pred = model(val_img)\n\n    plot_pred_img(samples, pred.detach())","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:15:19.633120Z","iopub.execute_input":"2023-07-27T10:15:19.633763Z","iopub.status.idle":"2023-07-27T10:15:19.646229Z","shell.execute_reply.started":"2023-07-27T10:15:19.633729Z","shell.execute_reply":"2023-07-27T10:15:19.645430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_train_progress(epoch, model, samples, pred):\n    if epoch == 25:\n        fig, (ax1, ax2, ax3) = plt.subplots(nrows=1, ncols=3, figsize=(12, 6))\n        fig.tight_layout()\n\n        ax1.axis('off')\n        ax1.set_title('input image')\n        ax1.imshow(np.transpose(vutils.make_grid(samples['img'], padding=2).numpy(),\n                                (1, 2, 0)))\n\n        ax2.axis('off')\n        ax2.set_title('input mask')\n        ax2.imshow(np.transpose(vutils.make_grid(samples['mask'], padding=2).numpy(),\n                                (1, 2, 0)), cmap='gray')\n\n        ax3.axis('off')\n        ax3.set_title('predicted mask')\n        ax3.imshow(np.transpose(vutils.make_grid(pred, padding=2).cpu().numpy(),\n                                (1, 2, 0)), cmap='gray')\n\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:16:49.949870Z","iopub.execute_input":"2023-07-27T10:16:49.950249Z","iopub.status.idle":"2023-07-27T10:16:49.959445Z","shell.execute_reply.started":"2023-07-27T10:16:49.950213Z","shell.execute_reply":"2023-07-27T10:16:49.958321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# LOSS FUNCTION","metadata":{}},{"cell_type":"code","source":"class DiceBCELoss(nn.Module):\n    def __init__(self, weight=None, size_average=True, gamma = 0.2, beta = 0.8):\n        super(DiceBCELoss, self).__init__()\n        self.gamma = gamma\n        self.beta = beta\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 = F.sigmoid(inputs)       \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_loss = 1 - (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n        BCE_with_Logits = F.binary_cross_entropy_with_logits(inputs, targets, reduction='mean')\n        Dice_BCE = BCE_with_Logits + dice_loss\n        \n        dice_loss = self.beta * dice_loss\n        BCE_with_Logits = self.gamma * BCE_with_Logits\n\n        Dice_BCE = BCE_with_Logits + dice_loss\n\n        \n        return Dice_BCE","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:16:51.391474Z","iopub.execute_input":"2023-07-27T10:16:51.391835Z","iopub.status.idle":"2023-07-27T10:16:51.400493Z","shell.execute_reply.started":"2023-07-27T10:16:51.391807Z","shell.execute_reply":"2023-07-27T10:16:51.399422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"def train(model, optimizer, criterion, scheduler=None, CFG=CFG):\n    train_losses = []\n    val_losses = []\n    train_accs = []\n    valid_accs = []\n    dice_scores = []  # List to store Dice coefficients\n    lr_rates = []\n\n    best_metric1 = BestMetric()\n    best_metric2 = BestMetric()\n    best_metric3 = BestMetric()\n    best_metric4 = BestMetric()\n    best_metric5 = BestMetric()\n\n    for epoch in range(CFG.epochs):\n        model.train()\n        train_total_loss = 0\n        train_iterations = 0\n        train_acc_total = 0\n\n        for idx, data in enumerate(tqdm(train_dataset_loader)):\n            train_iterations += 1\n            train_img = data['img'].to(device)\n            train_mask = data['mask'].to(device)\n            threshold = 0.5\n            train_mask1 = (train_mask >= threshold).float()\n\n            optimizer.zero_grad()\n            with torch.autocast(device_type='cuda'):\n                train_output_mask = model(train_img)\n                train_loss = criterion(train_output_mask, train_mask)\n                train_total_loss += train_loss.item()\n                train_acc_total += train_accuracy(train_output_mask, train_mask1)\n\n            train_loss.backward()\n            optimizer.step()\n\n        train_epoch_loss = train_total_loss / train_iterations\n        train_acc = train_acc_total / train_iterations\n        train_losses.append(train_epoch_loss)\n        train_accs.append(train_acc)\n\n        model.eval()\n        with torch.no_grad():\n            val_total_loss = 0\n            val_iterations = 0\n            total_valid_acc = 0\n            scores = 0\n            jaccard_scores = 0  # New variable to accumulate Jaccard scores\n            metric_scores = 0\n            metric1_scores = 0\n            metric2_scores = 0\n\n            for vidx, val_data in enumerate(tqdm(val_dataset_loader)):\n                val_iterations += 1\n                val_img = val_data['img'].to(device)\n                val_mask = val_data['mask'].to(device)\n                threshold = 0.5\n                val_mask1 = (val_mask >= threshold).float()\n\n                with torch.no_grad():\n                    pred = model(val_img)\n                    # Ensure the predictions and target tensors have the same shape\n                    pred = pred[:, 0:1, :, :]\n\n                    val_loss = criterion(pred, val_mask)\n                    val_total_loss += val_loss.item()\n                    total_valid_acc += valid_accuracy(pred, val_mask1)\n                    scores += dice_score(pred, val_mask)\n                    jaccard_scores += jaccard(pred, val_mask1) # Accumulate Jaccard scores\n                    metric_scores += metric(pred, val_mask1)\n                    metric1_scores += metric1(pred, val_mask1)\n                    metric2_scores += metric2(pred, val_mask1)\n            \n            val_epoch_loss = val_total_loss / val_iterations\n            val_acc = total_valid_acc / val_iterations\n            val_losses.append(val_epoch_loss)\n            valid_accs.append(val_acc)\n            dice_coef_score = scores / val_iterations\n            dice_scores.append(dice_coef_score)  # Store Dice coefficient\n            jaccard_score = jaccard_scores / val_iterations  # Average Jaccard index\n            metric_score = metric_scores / val_iterations\n            metric1_score = metric1_scores / val_iterations\n            metric2_score = metric2_scores / val_iterations \n\n            best_metric1.update(dice_coef_score)\n            best_metric2.update(jaccard_score)\n            best_metric3.update(metric1(pred, val_mask1))\n            best_metric4.update(train_acc)\n            best_metric5.update(val_acc)\n\n            if epoch == 25:\n                samples = next(iter(val_dataset_loader))\n                val_img = samples['img'].to(CFG.device)\n                val_mask = samples['mask'].to(CFG.device)\n                pred = model(val_img)\n                plot_train_progress(epoch, model, samples, pred.detach())\n\n            print('epochs - {}/{} [{}/{}], dice score: {}, train loss: {}, val loss: {}, iou score:{}, train_accuracy:{}, valid_accuracy:{}, Recall:{}, Precision:{}, Specificity:{}'.format(\n                epoch+1, CFG.epochs,\n                idx+1, len(train_dataset_loader),\n                dice_coef_score, train_epoch_loss, val_epoch_loss, jaccard_score, train_acc, val_acc, metric_score, metric1_score, metric2_score \n            ))\n\n        lr_rates.append(optimizer.param_groups[0]['lr'])\n        if scheduler:\n            scheduler.step()\n\n    best_DiceScore = best_metric1.get_best()\n    best_IoUscore = best_metric2.get_best()\n    best_precision = best_metric3.get_best()\n    best_TrainAcc = best_metric4.get_best()\n    best_ValidAcc = best_metric5.get_best()\n    print('Best Dice Score: {:.4f}'.format(best_DiceScore))\n    print('Best IoU Score: {:.4f}'.format(best_IoUscore))\n    print('Best Precision: {:.4f}'.format(best_precision))\n    print('Best Train Accuracy: {:.4f}'.format(best_TrainAcc))\n    print('Best Valid Accuracy: {:.4f}'.format(best_ValidAcc))\n\n    return {\n        'lr': lr_rates,\n        'train_loss': train_losses,\n        'valid_loss': val_losses,\n        'train_acc': train_accs,\n        'valid_acc': valid_accs,\n        'dice_scores': dice_scores  # Include dice scores in the returned dictionary\n    }","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:16:52.204944Z","iopub.execute_input":"2023-07-27T10:16:52.205340Z","iopub.status.idle":"2023-07-27T10:16:52.228567Z","shell.execute_reply.started":"2023-07-27T10:16:52.205287Z","shell.execute_reply":"2023-07-27T10:16:52.227440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = TransUNet(n_channels=3, n_classes=1).to(device)\n#criterion = nn.BCEWithLogitsLoss()\n# criterion = smp.losses.DiceLoss(mode='binary')\ncriterion = DiceBCELoss()\n#optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)\n#optimizer = optim.Adamax(model.parameters(), lr=0.0002)\noptimizer = AdaBound(model.parameters(), lr=0.002, final_lr=0.1)\n#optimizer = torch.optim.Adadelta(model.parameters(), lr=learning_rate)","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:16:53.272491Z","iopub.execute_input":"2023-07-27T10:16:53.272846Z","iopub.status.idle":"2023-07-27T10:16:53.482004Z","shell.execute_reply.started":"2023-07-27T10:16:53.272816Z","shell.execute_reply":"2023-07-27T10:16:53.481017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_img = samples['img'].to(CFG.device)\npred = model(val_img)\npred","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:16:53.574057Z","iopub.execute_input":"2023-07-27T10:16:53.574429Z","iopub.status.idle":"2023-07-27T10:16:53.624183Z","shell.execute_reply.started":"2023-07-27T10:16:53.574398Z","shell.execute_reply":"2023-07-27T10:16:53.623144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = train(model, optimizer, criterion) ","metadata":{"execution":{"iopub.status.busy":"2023-07-27T10:16:53.978110Z","iopub.execute_input":"2023-07-27T10:16:53.979024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TEST","metadata":{}},{"cell_type":"code","source":"def create_dir(path):\n    if not os.path.exists(path):\n        os.makedirs(path)","metadata":{"execution":{"iopub.status.busy":"2023-07-26T11:05:54.632931Z","iopub.execute_input":"2023-07-26T11:05:54.633424Z","iopub.status.idle":"2023-07-26T11:05:54.640341Z","shell.execute_reply.started":"2023-07-26T11:05:54.633383Z","shell.execute_reply":"2023-07-26T11:05:54.639080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"create_dir(\"new_data/test/image_TransUNET-DILATED/\")","metadata":{"execution":{"iopub.status.busy":"2023-07-26T11:05:55.046633Z","iopub.execute_input":"2023-07-26T11:05:55.047054Z","iopub.status.idle":"2023-07-26T11:05:55.052977Z","shell.execute_reply.started":"2023-07-26T11:05:55.047022Z","shell.execute_reply":"2023-07-26T11:05:55.051999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport hashlib\nimport torchvision.transforms as transforms\nfrom PIL import Image, ImageDraw, ImageFont","metadata":{"execution":{"iopub.status.busy":"2023-07-26T18:20:28.750309Z","iopub.execute_input":"2023-07-26T18:20:28.750690Z","iopub.status.idle":"2023-07-26T18:20:28.755984Z","shell.execute_reply.started":"2023-07-26T18:20:28.750658Z","shell.execute_reply":"2023-07-26T18:20:28.754857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"font_size = 22\nfont_path = \"/kaggle/input/fontss/Agdasima-Bold.ttf\"  # Update the path based on your font file location\nfont = ImageFont.truetype(font_path, size=font_size)","metadata":{"execution":{"iopub.status.busy":"2023-07-26T18:20:29.424196Z","iopub.execute_input":"2023-07-26T18:20:29.424978Z","iopub.status.idle":"2023-07-26T18:20:29.431735Z","shell.execute_reply.started":"2023-07-26T18:20:29.424942Z","shell.execute_reply":"2023-07-26T18:20:29.430688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test(model, test_dataset_loader, criterion, save_dir):\n    model.eval()\n    test_loss = 0\n    test_iterations = 0\n    dice_score_total = 0\n    iou_score_total = 0\n\n    transform = transforms.ToPILImage()\n    gap_size = 20  # Adjust the gap size between the images\n\n    with torch.no_grad():\n        for idx, data in enumerate(tqdm(test_dataset_loader)):\n            test_iterations += 1\n            test_img = data['img'].to(device)\n            test_mask = data['mask'].to(device)\n            threshold = 0.5  # You can adjust this threshold based on your needs\n            test_mask1 = (test_mask >= threshold).float()\n            \n            test_img = test_img.unsqueeze(0)\n\n            with torch.autocast(device_type='cuda'):\n                test_output_mask = model(test_img)\n                test_loss = criterion(test_output_mask, test_mask)\n                dice_score_batch = dice_score(test_output_mask, test_mask)\n                dice_score_total += dice_score_batch\n\n                for i in range(test_img.size(0)):\n                    img = transform(test_img[i].cpu())\n                    mask = transform(test_mask[i].cpu())\n                    predicted_mask = transform(test_output_mask[i].cpu())\n                    img_width, img_height = img.size\n                    combined_img_width = img_width * 3 + gap_size * 2\n                    combined_img_height = img_height + gap_size * 2\n                    combined_img = Image.new('RGB', (combined_img_width, combined_img_height), color='white')\n                    combined_img.paste(img, (gap_size, gap_size))\n                    combined_img.paste(mask, (img_width + gap_size * 2, gap_size))\n                    combined_img.paste(predicted_mask, (img_width * 2 + gap_size * 3, gap_size))\n\n                    draw = ImageDraw.Draw(combined_img)\n                    # Adjust the font size here\n                    # font = ImageFont.truetype(\"arial.ttf\", 16)\n\n                    # Get the original image name from the test_dataset_loader\n                    original_img_path = test_dataset_loader.dataset.images[idx * test_dataset_loader.batch_size + i]\n                    original_img_name = os.path.basename(original_img_path)\n\n                    # Generate a unique hash for the original image name\n                    img_name_hash = hashlib.md5(original_img_name.encode()).hexdigest()\n\n                    # Construct the predicted image name using the hash\n                    predicted_img_name = f'predicted_{img_name_hash}.png'\n\n                    draw.text((gap_size, 0), \"Original Image\", fill=\"black\", font=font)\n                    draw.text((img_width + gap_size * 2, 0), \"Ground Truth Mask\", fill=\"black\", font=font)\n                    draw.text((img_width * 2 + gap_size * 3, 0), \"Predicted Mask\", fill=\"black\", font=font)\n\n                    combined_img.save(os.path.join(save_dir, predicted_img_name))\n\n    test_epoch_loss = test_loss / test_iterations\n    average_dice_score = dice_score_total / test_iterations\n\n    print('Test loss: {}, Dice score: {}'.format(test_epoch_loss, average_dice_score))\n\n    return test_epoch_loss, average_dice_score\n","metadata":{"execution":{"iopub.status.busy":"2023-07-26T18:28:58.989029Z","iopub.execute_input":"2023-07-26T18:28:58.989386Z","iopub.status.idle":"2023-07-26T18:28:59.004335Z","shell.execute_reply.started":"2023-07-26T18:28:58.989356Z","shell.execute_reply":"2023-07-26T18:28:59.003216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_dir = '/kaggle/working/new_data/test/image_TransUNET-DILATED'\ntest_loss, test_accuracy = test(model, test_dataset, criterion, save_dir)","metadata":{"execution":{"iopub.status.busy":"2023-07-26T18:28:59.536443Z","iopub.execute_input":"2023-07-26T18:28:59.537160Z","iopub.status.idle":"2023-07-26T18:28:59.713339Z","shell.execute_reply.started":"2023-07-26T18:28:59.537126Z","shell.execute_reply":"2023-07-26T18:28:59.711849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# RESULTS","metadata":{}},{"cell_type":"code","source":"# Accuracy graph\ndef plot_accuracy(train_acc, valid_acc):\n    epochs = range(1, len(train_acc) + 1)\n    train_acc_cpu = [acc.cpu().numpy() for acc in train_acc]  # Convert train_acc to CPU\n    valid_acc_cpu = [acc.cpu().numpy() for acc in valid_acc]  # Convert valid_acc to CPU\n    plt.plot(epochs, train_acc_cpu, 'b', label='Train Accuracy')\n    plt.plot(epochs, valid_acc_cpu, 'r', label='Valid Accuracy')\n    plt.title('Train Accuracy vs Valid Accuracy')\n    plt.xlabel('Epochs')\n    plt.ylabel('Accuracy')\n    plt.legend()\n    plt.show()\n    \ntrain_acc = history['train_acc']\nvalid_acc = history['valid_acc']\nplot_accuracy(train_acc, valid_acc)","metadata":{"execution":{"iopub.status.busy":"2023-07-26T11:07:31.808763Z","iopub.execute_input":"2023-07-26T11:07:31.809134Z","iopub.status.idle":"2023-07-26T11:07:31.865192Z","shell.execute_reply.started":"2023-07-26T11:07:31.809103Z","shell.execute_reply":"2023-07-26T11:07:31.863755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plotting the Dice coefficients\nepochs = range(1, len(history['dice_scores']) + 1)\ndice_scores = history['dice_scores']\n\nplt.figure(figsize=(8, 6))\nplt.plot(epochs, dice_scores, marker='o')\nplt.xlabel('Epochs')\nplt.ylabel('Dice Coefficient')\nplt.title('Dice Coefficient Progress')\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Loss Graph\nplt.plot(history['train_loss'], label='constant LR train loss')\nplt.plot(history['valid_loss'], label='constant LR val loss')\nplt.xlabel('epoch')\nplt.ylabel('loss')\nplt.legend()\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# SAVING THE MODEL","metadata":{}},{"cell_type":"code","source":"import shutil\n\n# Create a zip file\nshutil.make_archive('/kaggle/working/new_data/test/image_TransUNET-DILATED', 'zip', save_dir)\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\n# Save the entire model\ntorch.save(model, 'model_TransUNET-DILATED.pth')\n\n# Save only the model state dictionary\ntorch.save(model.state_dict(), 'model_state_TransUNET-DILATED.pth')\n","metadata":{},"execution_count":null,"outputs":[]}]}