{"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-02-23T11:00:33.529115Z","iopub.execute_input":"2023-02-23T11:00:33.529604Z","iopub.status.idle":"2023-02-23T11:00:33.561753Z","shell.execute_reply.started":"2023-02-23T11:00:33.529527Z","shell.execute_reply":"2023-02-23T11:00:33.560951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"In this nootbok, I demostrated how to use UNet to perform the binary image segmentation task.\n\n## Steps:\n- Import Libaries\n- Load Data\n- Visualize the Sample Data\n- Define UNet Model\n- Training\n- Visualize Result","metadata":{}},{"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-02-23T11:00:52.031183Z","iopub.execute_input":"2023-02-23T11:00:52.031534Z","iopub.status.idle":"2023-02-23T11:01:12.674639Z","shell.execute_reply.started":"2023-02-23T11:00:52.031504Z","shell.execute_reply":"2023-02-23T11:01:12.673218Z"},"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-02-23T11:01:12.679501Z","iopub.execute_input":"2023-02-23T11:01:12.680534Z","iopub.status.idle":"2023-02-23T11:01:13.144397Z","shell.execute_reply.started":"2023-02-23T11:01:12.680489Z","shell.execute_reply":"2023-02-23T11:01:13.143446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"batch_size = 4\nn_iters = 10000\nepochs = 10\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-02-23T11:01:13.146017Z","iopub.execute_input":"2023-02-23T11:01:13.146374Z","iopub.status.idle":"2023-02-23T11:01:13.270514Z","shell.execute_reply.started":"2023-02-23T11:01:13.146338Z","shell.execute_reply":"2023-02-23T11:01:13.269147Z"},"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')\n    \n# # extract test data\n# with ZipFile('../input/carvana-image-masking-challenge/test.zip', 'r') as zf:\n#     zf.extractall('../working') ","metadata":{"execution":{"iopub.status.busy":"2023-02-23T11:01:13.272902Z","iopub.execute_input":"2023-02-23T11:01:13.273426Z","iopub.status.idle":"2023-02-23T11:01:24.505983Z","shell.execute_reply.started":"2023-02-23T11:01:13.273386Z","shell.execute_reply":"2023-02-23T11:01:24.504920Z"},"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-02-05T10:33:28.326733Z","iopub.execute_input":"2023-02-05T10:33:28.327185Z","iopub.status.idle":"2023-02-05T10:33:28.41265Z","shell.execute_reply.started":"2023-02-05T10:33:28.327147Z","shell.execute_reply":"2023-02-05T10:33:28.411582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"define data loader and manualy split them into train/valid dataset","metadata":{}},{"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-02-05T10:33:28.41446Z","iopub.execute_input":"2023-02-05T10:33:28.414897Z","iopub.status.idle":"2023-02-05T10:33:30.411006Z","shell.execute_reply.started":"2023-02-05T10:33:28.414855Z","shell.execute_reply":"2023-02-05T10:33:30.409738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Sampele images of the dataset","metadata":{}},{"cell_type":"markdown","source":"# Define Netowrk Model","metadata":{}},{"cell_type":"markdown","source":"![U-Net](https://media.springernature.com/lw685/springer-static/image/art%3A10.1007%2Fs10462-022-10152-1/MediaObjects/10462_2022_10152_Fig3_HTML.png?as=webp)","metadata":{}},{"cell_type":"code","source":"class ConvBlock(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int):\n        super(ConvBlock, self).__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            \n            nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n        \n    def forward(self, x: torch.Tensor):\n        return self.block(x)\n    \nclass CopyAndCrop(nn.Module):\n    def forward(self, x: torch.Tensor, encoded: torch.Tensor):\n        _, _, h, w = encoded.shape\n        crop = T.CenterCrop((h, w))(x)\n        output = torch.cat((x, crop), 1)\n        \n        return output\n\n    \nclass UNet(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int):\n        super(UNet, self).__init__()\n        \n        self.encoders = nn.ModuleList([\n            ConvBlock(in_channels, 64),\n            ConvBlock(64, 128),\n            ConvBlock(128, 256),\n            ConvBlock(256, 512),\n        ])\n        self.down_sample = nn.MaxPool2d(2)\n        self.copyAndCrop = CopyAndCrop()\n        self.decoders = nn.ModuleList([\n            ConvBlock(1024, 512),\n            ConvBlock(512, 256),\n            ConvBlock(256, 128),\n            ConvBlock(128, 64),\n        ])\n        \n        # PixelShuffle, UpSample will modify the output channel (you can add extra operation to update the channel, e.g.conv2d)\n        # preffer use convTranspose2d, it won't modify the output channel\n        self.up_samples = nn.ModuleList([\n            nn.ConvTranspose2d(1024, 512, kernel_size=2, stride=2),\n            nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2),\n            nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2),\n            nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)\n        ])\n\n        \n        self.bottleneck = ConvBlock(512, 1024)\n        self.final_conv = nn.Conv2d(64, out_channels, kernel_size=1, stride=1)\n        \n    def forward(self, x: torch.Tensor):\n        # encod\n        encoded_features = []\n        for enc in self.encoders:\n            x = enc(x)\n            encoded_features.append(x)\n            x = self.down_sample(x)\n            \n        \n        x = self.bottleneck(x)\n        \n        # decode\n        for idx, denc in enumerate(self.decoders):\n            x = self.up_samples[idx](x)\n            encoded = encoded_features.pop()\n            x = self.copyAndCrop(x, encoded)\n            x = denc(x)\n            \n        output = self.final_conv(x)\n        return output\n        \n","metadata":{"execution":{"iopub.status.busy":"2023-02-05T10:33:30.413238Z","iopub.execute_input":"2023-02-05T10:33:30.413813Z","iopub.status.idle":"2023-02-05T10:33:30.43068Z","shell.execute_reply.started":"2023-02-05T10:33:30.41375Z","shell.execute_reply":"2023-02-05T10:33:30.429794Z"},"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    pass\n\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","metadata":{"execution":{"iopub.status.busy":"2023-02-05T10:33:30.432427Z","iopub.execute_input":"2023-02-05T10:33:30.43352Z","iopub.status.idle":"2023-02-05T10:33:30.446665Z","shell.execute_reply.started":"2023-02-05T10:33:30.433481Z","shell.execute_reply":"2023-02-05T10:33:30.445761Z"},"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-02-05T10:33:30.448037Z","iopub.execute_input":"2023-02-05T10:33:30.448471Z","iopub.status.idle":"2023-02-05T10:33:30.460723Z","shell.execute_reply.started":"2023-02-05T10:33:30.448436Z","shell.execute_reply":"2023-02-05T10:33:30.459663Z"},"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\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\n\n            val_epoch_loss = val_total_loss / val_iterations\n            dice_coef_scroe = scores / 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: {}'.format(\n                epoch+1, epochs,\n                idx+1, len(train_dataset_loader),\n                dice_coef_scroe, train_epoch_loss, val_epoch_loss\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    }\n            ","metadata":{"execution":{"iopub.status.busy":"2023-02-05T10:33:30.465169Z","iopub.execute_input":"2023-02-05T10:33:30.465455Z","iopub.status.idle":"2023-02-05T10:33:30.479929Z","shell.execute_reply.started":"2023-02-05T10:33:30.46543Z","shell.execute_reply":"2023-02-05T10:33:30.478922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet(in_channels=3, out_channels=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-02-05T10:33:30.481789Z","iopub.execute_input":"2023-02-05T10:33:30.482842Z","iopub.status.idle":"2023-02-05T11:24:31.90546Z","shell.execute_reply.started":"2023-02-05T10:33:30.482808Z","shell.execute_reply":"2023-02-05T11:24:31.904385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet(in_channels=3, out_channels=1).to(device)\ncriterion = nn.BCEWithLogitsLoss()\n# criterion = smp.losses.DiceLoss(mode='binary')\noptimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)\n\n# try to find the best learning rate\nscheduler = StepLR(optimizer, step_size=2, gamma=0.1)\nhistory2 = train(model, optimizer, criterion, scheduler)","metadata":{"execution":{"iopub.status.busy":"2023-02-05T11:24:31.90775Z","iopub.execute_input":"2023-02-05T11:24:31.909033Z","iopub.status.idle":"2023-02-05T12:15:20.69979Z","shell.execute_reply.started":"2023-02-05T11:24:31.908986Z","shell.execute_reply":"2023-02-05T12:15:20.698701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Experiments\n\nAfter few experiments, I found *DiceLoss* is more stable than *BCEWithLogitsLoss* if the batch size is large (batch size: 8, 10, 32). *BCEWithLogitsLoss* will output *empty* result (all nan) after 9 epochs.(**version 6** of this notebook). Is it the gradient decent issue?\n\nIf batch size are 4, the *BCEWithLogitsLoss* perform better than *DiceLoss*.\n\nAnd, is the calculation method of the dice score incorrect? It looks weird.","metadata":{}},{"cell_type":"markdown","source":"# Visualize 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')\n\nplt.plot(history2['train_loss'], label='schdule LR train loss')\nplt.plot(history2['valid_loss'], label='schdule LR val loss')\n\nplt.xlabel('epoch')\nplt.ylabel('loss')\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-05T12:15:30.371458Z","iopub.execute_input":"2023-02-05T12:15:30.371868Z","iopub.status.idle":"2023-02-05T12:15:30.607562Z","shell.execute_reply.started":"2023-02-05T12:15:30.371816Z","shell.execute_reply":"2023-02-05T12:15:30.606632Z"},"trusted":true},"execution_count":null,"outputs":[]}]}