{"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":"markdown","source":"# Simple Unet Baseline (Train)\n\nThis is the training part of the two part Unet Baseline for this competition.\n#### Inference Notebook: [Simple Unet Baseline (Infer)][1]. \nYou can find the notebook to create the dataset used for training [here][2] to get a better understanding of how everything works.\n* Smp library is used to get the unet model.\n* EfficientNetB0 is used as the backbone initialized on imagenet weight.\n* Ash color images are used for training (With only the labeled frames and human_pixel_masks.\n* Custom implementation of dice score is used according to this competition.\n* After training, we find the best threshold for the valid set, which will then be used for the submission.\n* Wandb can also be used with this notebook to log experiments, just uncomment the wandb code snippets.\n\n**Version 5** Updates:\n* Added some Augmentations\n* Trained for more Epochs\n* Option to increase image size\n\n### Please upvote if you find this useful.\n\n[1]: https://www.kaggle.com/code/shashwatraman/simple-unet-pytorch-baseline-infer\n[2]: https://www.kaggle.com/code/shashwatraman/contrails-dataset-ash-color/notebook","metadata":{}},{"cell_type":"markdown","source":"## Import Libraries","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nimport os\nimport random\nimport math\nfrom collections import defaultdict\nimport cv2\nimport skimage\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport torch\nfrom torch import nn\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nimport torch.nn.functional as F\n\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nfrom transformers import get_cosine_schedule_with_warmup\n\ntorch.__version__","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-05T10:34:12.365593Z","iopub.execute_input":"2023-06-05T10:34:12.367963Z","iopub.status.idle":"2023-06-05T10:34:24.992562Z","shell.execute_reply.started":"2023-06-05T10:34:12.367922Z","shell.execute_reply":"2023-06-05T10:34:24.9916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install segmentation-models-pytorch\nimport segmentation_models_pytorch as smp","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-05T10:34:24.994235Z","iopub.execute_input":"2023-06-05T10:34:24.994864Z","iopub.status.idle":"2023-06-05T10:34:44.062043Z","shell.execute_reply.started":"2023-06-05T10:34:24.994835Z","shell.execute_reply":"2023-06-05T10:34:44.060929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install -qU wandb\n# import wandb\n# wandb.login(key='')","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:44.063822Z","iopub.execute_input":"2023-06-05T10:34:44.064208Z","iopub.status.idle":"2023-06-05T10:34:44.070232Z","shell.execute_reply.started":"2023-06-05T10:34:44.064172Z","shell.execute_reply":"2023-06-05T10:34:44.068646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Preparation","metadata":{}},{"cell_type":"code","source":"class Config:\n    train = True\n    train_aug=True\n    \n    num_epochs = 30\n    num_classes = 1\n    batch_size = 32\n    seed = 42\n    \n    encoder = 'efficientnet-b3'\n    pretrained = True\n    weights = 'imagenet'\n    classes = ['contrail']\n    activation = None\n    in_chans = 3\n    \n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    image_size = 256\n    warmup = 0\n    lr = 3e-3\n    \nclass Paths:\n    data_root = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\n    contrails = '/kaggle/input/contrails-images-ash-color/contrails/'\n    train_path = '/kaggle/input/contrails-images-ash-color/train_df.csv'\n    valid_path = '/kaggle/input/contrails-images-ash-color/valid_df.csv'","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:44.073519Z","iopub.execute_input":"2023-06-05T10:34:44.073871Z","iopub.status.idle":"2023-06-05T10:34:44.107297Z","shell.execute_reply.started":"2023-06-05T10:34:44.073839Z","shell.execute_reply":"2023-06-05T10:34:44.106195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=1234):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    \n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:44.108615Z","iopub.execute_input":"2023-06-05T10:34:44.109407Z","iopub.status.idle":"2023-06-05T10:34:44.117248Z","shell.execute_reply.started":"2023-06-05T10:34:44.109375Z","shell.execute_reply":"2023-06-05T10:34:44.116446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set_seed(9)","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:44.118514Z","iopub.execute_input":"2023-06-05T10:34:44.119025Z","iopub.status.idle":"2023-06-05T10:34:44.131539Z","shell.execute_reply.started":"2023-06-05T10:34:44.118994Z","shell.execute_reply":"2023-06-05T10:34:44.130807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Import dataframes\ntrain_df = pd.read_csv(Paths.train_path)\nvalid_df = pd.read_csv(Paths.valid_path)\n\ntrain_df['path'] = Paths.contrails + train_df['record_id'].astype(str) + '.npy'\nvalid_df['path'] = Paths.contrails + valid_df['record_id'].astype(str) + '.npy'\n\ntrain_df.shape, valid_df.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:44.132997Z","iopub.execute_input":"2023-06-05T10:34:44.134078Z","iopub.status.idle":"2023-06-05T10:34:44.200857Z","shell.execute_reply.started":"2023-06-05T10:34:44.134047Z","shell.execute_reply":"2023-06-05T10:34:44.199608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform_size = A.Compose([\n    A.Resize(Config.image_size, Config.image_size, interpolation=cv2.INTER_LANCZOS4, always_apply=True)\n])\n\ntrain_transform = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.RandomResizedCrop(height=256, width=256, scale=(0.75, 1.0), p=0.6)\n])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, train=True, transform=None):\n        \n        self.df = df\n        self.trn = train\n        self.transform = transform\n    \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        con_path = row.path\n        con = np.load(str(con_path))\n        \n        img = con[..., :-1]\n        label = con[..., -1]\n        \n        img = img.astype(np.float32)\n        label = label.astype(np.float32)\n        \n        if Config.train_aug:\n            if self.transform is not None:\n                augmented = self.transform(image=img, mask=label)\n                img = augmented['image']\n                label = augmented['mask']\n                \n        if Config.image_size != 256:\n            img = transform_size(image=img)[\"image\"]\n        \n        img = torch.tensor(img)\n        label = torch.tensor(label)\n        \n        img = img.permute(2, 0, 1)\n            \n        return img.float(), label.float()\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:44.202554Z","iopub.execute_input":"2023-06-05T10:34:44.202926Z","iopub.status.idle":"2023-06-05T10:34:44.210458Z","shell.execute_reply.started":"2023-06-05T10:34:44.202892Z","shell.execute_reply":"2023-06-05T10:34:44.209547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = ContrailsDataset(\n        train_df,\n        train=True,\n        transform=train_transform\n    )\n\nvalid_ds = ContrailsDataset(\n        valid_df,\n        train=False,\n        transform=None\n    )\n\ntrain_dl = DataLoader(train_ds, batch_size=Config.batch_size , shuffle=True, num_workers = 2)    \nvalid_dl = DataLoader(valid_ds, batch_size=Config.batch_size, num_workers = 2)","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:44.211821Z","iopub.execute_input":"2023-06-05T10:34:44.212791Z","iopub.status.idle":"2023-06-05T10:34:44.223459Z","shell.execute_reply.started":"2023-06-05T10:34:44.212758Z","shell.execute_reply":"2023-06-05T10:34:44.222488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, label = next(iter(train_dl))\nimg.shape, label.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:44.228871Z","iopub.execute_input":"2023-06-05T10:34:44.229448Z","iopub.status.idle":"2023-06-05T10:34:45.186963Z","shell.execute_reply.started":"2023-06-05T10:34:44.229401Z","shell.execute_reply":"2023-06-05T10:34:45.185628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, label = next(iter(valid_dl))\nimg.shape, label.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:45.188788Z","iopub.execute_input":"2023-06-05T10:34:45.189501Z","iopub.status.idle":"2023-06-05T10:34:46.12151Z","shell.execute_reply.started":"2023-06-05T10:34:45.189463Z","shell.execute_reply":"2023-06-05T10:34:46.120275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def display_random_images(dataset, n=10, seed=None):\n    if seed:\n        random.seed(seed)\n    random_samples_idx = random.sample(range(len(dataset)), k=n)\n    plt.figure(figsize=(30, 20))\n    \n    for i, targ_sample in enumerate(random_samples_idx):\n        targ_image, targ_label = dataset[targ_sample][0], dataset[targ_sample][1]\n        \n        targ_image = targ_image.permute(1, 2, 0)\n        \n        plt.subplot(1, n, i+1)\n        plt.imshow(targ_image)\n        plt.axis(False)","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:46.12381Z","iopub.execute_input":"2023-06-05T10:34:46.124548Z","iopub.status.idle":"2023-06-05T10:34:46.132319Z","shell.execute_reply.started":"2023-06-05T10:34:46.124505Z","shell.execute_reply":"2023-06-05T10:34:46.131048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_random_images(train_ds, 4, 42)","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:46.134347Z","iopub.execute_input":"2023-06-05T10:34:46.135193Z","iopub.status.idle":"2023-06-05T10:34:47.745094Z","shell.execute_reply.started":"2023-06-05T10:34:46.135159Z","shell.execute_reply":"2023-06-05T10:34:47.744196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_random_images(valid_ds, 4, 42)","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:47.746092Z","iopub.execute_input":"2023-06-05T10:34:47.746434Z","iopub.status.idle":"2023-06-05T10:34:49.049693Z","shell.execute_reply.started":"2023-06-05T10:34:47.746403Z","shell.execute_reply":"2023-06-05T10:34:49.048838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def dice_coef(y_true, y_pred, thr=0.5, epsilon=0.001):\n    y_true = y_true.flatten()\n    y_pred = (y_pred>thr).astype(np.float32).flatten()\n    inter = (y_true*y_pred).sum()\n    den = y_true.sum() + y_pred.sum()\n    dice = ((2*inter+epsilon)/(den+epsilon))\n    return dice","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:49.051097Z","iopub.execute_input":"2023-06-05T10:34:49.05173Z","iopub.status.idle":"2023-06-05T10:34:49.058513Z","shell.execute_reply.started":"2023-06-05T10:34:49.051674Z","shell.execute_reply":"2023-06-05T10:34:49.057777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, cfg):\n        super(UNet, self).__init__()\n        \n        self.cfg = cfg\n        self.training = True\n        \n        self.model = smp.Unet(\n            encoder_name=cfg.encoder, \n            encoder_weights=cfg.weights, \n            decoder_use_batchnorm=True,\n            classes=len(cfg.classes), \n            activation=cfg.activation,\n        )\n        \n        self.loss_fn = smp.losses.DiceLoss(mode='binary')\n    \n    def forward(self, imgs, targets):\n        \n        x = imgs\n        y = targets\n\n        logits = self.model(x)\n        \n        if Config.image_size != 256:\n            logits = F.interpolate(logits, size=(256, 256), mode='nearest-exact')\n        \n        loss = self.loss_fn(logits, y)\n        \n        return {\"loss\": loss, \"logits\": logits.sigmoid(), \"logits_raw\": logits, \"target\": y}","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:49.060043Z","iopub.execute_input":"2023-06-05T10:34:49.061104Z","iopub.status.idle":"2023-06-05T10:34:49.069799Z","shell.execute_reply.started":"2023-06-05T10:34:49.061061Z","shell.execute_reply":"2023-06-05T10:34:49.068934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_step(model, dataloader, optimizer, device):\n    \n    model.train()\n    \n    train_losses = []\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    \n    for step, (X, y) in pbar:\n        \n        X, y = X.to(device), y.to(device)\n        torch.set_grad_enabled(True)\n        \n        output_dict = model(X, y)\n        loss = output_dict[\"loss\"]\n        train_losses.append(loss.item())\n        \n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n        \n        if scheduler is not None:\n            scheduler.step()\n    \n    train_loss = np.sum(train_losses)\n    \n    return train_loss","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:49.071087Z","iopub.execute_input":"2023-06-05T10:34:49.071809Z","iopub.status.idle":"2023-06-05T10:34:49.084879Z","shell.execute_reply.started":"2023-06-05T10:34:49.071775Z","shell.execute_reply":"2023-06-05T10:34:49.084026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_step(model, dataloader, device):\n    \n    model.eval()\n    torch.set_grad_enabled(False)\n    \n    val_data = defaultdict(list)\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid')\n    for step, (X, y) in pbar: \n        X, y = X.to(device), y.to(device)\n\n        output = model(X, y)\n        for key, val in output.items():\n            val_data[key] += [output[key]]\n\n    for key, val in output.items():\n        value = val_data[key]\n        if len(value[0].shape) == 0:\n            val_data[key] = torch.stack(value)\n        else:\n            val_data[key] = torch.cat(value, dim=0).cpu().detach().numpy()\n    \n    val_losses = val_data[\"loss\"].cpu().numpy()\n    val_loss = np.sum(val_losses)\n    \n    val_dice = dice_coef(val_data['target'], val_data['logits'])\n    \n    return val_loss, val_dice","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:49.086236Z","iopub.execute_input":"2023-06-05T10:34:49.086908Z","iopub.status.idle":"2023-06-05T10:34:49.098901Z","shell.execute_reply.started":"2023-06-05T10:34:49.086873Z","shell.execute_reply":"2023-06-05T10:34:49.097917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.auto import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:49.10028Z","iopub.execute_input":"2023-06-05T10:34:49.101194Z","iopub.status.idle":"2023-06-05T10:34:49.110455Z","shell.execute_reply.started":"2023-06-05T10:34:49.101113Z","shell.execute_reply":"2023-06-05T10:34:49.109438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, train_dataloader, test_dataloader, optimizer, epochs, device):\n    results = {'train_loss': [],\n              'val_loss': [],\n              'val_dice': []}\n    for epoch in range(epochs):\n        \n        set_seed(Config.seed + epoch)\n        print(\"EPOCH:\", epoch)\n        \n        train_loss = train_step(model,\n                              train_dataloader,\n                              optimizer,\n                              device)\n        val_loss, val_dice = test_step(model,\n                            test_dataloader,\n                            device)\n        \n        train_loss = train_loss / len(train_ds)\n        val_loss = val_loss / len(valid_ds)\n        \n        print(f'Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | Val Dice: {val_dice:.4f}')\n        print(f\"Learning rate: {optimizer.param_groups[0]['lr']}\")\n        \n        results['train_loss'].append(train_loss)\n        results['val_loss'].append(val_loss)\n        results['val_dice'].append(val_dice)\n        \n#         wandb.log({\n#         \"Train Loss\": train_loss,\n#         \"Valid Loss\": val_loss,\n#         'Valid Dice': val_dice})\n        \n        PATH = f\"epoch-{epoch}.pth\"\n        torch.save(model.state_dict(), PATH)\n        \n#         wandb.save(PATH)\n\n    return results","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:49.112422Z","iopub.execute_input":"2023-06-05T10:34:49.112846Z","iopub.status.idle":"2023-06-05T10:34:49.123667Z","shell.execute_reply.started":"2023-06-05T10:34:49.112813Z","shell.execute_reply":"2023-06-05T10:34:49.122598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_optimizer(lr, params):\n    \n    model_optimizer = torch.optim.Adam(\n            filter(lambda p: p.requires_grad, params), \n            lr=lr,\n            weight_decay=0)\n    \n    return model_optimizer","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:49.125318Z","iopub.execute_input":"2023-06-05T10:34:49.125674Z","iopub.status.idle":"2023-06-05T10:34:49.136092Z","shell.execute_reply.started":"2023-06-05T10:34:49.125643Z","shell.execute_reply":"2023-06-05T10:34:49.134993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_scheduler(cfg, optimizer, total_steps):\n    scheduler = get_cosine_schedule_with_warmup(\n        optimizer,\n        num_warmup_steps= cfg.warmup * (total_steps // cfg.batch_size),\n        num_training_steps= cfg.num_epochs * (total_steps // cfg.batch_size)\n    )\n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:49.137825Z","iopub.execute_input":"2023-06-05T10:34:49.138156Z","iopub.status.idle":"2023-06-05T10:34:49.146572Z","shell.execute_reply.started":"2023-06-05T10:34:49.138127Z","shell.execute_reply":"2023-06-05T10:34:49.145607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_EPOCHS = Config.num_epochs\nmodel = UNet(Config).to(Config.device)\n\n# run = wandb.init(project='Google Contrails', \n#                      config={k:v for k, v in dict(vars(Config)).items() if '__' not in k},\n#                      name=f\"{Config.encoder}-{Config.num_epochs}epos-{Config.lr}-unet\"\n#                     )\n\ntotal_steps = len(train_ds)\noptimizer = get_optimizer(lr=Config.lr, params=model.parameters())\nscheduler = get_scheduler(Config, optimizer, total_steps)\n\n# wandb.watch(model, log_freq=100, log='all')\n\nfrom timeit import default_timer as timer\nstart_time = timer()\n\nmodel_results = train(model, train_dl, valid_dl, optimizer, NUM_EPOCHS, Config.device)\n\nend_time = timer()\n\n# run.finish()\nprint(f'Total Training Time: {end_time-start_time:.3f} seconds')","metadata":{"execution":{"iopub.status.busy":"2023-06-05T10:34:49.148294Z","iopub.execute_input":"2023-06-05T10:34:49.148778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Finding the Best Threshold","metadata":{}},{"cell_type":"code","source":"# Predicting the Valid Set\nmodel.eval()\ntorch.set_grad_enabled(False)\n\nval_data = defaultdict(list)\npbar = tqdm(enumerate(valid_dl), total=len(valid_dl), desc='Valid')\nfor step, (X, y) in pbar: \n    X, y = X.to(Config.device), y.to(Config.device)\n\n    output = model(X, y)\n    for key, val in output.items():\n        val_data[key] += [output[key]]\n\nfor key, val in output.items():\n    value = val_data[key]\n    if len(value[0].shape) == 0:\n        val_data[key] = torch.stack(value)\n    else:\n        val_data[key] = torch.cat(value, dim=0).cpu().detach().numpy()\n\nval_losses = val_data[\"loss\"].cpu().numpy()\nval_loss = np.sum(val_losses)\nval_loss = val_loss / len(valid_ds)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = val_data['logits']\nground_truths = val_data['target']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions.shape, ground_truths.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Finding the Best Threshold\nbdice = -1\nbi = None\nfor i in tqdm(np.arange(0, 1.01, 0.01)):\n    val_dice = dice_coef(ground_truths, predictions, i)\n    if val_dice > bdice:\n        bdice = val_dice\n        bi = i","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'Best Threshold: {bi}')\nprint(f'Best Validation Dice Score: {bdice}')","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}