{"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","execution":{"iopub.status.busy":"2022-11-16T09:06:03.677239Z","iopub.execute_input":"2022-11-16T09:06:03.678185Z","iopub.status.idle":"2022-11-16T09:06:03.694190Z","shell.execute_reply.started":"2022-11-16T09:06:03.678147Z","shell.execute_reply":"2022-11-16T09:06:03.692804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision\nimport torchvision.transforms.functional as TF\nfrom PIL import Image\nfrom torch.utils.data import Dataset\nimport torch.optim as optim\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nfrom torch.utils.data import DataLoader","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:18:39.404537Z","iopub.execute_input":"2022-11-16T09:18:39.404924Z","iopub.status.idle":"2022-11-16T09:18:39.411830Z","shell.execute_reply.started":"2022-11-16T09:18:39.404891Z","shell.execute_reply":"2022-11-16T09:18:39.410779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from zipfile import ZipFile\n\npath = \"../input/carvana-image-masking-challenge/\"\ndirs = ['train.zip', 'train_masks.zip']\ncdir = './'\n\nfor x in dirs:\n    with ZipFile(path+x, 'r') as zipf:\n        zipf.extractall(cdir)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:03.736085Z","iopub.execute_input":"2022-11-16T09:06:03.738246Z","iopub.status.idle":"2022-11-16T09:06:10.784836Z","shell.execute_reply.started":"2022-11-16T09:06:03.738218Z","shell.execute_reply":"2022-11-16T09:06:10.783444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Hyperparameter**","metadata":{}},{"cell_type":"code","source":"# Hyperparameters etc.\nLEARNING_RATE = 1e-4\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nBATCH_SIZE = 16\nNUM_EPOCHS = 3\nNUM_WORKERS = 2\nIMAGE_HEIGHT = 160 #1280 originally\nIMAGE_WIDTH = 240 #1918 originally\nPIN_MEMORY = True\nLOAD_MODEL = False\nTRAIN_IMG_DIR = \"./train\"\nTRAIN_MASK_DIR = \"./train_masks\"\n# VAL_IMG_DIR = \"data/val_images/\"\n# VAL_MASK_DIR = \"data/val_masks/\"","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:10.787134Z","iopub.execute_input":"2022-11-16T09:06:10.787500Z","iopub.status.idle":"2022-11-16T09:06:10.793775Z","shell.execute_reply.started":"2022-11-16T09:06:10.787463Z","shell.execute_reply":"2022-11-16T09:06:10.792802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_imgs = os.listdir(TRAIN_IMG_DIR)\ntrain_masks = os.listdir(TRAIN_MASK_DIR)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:10.794957Z","iopub.execute_input":"2022-11-16T09:06:10.795837Z","iopub.status.idle":"2022-11-16T09:06:11.011607Z","shell.execute_reply.started":"2022-11-16T09:06:10.795803Z","shell.execute_reply":"2022-11-16T09:06:11.010338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"train_transforms = A.Compose(\n        [\n            A.Resize(height=IMAGE_HEIGHT, width=IMAGE_WIDTH),\n            A.Rotate(limit=35, p=1.0),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.1),\n            A.Normalize(\n                mean=[0.0, 0.0, 0.0],\n                std=[1.0, 1.0, 1.0],\n                max_pixel_value=255.0,\n            ),\n            ToTensorV2(),\n        ],\n)\n\nval_transforms = A.Compose(\n        [\n            A.Resize(height=IMAGE_HEIGHT, width=IMAGE_WIDTH),\n            A.Normalize(\n                mean=[0.0, 0.0, 0.0],\n                std=[1.0, 1.0, 1.0],\n                max_pixel_value=255.0,\n            ),\n            ToTensorV2(),\n        ],\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:11.019251Z","iopub.execute_input":"2022-11-16T09:06:11.022353Z","iopub.status.idle":"2022-11-16T09:06:11.036518Z","shell.execute_reply.started":"2022-11-16T09:06:11.022305Z","shell.execute_reply":"2022-11-16T09:06:11.035465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CarvanaDataset(Dataset):\n    def __init__(self, images, image_dir, mask_dir, transform=None):\n        self.image_dir = image_dir\n        self.mask_dir = mask_dir\n        self.transform = transform\n        self.images = os.listdir(image_dir)\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, index):\n        img_path = os.path.join(self.image_dir, self.images[index])\n        mask_path = os.path.join(self.mask_dir, self.images[index].replace(\".jpg\", \"_mask.gif\"))\n        image = np.array(Image.open(img_path).convert(\"RGB\"))\n        mask = np.array(Image.open(mask_path).convert(\"L\"), dtype=np.float32)\n        mask[mask == 255.0] = 1.0\n\n        if self.transform is not None:\n            augmentations = self.transform(image=image, mask=mask)\n            image = augmentations[\"image\"]\n            mask = augmentations[\"mask\"]\n\n        return image, mask","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:11.043404Z","iopub.execute_input":"2022-11-16T09:06:11.046645Z","iopub.status.idle":"2022-11-16T09:06:11.056906Z","shell.execute_reply.started":"2022-11-16T09:06:11.046601Z","shell.execute_reply":"2022-11-16T09:06:11.055886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Split Dataset**","metadata":{}},{"cell_type":"code","source":"def train_test_split(images, splitSize=0.2):\n    imageLen = len(images)\n    val_len = int(splitSize*imageLen)\n    train_len = imageLen - val_len\n    train_images, val_images = images[:train_len], images[train_len:]\n    return train_images, val_images\n\ntrain_images, val_images = train_test_split(train_imgs)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:11.060093Z","iopub.execute_input":"2022-11-16T09:06:11.060815Z","iopub.status.idle":"2022-11-16T09:06:11.068453Z","shell.execute_reply.started":"2022-11-16T09:06:11.060730Z","shell.execute_reply":"2022-11-16T09:06:11.067476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_images)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:11.071693Z","iopub.execute_input":"2022-11-16T09:06:11.072006Z","iopub.status.idle":"2022-11-16T09:06:11.080599Z","shell.execute_reply.started":"2022-11-16T09:06:11.071982Z","shell.execute_reply":"2022-11-16T09:06:11.079641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(val_images)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:11.082067Z","iopub.execute_input":"2022-11-16T09:06:11.082641Z","iopub.status.idle":"2022-11-16T09:06:11.092444Z","shell.execute_reply.started":"2022-11-16T09:06:11.082597Z","shell.execute_reply":"2022-11-16T09:06:11.091334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = CarvanaDataset(\n        images = train_images,\n        image_dir= TRAIN_IMG_DIR,\n        mask_dir=TRAIN_MASK_DIR,\n        transform=train_transforms,\n)\n\ntrain_loader = DataLoader(\n        train_ds,\n        batch_size=BATCH_SIZE,\n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        shuffle=True,\n)\n\nval_ds = CarvanaDataset(\n        images = val_images,\n        image_dir=TRAIN_IMG_DIR,\n        mask_dir=TRAIN_MASK_DIR,\n        transform=val_transforms,\n)\n\nval_loader = DataLoader(\n        val_ds,\n        batch_size=BATCH_SIZE,\n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        shuffle=True,\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:11.094244Z","iopub.execute_input":"2022-11-16T09:06:11.094576Z","iopub.status.idle":"2022-11-16T09:06:11.108392Z","shell.execute_reply.started":"2022-11-16T09:06:11.094543Z","shell.execute_reply":"2022-11-16T09:06:11.107261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, 3, 1, 1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, 3, 1, 1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.conv(x)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:11.112240Z","iopub.execute_input":"2022-11-16T09:06:11.112568Z","iopub.status.idle":"2022-11-16T09:06:11.120227Z","shell.execute_reply.started":"2022-11-16T09:06:11.112530Z","shell.execute_reply":"2022-11-16T09:06:11.118717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UNET(nn.Module):\n    def __init__(\n            self, in_channels=3, out_channels=1, features=[64, 128, 256, 512],\n    ):\n        super(UNET, self).__init__()\n        self.ups = nn.ModuleList()\n        self.downs = nn.ModuleList()\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)\n\n        # Down part of UNET\n        for feature in features:\n            self.downs.append(DoubleConv(in_channels, feature))\n            in_channels = feature\n\n        # Up part of UNET\n        for feature in reversed(features):\n            self.ups.append(\n                nn.ConvTranspose2d(\n                    feature * 2, feature, kernel_size=2, stride=2,\n                )\n            )\n            self.ups.append(DoubleConv(feature * 2, feature))\n\n        self.bottleneck = DoubleConv(features[-1], features[-1] * 2)\n        self.final_conv = nn.Conv2d(features[0], out_channels, kernel_size=1)\n\n    def forward(self, x):\n        skip_connections = []\n\n        for down in self.downs:\n            x = down(x)\n            skip_connections.append(x)\n            x = self.pool(x)\n\n        x = self.bottleneck(x)\n        skip_connections = skip_connections[::-1]\n\n        for idx in range(0, len(self.ups), 2):\n            x = self.ups[idx](x)\n            skip_connection = skip_connections[idx // 2]\n\n            if x.shape != skip_connection.shape:\n                x = TF.resize(x, size=skip_connection.shape[2:])\n\n            concat_skip = torch.cat((skip_connection, x), dim=1)\n            x = self.ups[idx + 1](concat_skip)\n\n        return self.final_conv(x)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:11.121420Z","iopub.execute_input":"2022-11-16T09:06:11.121696Z","iopub.status.idle":"2022-11-16T09:06:11.135354Z","shell.execute_reply.started":"2022-11-16T09:06:11.121672Z","shell.execute_reply":"2022-11-16T09:06:11.134320Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test():\n    x = torch.randn((3, 1, 161, 161))\n    model = UNET(in_channels=1, out_channels=1)\n    preds = model(x)\n    print(preds.shape)\n    print(x.shape)\n    assert preds.shape == x.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:11.136922Z","iopub.execute_input":"2022-11-16T09:06:11.137377Z","iopub.status.idle":"2022-11-16T09:06:11.148529Z","shell.execute_reply.started":"2022-11-16T09:06:11.137249Z","shell.execute_reply":"2022-11-16T09:06:11.147499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test()","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:11.149769Z","iopub.execute_input":"2022-11-16T09:06:11.150244Z","iopub.status.idle":"2022-11-16T09:06:13.319216Z","shell.execute_reply.started":"2022-11-16T09:06:11.150210Z","shell.execute_reply":"2022-11-16T09:06:13.318175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utility","metadata":{}},{"cell_type":"code","source":"# save_dir = \"/kaggle/working/saved_images\"\n# os.mkdir(save_dir)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:13.320910Z","iopub.execute_input":"2022-11-16T09:06:13.321565Z","iopub.status.idle":"2022-11-16T09:06:13.326149Z","shell.execute_reply.started":"2022-11-16T09:06:13.321525Z","shell.execute_reply":"2022-11-16T09:06:13.324930Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_checkpoint(state, filename=\"my_checkpoint.pth.tar\"):\n    print(\"=> Saving checkpoint\")\n    torch.save(state, filename)\n\ndef load_checkpoint(checkpoint, model):\n    print(\"=> Loading checkpoint\")\n    model.load_state_dict(checkpoint[\"state_dict\"])","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:13.327583Z","iopub.execute_input":"2022-11-16T09:06:13.328202Z","iopub.status.idle":"2022-11-16T09:06:13.337985Z","shell.execute_reply.started":"2022-11-16T09:06:13.328166Z","shell.execute_reply":"2022-11-16T09:06:13.337040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def check_accuracy(loader, model, device=DEVICE):\n    num_correct = 0\n    num_pixels = 0\n    dice_score = 0\n    model.eval()\n\n    with torch.no_grad():\n        for x, y in loader:\n            x = x.to(device)\n            y = y.to(device).unsqueeze(1)\n            preds = torch.sigmoid(model(x))\n            preds = (preds > 0.5).float()\n            num_correct += (preds == y).sum()\n            num_pixels += torch.numel(preds)\n            dice_score += (2 * (preds * y).sum()) / (\n                (preds+y).sum() + 1e-8\n            )\n\n    print(\n        f\"Got {num_correct}/{num_pixels} with acc {num_correct/num_pixels*100:.2f}\"\n    )\n    print(f\"Dice score: {dice_score/len(loader)}\")\n    model.train()\n\ndef save_predictions_as_imgs(\n    loader, model, folder=save_dir, device=DEVICE\n):\n    model.eval()\n    for idx, (x, y) in enumerate(loader):\n        x = x.to(device=device)\n        with torch.no_grad():\n            preds = torch.sigmoid(model(x))\n            preds = (preds > 0.5).float()\n        torchvision.utils.save_image(\n            preds, f\"{folder}/pred_{idx}.png\"\n        )\n        torchvision.utils.save_image(y.unsqueeze(1), f\"{folder}{idx}.png\")\n\n    model.train()","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:13.339719Z","iopub.execute_input":"2022-11-16T09:06:13.340326Z","iopub.status.idle":"2022-11-16T09:06:13.354429Z","shell.execute_reply.started":"2022-11-16T09:06:13.340291Z","shell.execute_reply":"2022-11-16T09:06:13.353475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"def fit(model, dataloader, data, optimizer, criterion):\n    print(\"TRAINING: \")\n    model.train()\n    train_running_loss = 0.0\n    counter = 0\n    \n    #num of batches\n    num_batches = int(len(data)/dataloader.batch_size)\n    for i, data in tqdm(enumerate(dataloader), total=num_batches):\n        counter += 1\n        image, mask = data[0].to(DEVICE), data[1].to(DEVICE)\n        optimizer.zero_grad()\n        outputs = model(image)\n        outputs = outputs.squeeze(1)\n        loss = criterion(outputs, mask)\n        train_running_loss += loss.item()\n        loss.backward()\n        optimizer.step()\n        \n    train_loss = train_running_loss/counter\n    return train_loss\n\ndef validate(model, dataloader, data, criterion):\n    print(\"VALIDATION:\")\n    model.eval()\n    valid_running_loss = 0.0\n    counter = 0\n    \n    #num of batches\n    num_batches = int(len(data)/dataloader.batch_size)\n    with torch.no_grad():\n        for i, data in tqdm(enumerate(dataloader), total=num_batches):\n            counter += 1\n            image, mask = data[0].to(DEVICE), data[1].to(DEVICE)\n            outputs = model(image)\n            outputs = outputs.squeeze(1)\n            loss = criterion(outputs, mask)\n            valid_running_loss += loss.item()\n    valid_loss = valid_running_loss/counter\n    return valid_loss","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:06:13.357909Z","iopub.execute_input":"2022-11-16T09:06:13.358265Z","iopub.status.idle":"2022-11-16T09:06:13.368995Z","shell.execute_reply.started":"2022-11-16T09:06:13.358239Z","shell.execute_reply":"2022-11-16T09:06:13.367894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNET(in_channels=3, out_channels=1).to(DEVICE)\n# loss_fn = nn.BCEWithLogitsLoss()\noptimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE)\ncriterion = nn.BCEWithLogitsLoss()\n\nif LOAD_MODEL:\n    load_checkpoint(torch.load(\"my_checkpoint.pth.tar\"), model)\n\ntrain_loss = []\nval_loss = []\nfor epoch in range(NUM_EPOCHS):\n    #training and validation\n    print(f\"Epoch: {epoch+1} of {NUM_EPOCHS}\")\n    train_epoch_loss = fit(model, train_loader, train_ds, optimizer, criterion)\n    val_epoch_loss = validate(model, val_loader, val_ds, criterion)\n    train_loss.append(train_epoch_loss)\n    val_loss.append(val_epoch_loss)\n    print(f\"Train Loss: {train_epoch_loss:.4f}\")\n    print(f\"Val Loss: {val_epoch_loss:.4f}\")\n    \n    # save model\n    checkpoint = {\n        \"state_dict\": model.state_dict(),\n        \"optimizer\": optimizer.state_dict(),\n    }\n    save_checkpoint(checkpoint)\n\n    # check accuracy\n    check_accuracy(val_loader, model, device=DEVICE)\n\n    # print some examples to a folder\n    save_predictions_as_imgs(\n        val_loader, model, folder=save_dir, device=DEVICE\n    )\n\n","metadata":{"execution":{"iopub.status.busy":"2022-11-16T09:18:45.259270Z","iopub.execute_input":"2022-11-16T09:18:45.259643Z","iopub.status.idle":"2022-11-16T10:02:08.253016Z","shell.execute_reply.started":"2022-11-16T09:18:45.259609Z","shell.execute_reply":"2022-11-16T10:02:08.251835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}