{"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":"# imports ","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport os\n\nfrom argparse import Namespace\nfrom pathlib import Path\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import TensorDataset, DataLoader\nfrom torchvision import transforms\nimport albumentations as A\n\nfrom transformers import get_cosine_schedule_with_warmup\n\nfrom tqdm.notebook import tqdm\n\nif os.environ.get(\"KAGGLE_KERNEL_RUN_TYPE\", \"\"):\n    !pip install -q /kaggle/input/torchsummary/torchsummary-1.5.1-py3-none-any.whl\n    \nfrom torchsummary import summary\n\ntorch.__version__","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !mkdir -p ~/.torch/models\n# !cp /kaggle/input/torchvision-resnet-pretrained/resnet34-b627a593.pth ~/.torch/models/resnet34-333f7ec4.pth\n\n\nimport sys\n# sys.path.append(\"../input/pretrained-models-pytorch\")\nsys.path.append(\"../input/efficientnet-pytorch\")\n# sys.path.append(\"/kaggle/input/smp-github/segmentation_models.pytorch-master\")\n!pip install --no-index --find-links=\"/kaggle/input/segmentation-models-pytorch\" segmentation_models_pytorch \nimport segmentation_models_pytorch as smp\n\nprint(f\"Segmentation Models version: {smp.__version__}\")\n\n\n!mkdir -p /root/.cache/torch/hub/checkpoints\n# !cp /kaggle/input/torchvision-resnet-pretrained/resnet34-b627a593.pth /root/.cache/torch/hub/checkpoints/resnet34-333f7ec4.pth","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:08.343148Z","iopub.execute_input":"2023-08-01T04:03:08.344256Z","iopub.status.idle":"2023-08-01T04:03:27.614827Z","shell.execute_reply.started":"2023-08-01T04:03:08.344228Z","shell.execute_reply":"2023-08-01T04:03:27.613331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # !pip -q install segmentation_models_pytorch\n\n# # !pip download segmentation-models-pytorch -d .\n\n# !pip install --no-index --find-links=\"/kaggle/input/segmentation-models-pytorch\" segmentation_models_pytorch \n\n# import segmentation_models_pytorch as smp\n\n# print(f\"Segmentation Models version: {smp.__version__}\")","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:27.61714Z","iopub.execute_input":"2023-08-01T04:03:27.617598Z","iopub.status.idle":"2023-08-01T04:03:27.624469Z","shell.execute_reply.started":"2023-08-01T04:03:27.617555Z","shell.execute_reply":"2023-08-01T04:03:27.623353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# configs","metadata":{}},{"cell_type":"code","source":"if os.environ.get(\"KAGGLE_KERNEL_RUN_TYPE\", \"\"):\n    BASE_DIR = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\nelse:\n    BASE_DIR =  'data'","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:27.627898Z","iopub.execute_input":"2023-08-01T04:03:27.628591Z","iopub.status.idle":"2023-08-01T04:03:27.637788Z","shell.execute_reply.started":"2023-08-01T04:03:27.628542Z","shell.execute_reply":"2023-08-01T04:03:27.636666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"configs = Namespace(\n    base_dir= Path(BASE_DIR),\n\n    train = True,\n    train_aug = True,\n\n    batch_size= 32,\n    epochs= 30,\n\n    encoder = 'efficientnet-b3',\n#     weights= 'imagenet',\n    weights = \"/kaggle/input/contrails-trained-model/model_checkpoint_e20.pt\",\n    classes= ['contrail'],\n    activation= None,\n\n    img_size= 256,\n\n    num_worers= 2,\n    shuffle= True,\n\n    warmup = 0,\n    lr=3e-3,\n    device= torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"),\n\n)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:27.639654Z","iopub.execute_input":"2023-08-01T04:03:27.64008Z","iopub.status.idle":"2023-08-01T04:03:27.678545Z","shell.execute_reply.started":"2023-08-01T04:03:27.640033Z","shell.execute_reply":"2023-08-01T04:03:27.677755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# transform_size = A.Compose([\n#     A.Resize(256, 256, 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":{"iopub.status.busy":"2023-08-01T04:03:27.6812Z","iopub.execute_input":"2023-08-01T04:03:27.681777Z","iopub.status.idle":"2023-08-01T04:03:27.699839Z","shell.execute_reply.started":"2023-08-01T04:03:27.681738Z","shell.execute_reply":"2023-08-01T04:03:27.698886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# load data","metadata":{}},{"cell_type":"code","source":"def get_paths(data_type):\n\n    ids_list = os.listdir(os.path.join(configs.base_dir, data_type))\n\n    df = pd.DataFrame(ids_list, columns=['record_id'])\n\n    df['path'] = os.path.join(configs.base_dir, data_type ) +\"/\"+ df['record_id'].astype(str)\n\n    return df\n\ntrain_df = get_paths('train')\nval_df = get_paths('validation')","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:27.701791Z","iopub.execute_input":"2023-08-01T04:03:27.702914Z","iopub.status.idle":"2023-08-01T04:03:28.014281Z","shell.execute_reply.started":"2023-08-01T04:03:27.702879Z","shell.execute_reply":"2023-08-01T04:03:28.013268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, train=True, transforms=None):\n        self.df = df  # Initialize the instance variable df to store the DataFrame.\n        self.trn = train  # Initialize the instance variable trn to indicate if it is a training dataset.\n        self.transforms = transforms  # Initialize the instance variable transforms to store the transforms.\n\n    def read_record(self, directory):\n\n        record_data = {}  # Create a dictionary to store the record data.\n        for x in [\n            \"band_11\",\n            \"band_14\",\n            \"band_15\"\n        ]:\n            record_data[x] = np.load(os.path.join(directory, x + \".npy\"))  # Load data for each band and store it in the dictionary.\n\n        if self.trn:\n            record_data[\"mask\"] = np.load(os.path.join(directory, \"human_pixel_masks.npy\"))\n\n        return record_data\n\n    def normalize_range(self, data, bounds):\n        \"\"\"Normalize data to the range [0, 1].\"\"\"\n        return (data - bounds[0]) / (bounds[1] - bounds[0])\n\n    def get_false_color(self, record_data):\n\n        _T11_BOUNDS = (243, 303)\n        _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n        _TDIFF_BOUNDS = (-4, 2)\n\n        N_TIMES_BEFORE = 4\n        r = self.normalize_range(record_data[\"band_15\"] - record_data[\"band_14\"], _TDIFF_BOUNDS)\n        g = self.normalize_range(record_data[\"band_14\"] - record_data[\"band_11\"], _CLOUD_TOP_TDIFF_BOUNDS)\n        b = self.normalize_range(record_data[\"band_14\"], _T11_BOUNDS)\n        false_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n        img = false_color[..., N_TIMES_BEFORE]\n\n        if self.trn:\n            mask_img = record_data[\"mask\"]\n\n            return img, mask_img\n        \n        return img\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        con_path = row.path\n        data = self.read_record(con_path)  # dictionary with keys: band_11, band_14, band_15 and values: numpy arrays (height, width, channels)\n\n        if self.trn:\n            img, mask_img = self.get_false_color(data)\n\n            if configs.train_aug:\n                if self.transforms is not None:\n                    \n                    augmented = self.transforms(image=img, mask=mask_img)\n                    img = augmented['image']\n                    mask_img = augmented['mask']\n\n            img = torch.tensor(img).float()\n            mask_img = torch.tensor(mask_img).float()\n\n            img = img.permute(2, 0, 1)\n            mask_img = mask_img.permute(2, 0, 1)\n\n            return img, mask_img\n        \n        img = self.get_false_color(data)\n        \n        img = torch.tensor(img).float()\n\n        img = img.permute(2, 0, 1)\n        return img\n    \n\n    def __len__(self):\n        return len(self.df)\n","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:28.015898Z","iopub.execute_input":"2023-08-01T04:03:28.016298Z","iopub.status.idle":"2023-08-01T04:03:28.032389Z","shell.execute_reply.started":"2023-08-01T04:03:28.016265Z","shell.execute_reply":"2023-08-01T04:03:28.0312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = ContrailsDataset(\n        train_df,\n        train = True,\n        transforms = train_transform\n    )\n\ntrain_dl = DataLoader(train_ds, batch_size=configs.batch_size, shuffle=True, num_workers = configs.num_worers, pin_memory=True, prefetch_factor=4 )\n\nval_ds = ContrailsDataset(\n        val_df,\n        train = True,\n        transforms = None\n    )\n\nval_dl = DataLoader(val_ds, batch_size=configs.batch_size, num_workers = configs.num_worers, pin_memory=True, prefetch_factor=4 )","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:28.034253Z","iopub.execute_input":"2023-08-01T04:03:28.035007Z","iopub.status.idle":"2023-08-01T04:03:28.047471Z","shell.execute_reply.started":"2023-08-01T04:03:28.034976Z","shell.execute_reply":"2023-08-01T04:03:28.046639Z"},"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-08-01T04:03:28.050906Z","iopub.execute_input":"2023-08-01T04:03:28.051368Z","iopub.status.idle":"2023-08-01T04:03:40.614945Z","shell.execute_reply.started":"2023-08-01T04:03:28.05134Z","shell.execute_reply":"2023-08-01T04:03:40.613731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, label = next(iter(val_dl))\nimg.shape, label.shape","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:40.616703Z","iopub.execute_input":"2023-08-01T04:03:40.617492Z","iopub.status.idle":"2023-08-01T04:03:46.044281Z","shell.execute_reply.started":"2023-08-01T04:03:40.617455Z","shell.execute_reply":"2023-08-01T04:03:46.043155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Architecture","metadata":{}},{"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            encoder_weights = None,\n            decoder_use_batchnorm=True,\n            classes=len(cfg.classes), \n            activation=cfg.activation,\n        )\n        \n#         # load pre-trained weights from file\n#         pretrained_weights = torch.load(cfg.weights)\n        \n#         # set the weights of the encoder\n#         self.model.encoder.load_state_dict(pretrained_weights)\n\n    \n    def forward(self, imgs):\n        \n        x = imgs\n\n        logits = self.model(x)\n\n        return logits\n        \n        # if Config.image_size != 256:\n        #     logits = F.interpolate(logits, size=(256, 256), mode='nearest-exact')","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:46.048805Z","iopub.execute_input":"2023-08-01T04:03:46.05112Z","iopub.status.idle":"2023-08-01T04:03:46.061476Z","shell.execute_reply.started":"2023-08-01T04:03:46.051083Z","shell.execute_reply":"2023-08-01T04:03:46.060704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(UNet(configs).to(configs.device), input_size=(3, 256, 256))","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:46.066052Z","iopub.execute_input":"2023-08-01T04:03:46.068556Z","iopub.status.idle":"2023-08-01T04:03:51.408918Z","shell.execute_reply.started":"2023-08-01T04:03:46.068514Z","shell.execute_reply":"2023-08-01T04:03:51.40781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"class Dice(nn.Module):\n    def __init__(self, use_sigmoid=True):\n        super(Dice, self).__init__()\n        self.sigmoid = nn.Sigmoid()\n        self.use_sigmoid = use_sigmoid\n\n    def forward(self, inputs, targets, smooth=1):\n        if self.use_sigmoid:\n            inputs = self.sigmoid(inputs)       \n        \n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()\n        dice = (2.0 *intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n        \n        return dice\n    \ndice = Dice()","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:51.410482Z","iopub.execute_input":"2023-08-01T04:03:51.410861Z","iopub.status.idle":"2023-08-01T04:03:51.418924Z","shell.execute_reply.started":"2023-08-01T04:03:51.410825Z","shell.execute_reply":"2023-08-01T04:03:51.417732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyTrainer:\n    def __init__(self, model, optimizer, loss_fn, lr_scheduler):\n        self.validation_losses = []\n        self.batch_losses = []\n        self.epoch_losses = []\n        self.learning_rates = []\n        self.model = model\n        self.optimizer = optimizer\n        self.loss_fn = loss_fn\n        self.lr_scheduler = lr_scheduler\n        self._check_optim_net_aligned()\n\n    # Ensures that the given optimizer points to the given model\n    def _check_optim_net_aligned(self):\n        assert self.optimizer.param_groups[0]['params'] == list(self.model.parameters())\n\n    # Trains the model\n    def fit(self,\n            train_dataloader: DataLoader,\n            test_dataloader: DataLoader,\n            epochs: int = 10,\n            eval_every: int = 1,\n            ):\n  \n        for e in tqdm(range(epochs)):\n            print(\"New learning rate: {}\".format(self.lr_scheduler.get_last_lr()))\n            self.learning_rates.append(self.lr_scheduler.get_last_lr()[0])\n\n            # Stores data about the batch\n            batch_losses = []\n            sub_batch_losses = []\n\n            for i, data in enumerate(train_dataloader):\n                \n                self.model.train()\n                torch.set_grad_enabled(True)\n                if (i+1) % 150 == 0:\n                    print(f'epotch: {e} batch: {i}/{len(train_dataloader)} loss: {torch.Tensor(sub_batch_losses).mean()}')\n                    sub_batch_losses.clear()\n                # Every data instance is an input + label pair\n                images, mask = data\n                \n                if torch.cuda.is_available():\n                    images = images.cuda()\n                    mask = mask.cuda()\n\n                # Zero your gradients for every batch!\n                self.optimizer.zero_grad()\n                # Make predictions for this batch\n                outputs = self.model(images)\n                # Compute the loss and its gradients\n                loss = self.loss_fn(outputs, mask)\n                loss.backward()\n                # Adjust learning weights\n                self.optimizer.step()\n\n                # Saves data\n                self.batch_losses.append(loss.item())\n                batch_losses.append(loss)\n                sub_batch_losses.append(loss)\n            \n            \n\n            # Adjusts learning rate\n            if self.lr_scheduler is not None:\n                self.lr_scheduler.step()\n\n            # Reports on the path\n            mean_epoch_loss = torch.Tensor(batch_losses).mean()\n            self.epoch_losses.append(mean_epoch_loss.item())\n            print('Train Epoch: {} Average Loss: {:.6f}'.format(e, mean_epoch_loss))\n\n            # Reports on the training progress\n            if (e + 1) % eval_every == 0:\n                if not os.path.exists(\"segmentation_checkpoints\"):\n                    !mkdir segmentation_checkpoints\n                torch.save(self.model.state_dict(), \"segmentation_checkpoints/model_checkpoint_e\" + str(e) + \".pt\")\n                with torch.no_grad():\n                \n                    self.model.eval()\n                    torch.set_grad_enabled(False)\n                    losses = []\n                    for i, data in enumerate(test_dataloader):\n                        # Every data instance is an input + label pair\n                        images, mask = data\n\n                        if torch.cuda.is_available():\n                            images = images.cuda()\n                            mask = mask.cuda()\n\n                        output = self.model(images)\n                        loss = self.loss_fn(output, mask)\n                        losses.append(loss.item())\n                        \n                    avg_loss = torch.Tensor(losses).mean().item()\n                    self.validation_losses.append(avg_loss)\n                    print(\"Validation loss after\", (e + 1), \"epochs was\", round(avg_loss, 4))","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:57.426307Z","iopub.execute_input":"2023-08-01T04:03:57.426691Z","iopub.status.idle":"2023-08-01T04:03:57.4569Z","shell.execute_reply.started":"2023-08-01T04:03:57.426648Z","shell.execute_reply":"2023-08-01T04:03:57.455555Z"},"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-08-01T04:03:58.786937Z","iopub.execute_input":"2023-08-01T04:03:58.787304Z","iopub.status.idle":"2023-08-01T04:03:58.792611Z","shell.execute_reply.started":"2023-08-01T04:03:58.787275Z","shell.execute_reply":"2023-08-01T04:03:58.791486Z"},"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.epochs * (total_steps // cfg.batch_size)\n    )\n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:03:59.547337Z","iopub.execute_input":"2023-08-01T04:03:59.547739Z","iopub.status.idle":"2023-08-01T04:03:59.553787Z","shell.execute_reply.started":"2023-08-01T04:03:59.547702Z","shell.execute_reply":"2023-08-01T04:03:59.552642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = False\n\nif train:\n    model = UNet(configs).to(configs.device)\n    model.to(configs.device)\n\n    total_steps = len(train_ds)\n\n    # criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor(100))\n    criterion = smp.losses.DiceLoss(mode='binary')\n    optimizer = get_optimizer(params=model.parameters(), lr=configs.lr)\n#     scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.9)\n    scheduler = get_scheduler(configs, optimizer, total_steps)\n\n    trainer = MyTrainer(model, optimizer, criterion, scheduler)\n    trainer.fit(train_dl, val_dl, epochs=configs.epochs)\n\nelse:\n    \n    model = UNet(configs).to(configs.device)\n    model.load_state_dict(torch.load(os.path.join('/kaggle/input/contrails-trained-model/model_checkpoint_e20.pt')))\n    model.eval()\n    model.to(configs.device)","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:04:01.534978Z","iopub.execute_input":"2023-08-01T04:04:01.535392Z","iopub.status.idle":"2023-08-01T04:04:02.53187Z","shell.execute_reply.started":"2023-08-01T04:04:01.535364Z","shell.execute_reply":"2023-08-01T04:04:02.530858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Progress","metadata":{}},{"cell_type":"code","source":"if train:\n    df_data = pd.DataFrame({'Batch Losses': trainer.batch_losses})\n\n    sns.lineplot(data=df_data)\n    plt.xlabel('Batch')\n    plt.ylabel('Loss')\n    plt.title('Batch Loss')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T20:10:52.148377Z","iopub.execute_input":"2023-07-30T20:10:52.149096Z","iopub.status.idle":"2023-07-30T20:10:52.154945Z","shell.execute_reply.started":"2023-07-30T20:10:52.149061Z","shell.execute_reply":"2023-07-30T20:10:52.153754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    df_data = pd.DataFrame({'Loss': trainer.epoch_losses})\n\n    sns.lineplot(data=df_data)\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Model Argavgre Training Loss over Epochs')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T20:10:52.1563Z","iopub.execute_input":"2023-07-30T20:10:52.156793Z","iopub.status.idle":"2023-07-30T20:10:52.302128Z","shell.execute_reply.started":"2023-07-30T20:10:52.156759Z","shell.execute_reply":"2023-07-30T20:10:52.301099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    df_data = pd.DataFrame({'Loss': trainer.validation_losses})\n\n    sns.lineplot(data=df_data)\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Model Validation Loss over Epochs')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T20:10:52.303789Z","iopub.execute_input":"2023-07-30T20:10:52.304288Z","iopub.status.idle":"2023-07-30T20:10:52.310978Z","shell.execute_reply.started":"2023-07-30T20:10:52.304235Z","shell.execute_reply":"2023-07-30T20:10:52.309876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    df_data = pd.DataFrame({'Learning rates': trainer.learning_rates})\n\n    sns.lineplot(data=df_data)\n    plt.xlabel('Epoch')\n    plt.ylabel('Learinig Rate')\n    plt.title('Learinig Rate over Epochs')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T20:10:52.312636Z","iopub.execute_input":"2023-07-30T20:10:52.313099Z","iopub.status.idle":"2023-07-30T20:10:52.319957Z","shell.execute_reply.started":"2023-07-30T20:10:52.313065Z","shell.execute_reply":"2023-07-30T20:10:52.318822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimum Threshold","metadata":{}},{"cell_type":"code","source":"class DiceThresholdTester:\n    \n    def __init__(self, model: nn.Module, data_loader: torch.utils.data.DataLoader):\n        self.model = model\n        self.data_loader = data_loader\n        self.cumulative_mask_pred = []\n        self.cumulative_mask_true = []\n        \n    def precalculate_prediction(self) -> None:\n        sigmoid = nn.Sigmoid()\n        \n        for images, mask_true in self.data_loader:\n            if torch.cuda.is_available():\n                images = images.cuda()\n\n            mask_pred = sigmoid(model.forward(images))\n\n            self.cumulative_mask_pred.append(mask_pred.cpu().detach().numpy())\n            self.cumulative_mask_true.append(mask_true.cpu().detach().numpy())\n            \n        self.cumulative_mask_pred = np.concatenate(self.cumulative_mask_pred, axis=0)\n        self.cumulative_mask_true = np.concatenate(self.cumulative_mask_true, axis=0)\n\n        self.cumulative_mask_pred = torch.flatten(torch.from_numpy(self.cumulative_mask_pred))\n        self.cumulative_mask_true = torch.flatten(torch.from_numpy(self.cumulative_mask_true))\n    \n    def test_threshold(self, threshold: float) -> float:\n        _dice = Dice(use_sigmoid=False)\n        after_threshold = np.zeros(self.cumulative_mask_pred.shape)\n        after_threshold[self.cumulative_mask_pred[:] > threshold] = 1\n        after_threshold[self.cumulative_mask_pred[:] < threshold] = 0\n        after_threshold = torch.flatten(torch.from_numpy(after_threshold))\n        return _dice(self.cumulative_mask_true, after_threshold).item()","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:04:18.018216Z","iopub.execute_input":"2023-08-01T04:04:18.018621Z","iopub.status.idle":"2023-08-01T04:04:18.03073Z","shell.execute_reply.started":"2023-08-01T04:04:18.018589Z","shell.execute_reply":"2023-08-01T04:04:18.029735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train = True\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:04:28.85659Z","iopub.execute_input":"2023-08-01T04:04:28.857637Z","iopub.status.idle":"2023-08-01T04:04:28.887663Z","shell.execute_reply.started":"2023-08-01T04:04:28.8576Z","shell.execute_reply":"2023-08-01T04:04:28.886214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_dl = DataLoader(val_ds, batch_size=16, num_workers = configs.num_worers, pin_memory=True, prefetch_factor=4 )","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:04:39.521051Z","iopub.execute_input":"2023-08-01T04:04:39.521422Z","iopub.status.idle":"2023-08-01T04:04:39.526923Z","shell.execute_reply.started":"2023-08-01T04:04:39.521393Z","shell.execute_reply":"2023-08-01T04:04:39.525436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    dice_threshold_tester = DiceThresholdTester(model, val_dl)\n    dice_threshold_tester.precalculate_prediction()","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:04:40.505567Z","iopub.execute_input":"2023-08-01T04:04:40.506591Z","iopub.status.idle":"2023-08-01T04:06:05.271428Z","shell.execute_reply.started":"2023-08-01T04:04:40.506546Z","shell.execute_reply":"2023-08-01T04:06:05.270122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    thresholds_to_test = [round(x * 0.01, 2) for x in range(101)]\n\n    optim_threshold = 0.18\n    best_dice_score = -1\n\n    thresholds = []\n    dice_scores = []\n\n    for t in thresholds_to_test:\n        dice_score = dice_threshold_tester.test_threshold(t)\n        if dice_score > best_dice_score:\n            best_dice_score = dice_score\n            optim_threshold = t\n\n        thresholds.append(t)\n        dice_scores.append(dice_score)\n\n    print(f'Best Threshold: {optim_threshold} with dice: {best_dice_score}')\n    df_threshold_data = pd.DataFrame({'Threshold': thresholds, 'Dice Score': dice_scores})\nelse:\n    optim_threshold = 0.75","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:07:35.887549Z","iopub.execute_input":"2023-08-01T04:07:35.888494Z","iopub.status.idle":"2023-08-01T04:12:29.415774Z","shell.execute_reply.started":"2023-08-01T04:07:35.888456Z","shell.execute_reply":"2023-08-01T04:12:29.414596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    df_threshold_data.tail(), df_threshold_data.shape","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:14:11.186227Z","iopub.execute_input":"2023-08-01T04:14:11.186641Z","iopub.status.idle":"2023-08-01T04:14:11.192886Z","shell.execute_reply.started":"2023-08-01T04:14:11.186611Z","shell.execute_reply":"2023-08-01T04:14:11.191494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sigmoid(x):\n    return 1 / (1 + np.exp(-x+1e-8))\n\nbatches_to_show = 1\nmodel.eval()\n\nfor i, data in enumerate(train_dl):\n    images, mask = data\n    \n    # Predict mask for this instance\n    if torch.cuda.is_available():\n        images = images.cuda()\n    predicated_mask = sigmoid(model.forward(images[:, :, :, :]).cpu().detach().numpy())\n\n    \n    # Apply threshold\n    predicated_mask_with_threshold = np.zeros((images.shape[0], 256, 256))\n    predicated_mask_with_threshold[predicated_mask[:, 0, :, :] < optim_threshold] = 0\n    predicated_mask_with_threshold[predicated_mask[:, 0, :, :] > optim_threshold] = 1\n    \n    images = images.cpu()\n        \n    for img_num in range(0, images.shape[0]):\n        fig, axes = plt.subplots(nrows=1, ncols=4, figsize=(20,10))\n        axes = axes.flatten()\n        \n        # Show groud trought \n        axes[0].imshow(mask[img_num, 0, :, :])\n        axes[0].axis('off')\n        axes[0].set_title('Ground Truth')\n        \n        # Show ash color scheme input image\n        # axes[1].imshow( np.concatenate(\n        #     (\n        #     np.expand_dims(images[img_num, 0, :, :], axis=2),\n        #     np.expand_dims(images[img_num, 1, :, :], axis=2),\n        #     np.expand_dims(images[img_num, 2, :, :], axis=2)\n        # ), axis=2))\n        axes[1].imshow(images[img_num, :, :, :].permute(1, 2, 0))\n        axes[1].axis('off')\n        axes[1].set_title('Ash color scheeme input - Frame 4')\n\n        # Show predicted mask\n        axes[2].imshow(predicated_mask[img_num, 0, :, :], vmin=0, vmax=1)\n        axes[2].axis('off')\n        axes[2].set_title('Predicted probability mask')\n\n        # Show predicted mask after threshold\n        axes[3].imshow(predicated_mask_with_threshold[img_num, :, :])\n        axes[3].axis('off')\n        axes[3].set_title('Predicted mask with threshold')\n        plt.show()\n    \n    if i + 1 >= batches_to_show:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:15:43.732202Z","iopub.execute_input":"2023-08-01T04:15:43.732602Z","iopub.status.idle":"2023-08-01T04:16:08.825989Z","shell.execute_reply.started":"2023-08-01T04:15:43.732574Z","shell.execute_reply":"2023-08-01T04:16:08.824853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"# clear the cache\n\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:16:08.828186Z","iopub.execute_input":"2023-08-01T04:16:08.828481Z","iopub.status.idle":"2023-08-01T04:16:09.182204Z","shell.execute_reply.started":"2023-08-01T04:16:08.828456Z","shell.execute_reply":"2023-08-01T04:16:09.180822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = get_paths('test')\n\n# cast record_id to int\ntest_df[\"record_id\"] = test_df.record_id.astype(int)\n\ntest_ds = ContrailsDataset(\n        test_df,\n        train = False\n    )\n\ntest_batch_size = 1\n\ntest_dl = DataLoader(test_ds, batch_size=test_batch_size, num_workers = configs.num_worers)\n\ndel test_ds","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:16:09.184316Z","iopub.execute_input":"2023-08-01T04:16:09.185371Z","iopub.status.idle":"2023-08-01T04:16:09.200208Z","shell.execute_reply.started":"2023-08-01T04:16:09.185334Z","shell.execute_reply":"2023-08-01T04:16:09.199259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#source https://www.kaggle.com/code/inversion/contrails-rle-submission?scriptVersionId=128527711&cellId=4\n\ndef rle_encode(x, fg_val=1):\n    \"\"\"\n    Args:\n        x:  numpy array of shape (height, width), 1 - mask, 0 - background\n    Returns: run length encoding as list\n    \"\"\"\n\n    dots = np.where(\n        x.T.flatten() == fg_val)[0]  # .T sets Fortran order down-then-right\n    run_lengths = []\n    prev = -2\n    for b in dots:\n        if b > prev + 1:\n            run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n    return run_lengths\n\n\ndef list_to_string(x):\n    \"\"\"\n    Converts list to a string representation\n    Empty list returns '-'\n    \"\"\"\n    if x: # non-empty list\n        s = str(x).replace(\"[\", \"\").replace(\"]\", \"\").replace(\",\", \"\")\n    else:\n        s = '-'\n    return s\n","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:16:09.203222Z","iopub.execute_input":"2023-08-01T04:16:09.203869Z","iopub.status.idle":"2023-08-01T04:16:09.213393Z","shell.execute_reply.started":"2023-08-01T04:16:09.203825Z","shell.execute_reply":"2023-08-01T04:16:09.212433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv(os.path.join(configs.base_dir, \"sample_submission.csv\"), index_col='record_id')","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:16:09.215117Z","iopub.execute_input":"2023-08-01T04:16:09.215551Z","iopub.status.idle":"2023-08-01T04:16:09.229318Z","shell.execute_reply.started":"2023-08-01T04:16:09.21552Z","shell.execute_reply":"2023-08-01T04:16:09.228374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, images in enumerate(test_dl):\n    \n    \n    image_id = torch.tensor(test_df.iloc[i]['record_id'])\n    \n    # Predict mask for this instance\n    if torch.cuda.is_available():\n        images = images.cuda()\n    predicated_mask = sigmoid(model.forward(images[:, :, :, :]).cpu().detach().numpy())\n    \n    # Apply threshold\n    predicated_mask_with_threshold = np.zeros((images.shape[0], 256, 256))\n    predicated_mask_with_threshold[predicated_mask[:, 0, :, :] < optim_threshold] = 0\n    predicated_mask_with_threshold[predicated_mask[:, 0, :, :] > optim_threshold] = 1\n    \n    current_mask = predicated_mask_with_threshold[:, :, :]\n    current_image_id = image_id.item()\n    submission.loc[int(current_image_id), 'encoded_pixels'] = list_to_string(rle_encode(current_mask))","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:16:09.230852Z","iopub.execute_input":"2023-08-01T04:16:09.23122Z","iopub.status.idle":"2023-08-01T04:16:09.630301Z","shell.execute_reply.started":"2023-08-01T04:16:09.231186Z","shell.execute_reply":"2023-08-01T04:16:09.62888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.head()","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:16:09.632287Z","iopub.execute_input":"2023-08-01T04:16:09.633181Z","iopub.status.idle":"2023-08-01T04:16:09.645309Z","shell.execute_reply.started":"2023-08-01T04:16:09.633137Z","shell.execute_reply":"2023-08-01T04:16:09.643484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-08-01T04:16:09.647276Z","iopub.execute_input":"2023-08-01T04:16:09.647738Z","iopub.status.idle":"2023-08-01T04:16:09.660103Z","shell.execute_reply.started":"2023-08-01T04:16:09.647703Z","shell.execute_reply":"2023-08-01T04:16:09.659112Z"},"trusted":true},"execution_count":null,"outputs":[]}]}