{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":6927,"databundleVersionId":45059,"sourceType":"competition"}],"dockerImageVersionId":30512,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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-06-19T09:16:14.499385Z","iopub.execute_input":"2023-06-19T09:16:14.499865Z","iopub.status.idle":"2023-06-19T09:16:42.393573Z","shell.execute_reply.started":"2023-06-19T09:16:14.499815Z","shell.execute_reply":"2023-06-19T09:16:42.392482Z"},"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\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-06-19T09:17:25.544809Z","iopub.execute_input":"2023-06-19T09:17:25.545203Z","iopub.status.idle":"2023-06-19T09:17:26.208747Z","shell.execute_reply.started":"2023-06-19T09:17:25.545172Z","shell.execute_reply":"2023-06-19T09:17:26.207380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"batch_size = 4\nn_iters = 10000\nepochs = 50\nlearning_rate = 0.0002\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-06-19T09:17:30.847265Z","iopub.execute_input":"2023-06-19T09:17:30.848391Z","iopub.status.idle":"2023-06-19T09:17:30.935336Z","shell.execute_reply.started":"2023-06-19T09:17:30.848329Z","shell.execute_reply":"2023-06-19T09:17:30.934214Z"},"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-06-19T09:17:34.251659Z","iopub.execute_input":"2023-06-19T09:17:34.252039Z"},"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.7)\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\ntrain_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=True,\n                                num_workers=n_workers)\n\n        ","metadata":{"execution":{"iopub.status.busy":"2023-06-19T08:14:20.387374Z","iopub.execute_input":"2023-06-19T08:14:20.387762Z","iopub.status.idle":"2023-06-19T08:14:20.487597Z","shell.execute_reply.started":"2023-06-19T08:14:20.387732Z","shell.execute_reply":"2023-06-19T08:14:20.486650Z"},"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-06-19T08:14:23.970114Z","iopub.execute_input":"2023-06-19T08:14:23.970512Z","iopub.status.idle":"2023-06-19T08:14:25.223868Z","shell.execute_reply.started":"2023-06-19T08:14:23.970476Z","shell.execute_reply":"2023-06-19T08:14:25.222828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#unet part\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass DoubleConv(nn.Module):\n    \"\"\"(convolution => [BN] => ReLU) * 2\"\"\"\n\n    def __init__(self, in_channels, out_channels, mid_channels=None):\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=1),\n            nn.BatchNorm2d(mid_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=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    \"\"\"Downscaling with maxpool then double conv\"\"\"\n\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    \"\"\"Upscaling then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super().__init__()\n\n        # if bilinear, use the normal convolutions to reduce the number of channels\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\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        # input is CHW\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n\n        x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,\n                        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)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T08:14:27.038766Z","iopub.execute_input":"2023-06-19T08:14:27.039147Z","iopub.status.idle":"2023-06-19T08:14:27.056821Z","shell.execute_reply.started":"2023-06-19T08:14:27.039099Z","shell.execute_reply":"2023-06-19T08:14:27.055476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#unet_att\nimport torch.nn as nn\n\nclass MultiConv(nn.Module):\n    def __init__(self, in_ch, out_ch, attn=True):\n        super(MultiConv, self).__init__()\n\n        self.fuse_attn = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.PReLU(),\n            nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.PReLU(),\n            nn.Conv2d(out_ch, out_ch, kernel_size=1),\n            nn.BatchNorm2d(out_ch),\n            nn.Softmax2d() if attn else nn.PReLU()\n        )\n\n    def forward(self, x):\n        return self.fuse_attn(x)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T08:14:28.459290Z","iopub.execute_input":"2023-06-19T08:14:28.460253Z","iopub.status.idle":"2023-06-19T08:14:28.469615Z","shell.execute_reply.started":"2023-06-19T08:14:28.460218Z","shell.execute_reply":"2023-06-19T08:14:28.468616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#unet_transformer\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn\n\n\nclass PAM_Module(nn.Module):\n    def __init__(self, in_dim):\n        super(PAM_Module, self).__init__()\n        self.chanel_in = in_dim\n\n        self.query_conv = nn.Conv2d(in_channels=in_dim, out_channels=in_dim // 8, kernel_size=1)\n        self.key_conv = nn.Conv2d(in_channels=in_dim, out_channels=in_dim // 8, kernel_size=1)\n        self.value_conv = nn.Conv2d(in_channels=in_dim, out_channels=in_dim, kernel_size=1)\n        self.gamma = nn.Parameter(torch.zeros(1))\n        self.softmax = nn.Softmax(dim=-1)\n\n    def forward(self, x):\n        m_batchsize, C, height, width = x.size()\n        proj_query = self.query_conv(x).view(m_batchsize, -1, width * height).permute(0, 2, 1)\n        proj_key = self.key_conv(x).view(m_batchsize, -1, width * height)\n\n        energy = torch.bmm(proj_query, proj_key)\n        attention = self.softmax(energy)\n        proj_value = self.value_conv(x).view(m_batchsize, -1, width * height)\n\n        out = torch.bmm(proj_value, attention.permute(0, 2, 1))\n        out = out.view(m_batchsize, C, height, width)\n\n        out = self.gamma * out + x\n        return out\n\nclass PositionEmbeddingLearned(nn.Module):\n    \n    def __init__(self, num_pos_feats=256, len_embedding=32):\n        super().__init__()\n        self.row_embed = nn.Embedding(len_embedding, num_pos_feats)\n        self.col_embed = nn.Embedding(len_embedding, num_pos_feats)\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        nn.init.uniform_(self.row_embed.weight)\n        nn.init.uniform_(self.col_embed.weight)\n\n    def forward(self, tensor_list):\n        x = tensor_list\n        h, w = x.shape[-2:]\n        i = torch.arange(w, device=x.device)\n        j = torch.arange(h, device=x.device)\n\n        x_emb = self.col_embed(i)\n        y_emb = self.row_embed(j)\n\n        pos = torch.cat([\n            x_emb.unsqueeze(0).repeat(h, 1, 1),\n            y_emb.unsqueeze(1).repeat(1, w, 1),\n        ], dim=-1).permute(2, 0, 1).unsqueeze(0).repeat(x.shape[0], 1, 1, 1)\n\n        return pos\n\nclass ScaledDotProductAttention(nn.Module):\n\n    def __init__(self, temperature, attn_dropout=0.1):\n        super().__init__()\n        self.temperature = temperature ** 0.5\n        self.dropout = nn.Dropout(attn_dropout)\n\n    def forward(self, x, mask=None):\n        m_batchsize, d, height, width = x.size()\n        q = x.view(m_batchsize, d, -1)\n        k = x.view(m_batchsize, d, -1)\n        k = k.permute(0, 2, 1)\n        v = x.view(m_batchsize, d, -1)\n\n        attn = torch.matmul(q / self.temperature, k)\n\n        if mask is not None:\n            attn = attn.masked_fill(mask == 0, -1e9)\n\n        attn = self.dropout(F.softmax(attn, dim=-1))\n        output = torch.matmul(attn, v)\n        output = output.view(m_batchsize, d, height, width)\n\n        return output\n\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels, mid_channels=None):\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=1),\n            nn.BatchNorm2d(mid_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T08:14:29.785216Z","iopub.execute_input":"2023-06-19T08:14:29.785623Z","iopub.status.idle":"2023-06-19T08:14:29.809099Z","shell.execute_reply.started":"2023-06-19T08:14:29.785588Z","shell.execute_reply":"2023-06-19T08:14:29.808124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UNet_Attention_Transformer_Multiscale(nn.Module):\n    def __init__(self, n_channels, n_classes, bilinear=True):\n        super(UNet_Attention_Transformer_Multiscale, 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)\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(1024, 256 // factor, bilinear)\n        self.up3 = Up(512, 128 // factor, bilinear)\n        self.up4 = Up(256, 64, bilinear)\n        self.outc = OutConv(128, n_classes)\n\n        self.pos = PositionEmbeddingLearned(512 // factor)\n\n        self.pam = PAM_Module(512)\n\n        self.sdpa = ScaledDotProductAttention(512)\n        \n        self.fuse1 = MultiConv(768, 256)\n        self.fuse2 = MultiConv(384, 128)\n        self.fuse3 = MultiConv(192, 64)\n        self.fuse4 = MultiConv(128, 64)\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\n        x5_pam = self.pam(x5)\n\n        x5_pos = self.pos(x5)\n        x5 = x5 + x5_pos\n\n\n        x5_sdpa = self.sdpa(x5)\n        x5 = x5_sdpa + x5_pam\n        \n\n        x6 = self.up1(x5, x4)\n        x5_scale = F.interpolate(x5, size=x6.shape[2:], mode='bilinear', align_corners=True)\n        x6_cat = torch.cat((x5_scale, x6), 1)\n\n        x7 = self.up2(x6_cat, x3)\n        x6_scale = F.interpolate(x6, size=x7.shape[2:], mode='bilinear', align_corners=True)\n        x7_cat = torch.cat((x6_scale, x7), 1)\n\n        x8 = self.up3(x7_cat, x2)\n        x7_scale = F.interpolate(x7, size=x8.shape[2:], mode='bilinear', align_corners=True)\n        x8_cat = torch.cat((x7_scale, x8), 1)\n\n        x9 = self.up4(x8_cat, x1)\n        x8_scale = F.interpolate(x8, size=x9.shape[2:], mode='bilinear', align_corners=True)\n        x9 = torch.cat((x8_scale, x9), 1)\n\n        logits = self.outc(x9)\n        #logits = torch.sigmoid(logits)\n        return logits","metadata":{"execution":{"iopub.status.busy":"2023-06-19T08:14:38.595595Z","iopub.execute_input":"2023-06-19T08:14:38.596729Z","iopub.status.idle":"2023-06-19T08:14:38.612999Z","shell.execute_reply.started":"2023-06-19T08:14:38.596683Z","shell.execute_reply":"2023-06-19T08:14:38.611988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"It is a binary segmentation problem. We use *BCEWithLogitsLoss* loss function. You can change the loss function to *BCELoss* by adding sigmoid layer to the output.","metadata":{}},{"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 np.mean(dice.cpu().numpy())\n\ndef iou_score(pred: torch.Tensor, mask: torch.Tensor):\n    intersection = torch.logical_and(pred, mask).sum()\n    union = torch.logical_or(pred, mask).sum()\n    iou = intersection.float() / union.float()\n    return iou\n\ndef pixel_accuracy(pred: torch.Tensor, mask: torch.Tensor):\n    correct = torch.eq(pred, mask).int()\n    return float(correct.sum()) / float(correct.numel())","metadata":{"execution":{"iopub.status.busy":"2023-06-19T08:14:41.477852Z","iopub.execute_input":"2023-06-19T08:14:41.478625Z","iopub.status.idle":"2023-06-19T08:14:41.487122Z","shell.execute_reply.started":"2023-06-19T08:14:41.478585Z","shell.execute_reply":"2023-06-19T08:14:41.486138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","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(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\n    plot_pred_img(samples, pred.detach())","metadata":{"execution":{"iopub.status.busy":"2023-06-19T08:14:44.917530Z","iopub.execute_input":"2023-06-19T08:14:44.917990Z","iopub.status.idle":"2023-06-19T08:14:44.929722Z","shell.execute_reply.started":"2023-06-19T08:14:44.917961Z","shell.execute_reply":"2023-06-19T08:14:44.928741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, optimizer, criteration, scheduler=None):\n    train_losses = []\n    val_lossess = []\n    lr_rates = []\n    \n    # calculate train epochs\n    epochs = int(n_iters / (len(train_dataset) / batch_size))\n\n\n    for epoch in range(epochs):\n        model.train()        \n        train_total_loss = 0\n        train_iterations = 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            \n            optimizer.zero_grad()\n            # speed up the training\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\n            train_loss.backward()\n            optimizer.step()\n\n\n\n        train_epoch_loss = train_total_loss / train_iterations\n        train_losses.append(train_epoch_loss)\n        \n        # evaluate mode\n        model.eval()\n        with torch.no_grad():\n            val_total_loss = 0\n            val_iterations = 0\n            scores = 0\n            score_iou = 0\n            score_acc = 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\n                with torch.autocast(device_type='cuda'):\n                    pred = model(val_img)\n                    val_loss = criterion(pred, val_mask)\n                    val_total_loss += val_loss.item()\n                    scores += dice_score(pred, val_mask)\n                    score_iou += iou_score(pred, val_mask)\n                    \n\n\n            val_epoch_loss = val_total_loss / val_iterations\n            dice_coef_scroe = scores / val_iterations\n            score_iou_coef = score_iou/val_iterations\n\n            val_lossess.append(val_epoch_loss)           \n\n            plot_train_progress(model)\n            print('epochs - {}/{} [{}/{}], dice score: {}, train loss: {}, val loss: {}, iou score:{}'.format(\n                epoch+1, epochs,\n                idx+1, len(train_dataset_loader),\n                dice_coef_scroe, train_epoch_loss, val_epoch_loss, score_iou_coef\n            )) \n            \n        lr_rates.append(optimizer.param_groups[0]['lr'])\n        if scheduler:\n            scheduler.step() # decay learning rate\n            print('LR rate:', scheduler.get_last_lr())\n            \n    return {\n        'lr': lr_rates,\n        'train_loss': train_losses,\n        'valid_loss': val_lossess\n    }","metadata":{"execution":{"iopub.status.busy":"2023-06-19T08:15:08.919794Z","iopub.execute_input":"2023-06-19T08:15:08.920172Z","iopub.status.idle":"2023-06-19T08:15:08.934388Z","shell.execute_reply.started":"2023-06-19T08:15:08.920142Z","shell.execute_reply":"2023-06-19T08:15:08.933200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet_Attention_Transformer_Multiscale(n_channels=3, n_classes=1).to(device)\ncriterion = nn.BCEWithLogitsLoss()\n#criterion = smp.losses.DiceLoss(mode='binary')\noptimizer = torch.optim.SGD(model.parameters(), lr=learning_rate)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-19T08:15:10.424241Z","iopub.execute_input":"2023-06-19T08:15:10.424930Z","iopub.status.idle":"2023-06-19T08:15:15.076285Z","shell.execute_reply.started":"2023-06-19T08:15:10.424890Z","shell.execute_reply":"2023-06-19T08:15:15.075310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_img = samples['img'].to(device)\npred = model(val_img)\npred","metadata":{"execution":{"iopub.status.busy":"2023-06-19T08:15:22.241025Z","iopub.execute_input":"2023-06-19T08:15:22.241388Z","iopub.status.idle":"2023-06-19T08:15:32.050343Z","shell.execute_reply.started":"2023-06-19T08:15:22.241359Z","shell.execute_reply":"2023-06-19T08:15:32.049385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet_Attention_Transformer_Multiscale(n_channels=3, n_classes=1).to(device)\ncriterion = nn.BCEWithLogitsLoss()\n#criterion = smp.losses.DiceLoss(mode='binary')\noptimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)\n\nhistory = train(model, optimizer, criterion)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T06:31:12.443842Z","iopub.execute_input":"2023-06-19T06:31:12.444204Z","iopub.status.idle":"2023-06-19T06:41:18.106964Z","shell.execute_reply.started":"2023-06-19T06:31:12.444174Z","shell.execute_reply":"2023-06-19T06:41:18.104523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#accuracy\ndef pixel_accuracy(pred: torch.Tensor, mask: torch.Tensor):\n    correct = torch.eq(pred, val_mask).int()\n    return float(correct.sum()) / float(correct.numel())\n\ndef train_A(model, optimizer, criteration, scheduler=None):\n    train_losses = []\n    val_lossess = []\n    lr_rates = []\n    \n    # calculate train epochs\n    epochs = int(n_iters / (len(train_dataset) / batch_size))\n\n\n    for epoch in range(epochs):\n        model.train()        \n        train_total_loss = 0\n        train_iterations = 0\n        \n        for idx, img_mask in enumerate(tqdm(train_dataset_loader)):\n            train_iterations += 1\n            train_img = samples['img'].float().to(device)\n            train_mask = samples['mask'].float().to(device)\n            \n            optimizer.zero_grad()\n            # speed up the training\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\n            train_loss.backward()\n            optimizer.step()\n\n\n\n        train_epoch_loss = train_total_loss / train_iterations\n        train_losses.append(train_epoch_loss)\n        \n        # evaluate mode\n        model.eval()\n        with torch.no_grad():\n            val_total_loss = 0\n            val_iterations = 0\n            scores_acc = 0\n\n            for vidx, img_mask in enumerate(tqdm(val_dataset_loader)):\n                val_iterations += 1\n                val_img = samples['img'].float().to(device)\n                val_mask = samples['mask'].float().to(device)\n                def pixel_accuracy(pred: torch.Tensor, mask: torch.Tensor):\n                    correct = torch.eq(pred, val_mask).int()\n                    return float(correct.sum()) / float(correct.numel())\n\n                with torch.autocast(device_type='cuda'):\n                    pred = model(val_img)\n                    scores_acc += pixel_accuracy(pred, val_mask)\n\n\n            scores_acc_coeff = scores_acc/ val_iterations\n                      \n\n            #plot_train_progress(model)\n            print('epochs - {}/{} [{}/{}], Accuracy : {}'.format(\n                epoch+1, epochs,\n                idx+1, len(train_dataset_loader),\n                scores_acc_coeff\n            )) \n            \n        lr_rates.append(optimizer.param_groups[0]['lr'])\n        if scheduler:\n            scheduler.step() # decay learning rate\n            print('LR rate:', scheduler.get_last_lr())\n            \n    return {\n        'lr': lr_rates,\n        'train_loss': train_losses,\n        'valid_loss': val_lossess\n    }","metadata":{"execution":{"iopub.status.busy":"2023-06-19T05:32:10.456349Z","iopub.status.idle":"2023-06-19T05:32:10.456975Z","shell.execute_reply.started":"2023-06-19T05:32:10.456707Z","shell.execute_reply":"2023-06-19T05:32:10.456741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet_Attention_Transformer_Multiscale(n_channels=3, n_classes=1).to(device)\n#criterion = nn.BCEWithLogitsLoss()\ncriterion = smp.losses.DiceLoss(mode='binary')\noptimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)\n\nhistory = train_A(model, optimizer, criterion)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T05:32:10.458772Z","iopub.status.idle":"2023-06-19T05:32:10.459262Z","shell.execute_reply.started":"2023-06-19T05:32:10.459001Z","shell.execute_reply":"2023-06-19T05:32:10.459024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Result","metadata":{}},{"cell_type":"code","source":"plt.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":{"iopub.status.busy":"2023-06-19T05:32:10.460892Z","iopub.status.idle":"2023-06-19T05:32:10.461350Z","shell.execute_reply.started":"2023-06-19T05:32:10.461111Z","shell.execute_reply":"2023-06-19T05:32:10.461144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}