{"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":"import zipfile\n\ndirs = ['train.zip', 'train_masks.zip', 'train_masks.csv.zip']\n\nfor x in dirs:\n    with zipfile.ZipFile(\"../input/carvana-image-masking-challenge/\"+ x,'r') as z:\n        z.extractall(\"/kaggle/temp\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-05-05T13:31:32.385882Z","iopub.execute_input":"2022-05-05T13:31:32.386702Z","iopub.status.idle":"2022-05-05T13:31:43.88143Z","shell.execute_reply.started":"2022-05-05T13:31:32.3866Z","shell.execute_reply":"2022-05-05T13:31:43.880695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir \"./saved_images\"","metadata":{"execution":{"iopub.status.busy":"2022-05-05T14:27:18.152232Z","iopub.execute_input":"2022-05-05T14:27:18.152584Z","iopub.status.idle":"2022-05-05T14:27:18.932669Z","shell.execute_reply.started":"2022-05-05T14:27:18.152491Z","shell.execute_reply":"2022-05-05T14:27:18.931539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\n\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nimport torchvision\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader, Dataset\n\nimport torchmetrics\nfrom torchvision import models\nimport torchvision.transforms.functional as TF\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\n\nimport pytorch_lightning as pl\nfrom pytorch_lightning.loggers import WandbLogger\nfrom pytorch_lightning import Callback, LightningModule, Trainer, LightningDataModule\nfrom pytorch_lightning.callbacks import ModelCheckpoint, LearningRateMonitor\n\nfrom PIL import Image","metadata":{"execution":{"iopub.status.busy":"2022-05-05T13:31:43.882822Z","iopub.execute_input":"2022-05-05T13:31:43.883093Z","iopub.status.idle":"2022-05-05T13:31:50.68107Z","shell.execute_reply.started":"2022-05-05T13:31:43.883056Z","shell.execute_reply":"2022-05-05T13:31:50.680172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_DIR = Path('/') / \"kaggle\"/ \"temp\"","metadata":{"execution":{"iopub.status.busy":"2022-05-05T13:31:50.682595Z","iopub.execute_input":"2022-05-05T13:31:50.682849Z","iopub.status.idle":"2022-05-05T13:31:50.688458Z","shell.execute_reply.started":"2022-05-05T13:31:50.682811Z","shell.execute_reply":"2022-05-05T13:31:50.687763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(DATA_DIR / \"train_masks.csv\")\ndf[\"img_path\"] = df[\"img\"].apply(lambda x: str(DATA_DIR / \"train\" / x))\ndf[\"img_mask\"] = df[\"img\"].apply(lambda x: str(DATA_DIR / \"train_masks\" / x.replace(\".jpg\", \"_mask.gif\")))\n\ntrain_df, valid_df = train_test_split(df, test_size=0.2)","metadata":{"execution":{"iopub.status.busy":"2022-05-05T13:31:50.690319Z","iopub.execute_input":"2022-05-05T13:31:50.69097Z","iopub.status.idle":"2022-05-05T13:31:51.746622Z","shell.execute_reply.started":"2022-05-05T13:31:50.69093Z","shell.execute_reply":"2022-05-05T13:31:51.745939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CarvanaDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.transform = transform\n        self.df = df\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        img_path = df[\"img_path\"].iloc[index]\n        mask_path = df[\"img_mask\"].iloc[index]\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-05-05T13:33:22.400744Z","iopub.execute_input":"2022-05-05T13:33:22.401468Z","iopub.status.idle":"2022-05-05T13:33:22.408264Z","shell.execute_reply.started":"2022-05-05T13:33:22.401431Z","shell.execute_reply":"2022-05-05T13:33:22.407419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Hyperparameters etc.\nLEARNING_RATE = 1e-4\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nBATCH_SIZE = 4\nNUM_EPOCHS = 4\nNUM_WORKERS = 2\nIMAGE_HEIGHT = 576  # 1280 originally ,  640 \nIMAGE_WIDTH = 863  # 1918 originally , 959\nPIN_MEMORY = True\nLOAD_MODEL = False","metadata":{"execution":{"iopub.status.busy":"2022-05-05T13:33:22.576435Z","iopub.execute_input":"2022-05-05T13:33:22.576945Z","iopub.status.idle":"2022-05-05T13:33:22.638173Z","shell.execute_reply.started":"2022-05-05T13:33:22.576893Z","shell.execute_reply":"2022-05-05T13:33:22.637291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n\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)\n\n\nclass UNET(nn.Module):\n    def __init__(self, in_channels=3, out_channels=1, features=[64, 128, 256, 512]):\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        for feature in features:\n            self.downs.append(DoubleConv(in_channels, feature))\n            in_channels = feature\n\n        for feature in reversed(features):\n            self.ups.append(\n                nn.ConvTranspose2d(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-05-05T13:33:23.089187Z","iopub.execute_input":"2022-05-05T13:33:23.089593Z","iopub.status.idle":"2022-05-05T13:33:23.105522Z","shell.execute_reply.started":"2022-05-05T13:33:23.089561Z","shell.execute_reply":"2022-05-05T13:33:23.104791Z"},"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\"])\n\ndef get_loaders(\n    batch_size,\n    train_transform,\n    val_transform,\n    num_workers=4,\n    pin_memory=True,\n):\n    train_ds = CarvanaDataset(\n        df=train_df,\n        transform=train_transform,\n    )\n\n    train_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\n    val_ds = CarvanaDataset(\n        df=valid_df,\n        transform=val_transform,\n    )\n\n    val_loader = DataLoader(\n        val_ds,\n        batch_size=batch_size,\n        num_workers=num_workers,\n        pin_memory=pin_memory,\n        shuffle=False,\n    )\n\n    return train_loader, val_loader\n\ndef check_accuracy(loader, model, device=\"cuda\"):\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=\"./saved_images\", device=\"cuda\"\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-05-05T13:33:24.326661Z","iopub.execute_input":"2022-05-05T13:33:24.32719Z","iopub.status.idle":"2022-05-05T13:33:24.340219Z","shell.execute_reply.started":"2022-05-05T13:33:24.327149Z","shell.execute_reply":"2022-05-05T13:33:24.339221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fn(loader, model, optimizer, loss_fn, scaler):\n    loop = tqdm(loader)\n\n    for batch_idx, (data, targets) in enumerate(loop):\n        data = data.to(device=DEVICE)\n        targets = targets.float().unsqueeze(1).to(device=DEVICE)\n\n        # forward\n        with torch.cuda.amp.autocast():\n            predictions = model(data)\n            loss = loss_fn(predictions, targets)\n\n        # backward\n        if batch_idx % 8 == 0:\n            optimizer.zero_grad()\n            \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        # update tqdm loop\n        loop.set_postfix(loss=loss.item())","metadata":{"execution":{"iopub.status.busy":"2022-05-05T13:37:56.148605Z","iopub.execute_input":"2022-05-05T13:37:56.148887Z","iopub.status.idle":"2022-05-05T13:37:56.156068Z","shell.execute_reply.started":"2022-05-05T13:37:56.148855Z","shell.execute_reply":"2022-05-05T13:37:56.155299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main():\n    train_transform = 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\n    val_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    )\n\n    model = UNET(in_channels=3, out_channels=1).to(DEVICE)\n    loss_fn = nn.BCEWithLogitsLoss()\n    optimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE)\n\n    train_loader, val_loader = get_loaders(\n        BATCH_SIZE,\n        train_transform,\n        val_transforms,\n        NUM_WORKERS,\n        PIN_MEMORY,\n    )\n\n    if LOAD_MODEL:\n        load_checkpoint(torch.load(\"my_checkpoint.pth.tar\"), model)\n\n\n    check_accuracy(val_loader, model, device=DEVICE)\n    scaler = torch.cuda.amp.GradScaler()\n\n    for epoch in range(NUM_EPOCHS):\n        train_fn(train_loader, model, optimizer, loss_fn, scaler)\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        torch.cuda.empty_cache()\n        gc.collect()\n\n        # print some examples to a folder\n        save_predictions_as_imgs(\n            val_loader, model, folder=\"./saved_images\", device=DEVICE\n        )","metadata":{"execution":{"iopub.status.busy":"2022-05-05T13:37:57.59387Z","iopub.execute_input":"2022-05-05T13:37:57.594354Z","iopub.status.idle":"2022-05-05T13:37:57.606873Z","shell.execute_reply.started":"2022-05-05T13:37:57.594317Z","shell.execute_reply":"2022-05-05T13:37:57.606207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"main()","metadata":{"execution":{"iopub.status.busy":"2022-05-05T13:38:06.177206Z","iopub.execute_input":"2022-05-05T13:38:06.177457Z","iopub.status.idle":"2022-05-05T13:44:11.612015Z","shell.execute_reply.started":"2022-05-05T13:38:06.177428Z","shell.execute_reply":"2022-05-05T13:44:11.610778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}