{"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":"# Setup","metadata":{}},{"cell_type":"markdown","source":"## Library Imports","metadata":{}},{"cell_type":"code","source":"# Adding paths to segmentation libraries\nimport sys\nsys.path.extend([\n    \"../input/pretrained-models-pytorch\",\n    \"../input/efficientnet-pytorch\",\n    \"../input/smp-github/segmentation_models.pytorch-master\",\n    \"/kaggle/input/timm-pretrained-resnest/resnest/\"\n])\n\n# Deep learning, computer vision and data handling\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn.modules.loss import _Loss\nimport torchvision.transforms as T\nimport segmentation_models_pytorch as smp\nimport numpy as np\nfrom timm.scheduler import CosineLRScheduler\nimport albumentations as A\nimport cv2\nimport pandas as pd\n\n# Visualization and logging\nimport wandb\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\n\n# Standard\nimport os\nimport time\nimport copy\nfrom typing import Optional, List","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:11:22.874501Z","iopub.execute_input":"2023-08-25T09:11:22.874913Z","iopub.status.idle":"2023-08-25T09:11:29.265955Z","shell.execute_reply.started":"2023-08-25T09:11:22.87488Z","shell.execute_reply":"2023-08-25T09:11:29.264983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Weights&Biases Configuration\nWeights and biases provide an efficient tool for data logging. Let's configure it:","metadata":{}},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"WANDB\")\nwandb.login(key=secret_value_0)","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:11:29.267824Z","iopub.execute_input":"2023-08-25T09:11:29.268485Z","iopub.status.idle":"2023-08-25T09:11:33.473072Z","shell.execute_reply.started":"2023-08-25T09:11:29.268448Z","shell.execute_reply":"2023-08-25T09:11:33.472059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def wandb_init(name, group):\n    \"\"\"Initialize Weights & Biases session.\"\"\"\n    run = wandb.init(project=\"Contrail-Public\",\n                     entity='ralf-c-kinkel',\n                     name=name,\n                     group=group,\n                     save_code=False)\n    return run\n\ndef log_wb(global_val_dice, global_dice, val_loss, loss, lr, elapsed):\n    \"\"\"Log metrics to Weights & Biases.\"\"\"\n    elapsed_time = time.time()\n    wandb.log({\n        'global_val_dice': global_val_dice,\n        'global_dice': global_dice,\n        'val_loss': val_loss,\n        'loss': loss,\n        'learning_rate': lr,\n        'elapsed_time': elapsed,\n    })","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:11:33.474661Z","iopub.execute_input":"2023-08-25T09:11:33.475291Z","iopub.status.idle":"2023-08-25T09:11:33.482085Z","shell.execute_reply.started":"2023-08-25T09:11:33.475255Z","shell.execute_reply":"2023-08-25T09:11:33.481073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Configuration","metadata":{}},{"cell_type":"code","source":"base_dir = \"../input/shifted-ash-for-contrails-05-pixels\"\ntrain_df = pd.read_csv(base_dir + '/train_df.csv')\nval_df = pd.read_csv(base_dir + '/valid_df.csv')","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:11:33.485125Z","iopub.execute_input":"2023-08-25T09:11:33.485812Z","iopub.status.idle":"2023-08-25T09:11:33.527651Z","shell.execute_reply.started":"2023-08-25T09:11:33.485778Z","shell.execute_reply":"2023-08-25T09:11:33.5268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dataset(torch.utils.data.Dataset):\n    def __init__(self, data_path, image_size, mode='train'):\n        self.data_path = data_path \n        if mode == 'train':\n            self.file_names = train_df['record_id'].values\n        elif mode == 'val':\n            self.file_names = val_df['record_id'].values\n        self.image_size = image_size\n        self.resize_image = T.transforms.Resize(image_size) \n        self.transform = A.Compose(train_aug_list)\n        self.mode = mode\n\n    def __len__(self):\n        return len(self.file_names)\n\n    def __getitem__(self, i):\n        file_name = self.file_names[i]\n        z = np.load(f'{self.data_path}/{file_name}.npy')\n        x = np.float32(z[:,:,:3])\n        y = np.float32(z[:,:,3])\n\n        if self.mode == 'train':\n            data = self.transform(image=x, mask=y)\n            x = data['image']\n            y = data['mask']\n        if image_size != 256:\n            x = np.array(self.resize_image(torch.tensor(x.transpose(2,0,1))))\n        else:\n            x = x.transpose(2,0,1)\n        return x, y","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:11:33.528954Z","iopub.execute_input":"2023-08-25T09:11:33.529876Z","iopub.status.idle":"2023-08-25T09:11:33.539914Z","shell.execute_reply.started":"2023-08-25T09:11:33.529844Z","shell.execute_reply":"2023-08-25T09:11:33.539046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Data Augmentation\ntrain_aug_list = [\n    A.RandomRotate90(p=1),\n    A.HorizontalFlip(p=0.5),\n    A.ShiftScaleRotate(rotate_limit=30, scale_limit=0.2)\n]\n\n# Dataset upscaling size\nimage_size = 384 \n\n# Dataloader settings\nbatch_size = 16\nnum_workers = 2\n\n# Datasets\ntrain_dataset = Dataset(base_dir+'/contrails/', image_size, mode='train')\nval_dataset = Dataset(base_dir+'/contrails/', image_size, mode='val')  # band away and mode to val\n\n# Dataloader\ntrain_loader = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers, drop_last=True, pin_memory=True)\nval_loader = torch.utils.data.DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers, drop_last=True, pin_memory=True)\n\n# Show augmentated image\nfor i in range(5):\n    ash_image, label = train_dataset[3]\n    print(ash_image.shape)\n    plt.figure(figsize=(12, 6))\n    ax = plt.subplot(1, 2, 1)\n    ax.imshow(np.transpose(ash_image[:3,:,:].astype(float),(1,2,0)))\n    ax = plt.subplot(1, 2, 2)\n    ax.imshow(label, interpolation='none') \n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:11:33.541439Z","iopub.execute_input":"2023-08-25T09:11:33.541779Z","iopub.status.idle":"2023-08-25T09:11:36.653822Z","shell.execute_reply.started":"2023-08-25T09:11:33.541748Z","shell.execute_reply":"2023-08-25T09:11:36.652881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluation Metric","metadata":{}},{"cell_type":"code","source":"class Dice(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(Dice, self).__init__()\n        self.sigmoid = nn.Sigmoid()\n    def forward(self, inputs, targets, smooth=1, th=0.5):\n        with torch.no_grad():\n            inputs = self.sigmoid(inputs) > th\n            inputs = inputs.view(-1)\n            targets = targets.view(-1)\n            intersection = (inputs * targets).sum()                            \n            dice = (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n        return dice, intersection.item(), inputs.sum().item(), targets.sum().item()\n        \ndice = Dice()","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:11:36.65539Z","iopub.execute_input":"2023-08-25T09:11:36.656026Z","iopub.status.idle":"2023-08-25T09:11:36.665153Z","shell.execute_reply.started":"2023-08-25T09:11:36.655978Z","shell.execute_reply":"2023-08-25T09:11:36.664212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Training","metadata":{}},{"cell_type":"markdown","source":"### Custom Dice Loss\nExplained in ","metadata":{}},{"cell_type":"code","source":"# Credits to segmentation models pytorch\nclass DiceLoss(_Loss):\n    def __init__(\n        self,\n        mode: str,\n        classes: Optional[List[int]] = None,\n        log_loss: bool = False,\n        from_logits: bool = True,\n        smooth: float = 0.0,\n        ignore_index: Optional[int] = None,\n        eps: float = 1e-7,):\n            \n        super(DiceLoss, self).__init__()\n        self.classes = classes\n        self.from_logits = from_logits\n        self.smooth = smooth\n        self.eps = eps\n        self.log_loss = log_loss\n        self.ignore_index = ignore_index\n\n    def forward(self, y_pred: torch.Tensor, y_true: torch.Tensor) -> torch.Tensor:\n        assert y_true.size(0) == y_pred.size(0)\n        if self.from_logits:\n            # Apply activations to get [0..1] class probabilities\n            # Using Log-Exp as this gives more numerically stable result and does not cause vanishing gradient on\n            # extreme values 0 and 1\n            y_pred = F.logsigmoid(y_pred).exp()\n        bs = y_true.size(0)\n        num_classes = y_pred.size(1)\n        dims = (0, 2)\n        y_true = y_true.view(bs, 1, -1)\n        y_pred = y_pred.view(bs, 1, -1)\n        if self.ignore_index is not None:\n            mask = y_true != self.ignore_index\n            y_pred = y_pred * mask\n            y_true = y_true * mask\n        scores = self.compute_score(y_pred, y_true.type_as(y_pred), smooth=self.smooth, eps=self.eps, dims=dims)\n        if self.log_loss:\n            loss = -torch.log(scores.clamp_min(self.eps))\n        else:\n            loss = 1.0 - scores\n        mask = y_true.sum(dims) > 0\n        loss *= mask.to(loss.dtype)\n        if self.classes is not None:\n            loss = loss[self.classes]\n        return self.aggregate_loss(loss)\n\n    def aggregate_loss(self, loss):\n        return loss.mean()\n\n    def compute_score(self, output, target, smooth=0.0, eps=1e-7, dims=None) -> torch.Tensor:\n        return soft_dice_score(output, target, smooth, eps, dims)\n\ndef soft_dice_score(\n    output: torch.Tensor,\n    target: torch.Tensor,\n    smooth: float = 0.0,\n    eps: float = 1e-7,\n    dims=None,\n) -> torch.Tensor:\n    assert output.size() == target.size()\n    if dims is not None:\n        intersection = torch.sum((1-torch.abs(target - output)) * torch.minimum(target, output), dim=dims)  # CUSTOM PART\n        cardinality = torch.sum(output + target, dim=dims)\n    else:\n        intersection = torch.sum((1-torch.abs(target - output)) * torch.minimum(target, output))            # CUSTOM PART\n        cardinality = torch.sum(output + target)\n    dice_score = (2.0 * intersection + smooth) / (cardinality + smooth).clamp_min(eps)\n    return dice_score","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:11:36.666726Z","iopub.execute_input":"2023-08-25T09:11:36.667424Z","iopub.status.idle":"2023-08-25T09:11:36.686162Z","shell.execute_reply.started":"2023-08-25T09:11:36.667392Z","shell.execute_reply":"2023-08-25T09:11:36.685262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = DiceLoss(mode=\"binary\", smooth=1.0, from_logits=True) ","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:11:36.68746Z","iopub.execute_input":"2023-08-25T09:11:36.687954Z","iopub.status.idle":"2023-08-25T09:11:36.700263Z","shell.execute_reply.started":"2023-08-25T09:11:36.68792Z","shell.execute_reply":"2023-08-25T09:11:36.699301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Train Pipeline","metadata":{}},{"cell_type":"code","source":"def train(model=None):\n    # WB Init\n    try:\n        wandb.run.finish()                   #WB         \n        run = wandb_init(name, group)        #WB\n    except:\n        run = wandb_init(name, group)        #WB\n    \n    # Model Init\n    if model == None:\n        model = smp.Unet(\n            encoder_name=encoder,           # choose encoder, e.g. mobilenet_v2 or efficientnet-b0\n            encoder_weights=\"imagenet\",     # use `imagenet` pretrained weights for encoder initialization\n            in_channels=3,                  # model input channels (1 for grayscale images, 3 for RGB, etc.)\n            classes=1,                      # model output channels (number of classes in your dataset)\n        )\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    model.to(device)\n    print(f\"Model at {device}\")\n    \n    # Optimizer and scheduler Init\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr_max, weight_decay=weight_decay)\n    nbatch = len(train_loader) // train_divisor\n    warmup = epochs_warmup * nbatch\n    nsteps = num_epochs * nbatch\n    scheduler = CosineLRScheduler(optimizer,\n                          warmup_t=warmup, warmup_lr_init=warmup_lr_init, warmup_prefix=True,\n                          t_initial=(nsteps - warmup), lr_min=lr_min) \n    \n    # Train Loop\n    val_best_dice = 0.0\n    i_scheduler = 0\n    th = 0.5\n    for epoch in range(num_epochs):\n        t = time.time()\n        train_loss, val_loss = 0, 0\n        train_dice, val_dice = 0, 0\n        global_intersection_train, global_sum_train = 0, 0\n        global_intersection_val, global_sum_val = 0, 0\n        area_pred, area_target = 0, 0\n        model.train()\n        for i_train, (X, y) in enumerate(tqdm(train_loader)): \n            # Load Data\n            X = X.to(device)\n            if label == 'soft': \n                y = y.to(device)     \n            elif label == 'hard': \n                y = (y > 0.51).to(device)\n    \n            # Train Model\n            optimizer.zero_grad()\n            preds = model(X)\n            if image_size != 256: \n                preds = torch.nn.functional.interpolate(preds, size=256, mode='bilinear')\n            loss = criterion(preds.squeeze(), y.squeeze())\n            loss.backward()\n            optimizer.step()\n            scheduler.step(i_scheduler)\n            i_scheduler +=1\n    \n            # Log Results\n            dice_temp, intersection, pred_sum, target_sum = dice(preds, y > 0.51, th=th)      # This cutoff is for soft label to hard label\n            global_intersection_train += intersection\n            global_sum_train += pred_sum + target_sum\n            area_pred += pred_sum\n            area_target += target_sum\n            train_loss += loss.item()\n            train_dice += dice_temp.item()\n            if pred_sum > target_sum and th < 0.999: # Adjust optimal threshold\n                th += 0.001\n            else: \n                th -= 0.001\n                \n            if i_train == nbatch: \n                break\n                \n        # Illustrate Results\n        if show:\n            fig, axs = plt.subplots(1, 3, figsize=(15, 5))  # Create 3 subplots side by side\n            y = y[0].detach().cpu().numpy()  # Detach from computation graph and convert to numpy\n            axs[0].imshow(y)  # Visualize the tensor\n            axs[0].set_title('Ground Truth') \n            X_img = np.transpose(X.cpu().float()[0,:3], (1, 2, 0))\n            axs[1].imshow(X_img)  # Visualize the tensor\n            axs[1].set_title('Input Image')\n            x = preds[0].detach().cpu().numpy()  # Detach from computation graph and convert to numpy\n            x = np.transpose(x, (1, 2, 0))  # Make sure channels dimension comes last\n            axs[2].imshow(x)  # Visualize the tensor\n            axs[2].set_title('Predicted Mask')\n            plt.show()\n    \n        model.eval()\n        with torch.no_grad():\n            for i_val, (X, y) in enumerate(tqdm(val_loader)):   \n                # Load Data\n                X = X.to(device)\n                y = y.to(device)\n    \n                # Inference\n                preds = model(X)\n                if image_size != 256: preds = torch.nn.functional.interpolate(preds, size=256, mode='bilinear')\n                loss = criterion(preds.squeeze(), y.squeeze())\n    \n                # Log Results\n                dice_temp, intersection, pred_sum, target_sum = dice(preds, y, th=th)\n                global_intersection_val += intersection\n                global_sum_val += pred_sum + target_sum\n                val_loss += loss.item()\n                val_dice += dice_temp.item()\n    \n        # Display Log and Saving\n        train_loss = f'{train_loss/nbatch:.5f}'\n        val_loss = f'{val_loss/len(val_loader):.5f}'\n        train_dice = f'{(train_dice+1e-23)/i_train:.5f}'\n        val_dice = f'{(val_dice+1e-23)/i_val:.5f}'\n        global_train_dice = 2*global_intersection_train/global_sum_train\n        global_val_dice = 2*global_intersection_val/global_sum_val\n        lr = optimizer.param_groups[0][\"lr\"]\n        elapsed = time.time()-t\n        if global_val_dice > val_best_dice:\n            val_best_dice = global_val_dice\n            torch.save(model, 'model.pth')\n            best_model = copy.deepcopy(model)\n        print(f'Global Dice Train: {2*global_intersection_train/global_sum_train} at threshold = {th}') \n        print(f'Global Dice Val: {2*global_intersection_val/global_sum_val} at threshold = {th}')\n        print (f'Epoch [{(epoch+1)}/{num_epochs}], loss: {train_loss}, val_loss: {val_loss}')\n        log_wb(global_val_dice, global_train_dice, float(val_loss), float(train_loss), lr, elapsed)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:11:36.703882Z","iopub.execute_input":"2023-08-25T09:11:36.704366Z","iopub.status.idle":"2023-08-25T09:11:36.729258Z","shell.execute_reply.started":"2023-08-25T09:11:36.704342Z","shell.execute_reply":"2023-08-25T09:11:36.728344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Settings and Train","metadata":{}},{"cell_type":"code","source":"num_epochs = 100\nnum_workers = 2\n\n# Optimizer & Scheduler\nmult = 0.3\nlr_max = 3e-4 * mult  \nepochs_warmup = 5\nwarmup_lr_init = 5e-4 * mult\nlr_min = 1e-6 * mult                   \nscheduler_name = \"CosineAnnealingLR\"\nweight_decay = 0.1\n\ntrain_divisor = 5   # Validation is done x times per epoch, or each 'epoch' only contains 1/x of train_sampls   \n\nencoder = \"timm-resnest101e\" #timm-resnest200e best\nshow = True\nlabel = \"soft\"\n\nname = f'u101e,100,big384'\ngroup = 'big'\n\nmodel = None\n#model = torch.load('submit/model.pth')\nmodel = train(model=model)","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:14:11.811481Z","iopub.execute_input":"2023-08-25T09:14:11.812033Z","iopub.status.idle":"2023-08-25T09:30:54.036323Z","shell.execute_reply.started":"2023-08-25T09:14:11.811978Z","shell.execute_reply":"2023-08-25T09:30:54.032898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Evaluation","metadata":{}},{"cell_type":"code","source":"device = 'cuda'\nbest_model = torch.load('/kaggle/working/model.pth')","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:31:18.244689Z","iopub.execute_input":"2023-08-25T09:31:18.245106Z","iopub.status.idle":"2023-08-25T09:31:18.556814Z","shell.execute_reply.started":"2023-08-25T09:31:18.245075Z","shell.execute_reply":"2023-08-25T09:31:18.555721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"outputs = []\nlabels = []\nwith torch.no_grad():\n    for i_val, (X, y) in enumerate(tqdm(val_loader)):\n        X = X.to(device).float()\n        y = y.to(device).float()\n        pred = best_model(X)                     \n        if image_size != 256: \n            pred = torch.nn.functional.interpolate(pred, size=256, mode='bilinear')\n        outputs.append(pred)\n        labels.append(y)\nall_preds = torch.cat(outputs)\nall_labels = torch.cat(labels)\noutputs = 0\nbest = 0\nfor i in range(1,1000):\n    th = i * 0.001\n    dice_temp, _, pred_sum, _ = dice(all_preds, all_labels, th=th)\n    if dice_temp > best:\n        best = dice_temp\n        best_th = th\nprint(best, best_th)","metadata":{"execution":{"iopub.status.busy":"2023-08-25T09:32:38.018106Z","iopub.execute_input":"2023-08-25T09:32:38.018488Z","iopub.status.idle":"2023-08-25T09:33:26.777867Z","shell.execute_reply.started":"2023-08-25T09:32:38.01846Z","shell.execute_reply":"2023-08-25T09:33:26.776604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}