{"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":"## Summary\n- Baseline written using Pytorch Lightning\n- resnest26d as encoder\n- Dataset is taken from: https://www.kaggle.com/datasets/shashwatraman/contrails-images-ash-color\n\n#### Improvements over [previous previous version](https://www.kaggle.com/code/egortrushin/gr-icrgw-pytorch-lightning-baseline-unet-resnest) (please upvote).\n- Option to change image size\n- Mixed precision training (only useful with T4x2, on P100 this slows down training). This helps to use GPU memory more efficiently\n- Training using 2 GPUs - with 2 GPUs we have more memory and higher speed\n- Other numerous small changes\n#### Improvements over previous version\n- Organized code, Separate classes with repeated names, added docstring and explanation on most code steps.  ","metadata":{}},{"cell_type":"markdown","source":"## Setup Environment and libraries","metadata":{}},{"cell_type":"markdown","source":"First of all lets setup the environment, librares, input and output files... etc\nwe would be using pythorch lighting and some custom models\n- [ ] Import standard libraries\n- [ ] Define paths to everything\n- [ ] Import custom libraries and models\n- [ ] Save YAML file with configurations\n\nPyTorch Lightning is a Python library that makes training and deploying deep learning models easier by providing a simple and standardized way to write code. It handles the repetitive parts of training, like setting up the training loop and handling gradients, so you don't have to worry about them. This makes the code cleaner and easier to understand and maintain.","metadata":{}},{"cell_type":"markdown","source":"- [X] Import libraries\n\n  - **sys** :  system-specific parameters and functions in this case we will use it to add path to the environment variables.📝\n  - **os** :  interacting with the operating system, in this case for creating directories in the filesystem and joining paths.🖥️\n  - **shutil** : high-level file operations in Python, like copying and moving files.📁\n  - **torch** : Our deep learning framework for building and training neural networks. 🔥\n  - **numpy** : Numerical computing and array operations in Python. 🔢\n  - **torchvision.transforms** : For computer vision tasks, providing datasets, models, and transforms 🌄.","metadata":{}},{"cell_type":"code","source":"#Standard library imports   \nimport os\nimport sys\nimport shutil\nimport warnings\nfrom pprint import pprint\nwarnings.filterwarnings(\"ignore\")\n\n#third-party library imports\nimport numpy as np\nimport pandas as pd\nimport yaml\n\n#Torch an PTLightning imports\nimport torch\nimport torch.nn as nn\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, ReduceLROnPlateau\nfrom torch.utils.data import DataLoader\nfrom torchmetrics.functional import dice\nimport torchvision.transforms as T\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, TQDMProgressBar","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-08T18:12:17.910876Z","iopub.execute_input":"2023-08-08T18:12:17.91129Z","iopub.status.idle":"2023-08-08T18:12:33.757609Z","shell.execute_reply.started":"2023-08-08T18:12:17.911257Z","shell.execute_reply":"2023-08-08T18:12:33.756545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- [X] Define paths to everything\n  -  Pretrained Pytorch models: [Pretrained-models](https://github.com/Cadene/pretrained-models.pytorch/tree/master) This repository aims to facilitate the reproduction of research paper results, particularly in transfer learning setups, and provides access to pretrained Convolutional Neural Networks (ConvNets) through a user-friendly interface/API inspired by torchvision.\n  -  EfficientNet-Pytorch: [EfficientNetV2](https://github.com/lukemelas/EfficientNet-PyTorch) is a new family of convolutional networks that have faster training speed and better parameter efficiency than previous models. \n  -  Segmentation Models: [smp-github](https://github.com/qubvel/segmentation_models.pytorch) This library offers a high-level API for neural network creation in just two lines.\n  -  Pretrained RestNet: [timm-pretrained-restnet](https://www.kaggle.com/datasets/ar90ngas/timm-pretrained-resnest) A dataset with weights of a pretrained restnet in timm format (Torch Image Model)\n  -  Other paths to work with\n\n","metadata":{}},{"cell_type":"code","source":"pretrained_models_path = \"../input/pretrained-models-pytorch\"\nefficientnet_path = \"../input/efficientnet-pytorch\"\nsmp_github_path = \"/kaggle/input/smp-github/segmentation_models.pytorch-master\"\ntimm_pretrained_resnest_path = \"/kaggle/input/timm-pretrained-resnest/resnest/\"\nmodel_source_path = \"/kaggle/input/timm-pretrained-resnest/resnest/gluon_resnest26-50eb607c.pth\"\nmodel_destination_path = \"/root/.cache/torch/hub/checkpoints/gluon_resnest26-50eb607c.pth\"","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-08T18:12:33.759757Z","iopub.execute_input":"2023-08-08T18:12:33.760066Z","iopub.status.idle":"2023-08-08T18:12:33.765348Z","shell.execute_reply.started":"2023-08-08T18:12:33.760041Z","shell.execute_reply":"2023-08-08T18:12:33.764433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- [X] Import custom libraries and models\n\nAdd the needed paths to env variables, and import the segmentation models. ","metadata":{}},{"cell_type":"code","source":"sys.path.extend([\n    pretrained_models_path,\n    efficientnet_path,\n    smp_github_path,\n    timm_pretrained_resnest_path\n])\n\nimport segmentation_models_pytorch as smp","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-08T18:12:33.766928Z","iopub.execute_input":"2023-08-08T18:12:33.767592Z","iopub.status.idle":"2023-08-08T18:12:36.15925Z","shell.execute_reply.started":"2023-08-08T18:12:33.767543Z","shell.execute_reply":"2023-08-08T18:12:36.158276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Copy the model from input directory to a working directory ","metadata":{}},{"cell_type":"code","source":"destination_dir = os.path.dirname(model_destination_path)\nos.makedirs(destination_dir, exist_ok=True)\n\nshutil.copyfile(model_source_path, model_destination_path)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-07T01:23:33.182319Z","iopub.execute_input":"2023-08-07T01:23:33.184551Z","iopub.status.idle":"2023-08-07T01:23:34.155707Z","shell.execute_reply.started":"2023-08-07T01:23:33.184522Z","shell.execute_reply":"2023-08-07T01:23:34.154796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### YAML configuration \n- [X] Save YAML file with configurations\nYAML file containing configuration for the model training.\n\n**Paths** for loading the data and to save the models, random integer **seed**, **batch size** (train and test), and how often to **refresh progress bar** (in steps). Value 0 disables progress bar. Ignored when a custom callback is passed to callbacks.\n```yaml\n\ndata_path : \"/kaggle/input/contrails-images-ash-color\"\noutput_dir: \"models\"\n\nseed: 0xDEADBEEF\n\ntrain_batch_size: 64\nvalid_batch_size: 128\nworkers: 2\n\nprogress_bar_refresh_rate: 1\n```\n**Early stop** parameters see [Early stop class docs](https://lightning.ai/docs/pytorch/2.0.6/api/pytorch_lightning.callbacks.early_stopping.html). \n```yaml\nearly_stop:\n    monitor: \"val_loss\"\n    mode: \"min\"\n    patience: 999\n    verbose: 1\n```\n**Trainer** parameters see [Trainer class docs](https://lightning.ai/docs/pytorch/stable/common/trainer.html)\n```yaml\ntrainer:\n    max_epochs: 29\n    min_epochs: 27\n    enable_progress_bar: True\n    precision: \"16-mixed\"\n    devices: 2\n```\nThis **model** configuration is for a segmentation task with the Unet model, using a resnest26d encoder, specific optimization and scheduler settings.\n```yaml\nmodel:\n    seg_model: \"Unet\"\n    encoder_name: \"timm-resnest26d\"\n    loss_smooth: 1.0\n    image_size: 384\n    optimizer_params:\n        lr: 0.0005\n        weight_decay: 0.01 #0.0\n    scheduler:\n        name: \"CosineAnnealingLR\"\n        params:\n            CosineAnnealingLR:\n                T_max: 2\n                eta_min: 1.0e-6\n                last_epoch: -1\n            ReduceLROnPlateau:\n                mode: \"min\"\n                factor: 0.31622776602\n                patience: 4\n                verbose: True\n```\nThis info will be saved in **config.yaml** file.","metadata":{"execution":{"iopub.status.busy":"2023-07-31T23:03:33.088999Z","iopub.execute_input":"2023-07-31T23:03:33.089406Z","iopub.status.idle":"2023-07-31T23:03:33.104314Z","shell.execute_reply.started":"2023-07-31T23:03:33.089376Z","shell.execute_reply":"2023-07-31T23:03:33.099165Z"}}},{"cell_type":"code","source":"%%writefile config.yaml\n\ndata_path: \"/kaggle/input/contrails-images-ash-color\"\noutput_dir: \"models\"\n\nseed: 0xDEADBEEF\n\ntrain_batch_size: 32\nvalid_batch_size: 128\nworkers: 2\n\nprogress_bar_refresh_rate: 1\n\nearly_stop:\n    monitor: \"val_loss\"\n    mode: \"min\"\n    patience: 999\n    verbose: 1\n\ntrainer:\n    max_epochs: 29\n    min_epochs: 27\n    enable_progress_bar: True\n    precision: \"16-mixed\"\n    devices: 2\n\nmodel:\n    seg_model: \"Unet\"\n    encoder_name: \"timm-resnest26d\"\n    loss_smooth: 1.0\n    image_size: 384\n    optimizer_params:\n        lr: 0.0005\n        weight_decay: 0.01 #0.0\n    scheduler:\n        name: \"CosineAnnealingLR\"\n        params:\n            CosineAnnealingLR:\n                T_max: 2\n                eta_min: 1.0e-6\n                last_epoch: -1\n            ReduceLROnPlateau:\n                mode: \"min\"\n                factor: 0.31622776602\n                patience: 4\n                verbose: True","metadata":{"execution":{"iopub.status.busy":"2023-08-07T01:23:34.156997Z","iopub.execute_input":"2023-08-07T01:23:34.157342Z","iopub.status.idle":"2023-08-07T01:23:34.165142Z","shell.execute_reply.started":"2023-08-07T01:23:34.157309Z","shell.execute_reply":"2023-08-07T01:23:34.16426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Loading and Model Training \n\nAt this point our environment is ready to get everything running. now we need to create the classes for loading the data and execute the model training. Some of the previous configurations are the resultos of iterations to this point. \n\nWe will be mainly using 3 classes for this: \n- [ ] Dataset  ``` torch.utils.data.Dataset``` The Dataset retrieves our dataset’s features and labels one sample at a time.\n\n- [ ] DataLoader  ``` torch.utils.data.DataLoader``` When we train a model, we usually give it groups of examples called \"minibatches.\" To prevent the model from memorizing the data too much, we mix up the data after each round of training, and we use Python's multiprocessing to make getting the data faster.The DataLoader provides a user-friendly interface that simplifies the process by abstracting away the complex task in a simple API.\n\n- [ ] Lightning Module ``` pytorch_lightning.LightningModule``` A module for standarizing model training. ","metadata":{}},{"cell_type":"markdown","source":"**Loading the data from disk:**\n- [X] Dataset\nWe will need to define 3 methods, the constructor, get item and length. \nAlso we need to have two different Dataset objects as we are using two different datasets one for training and one for validation.\nIn the second one we will need to define extra functions to reduce the data shift from the data and convert validation images to equivalent formatof the training dataset images. ","metadata":{}},{"cell_type":"code","source":"class ContrailsDataset(torch.utils.data.Dataset):\n    \"\"\"\n    A custom PyTorch Dataset class for loading and \n    processing contrails data.\n\n    Args:\n        data (pandas.DataFrame): DataFrame containing the paths \n         to the contrails data and labels.\n        image_size (int, optional): The size to which images \n         will be resized. Default is 256.\n        train (bool, optional): Flag indicating if the dataset \n         is for training. Default is True.\n    \"\"\"\n\n    def __init__(self, data, image_size=256, train=True):\n        self.dataframe = data\n        self.is_train = train\n        self.normalize_image = T.Normalize(\n            (0.485, 0.456, 0.406),\n            (0.229, 0.224, 0.225))\n        self.image_size = image_size\n        if image_size != 256:\n            self.resize_image = T.transforms.Resize(image_size)\n\n    def __getitem__(self, index):\n        \"\"\"\n        Retrieves and processes an item from the dataset.\n\n        Args:\n            index (int): Index of the item to retrieve.\n\n        Returns:\n            tuple: A tuple containing the processed image \n             (torch.Tensor) and its label (torch.Tensor).\n        \"\"\"\n        row = self.dataframe.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        label = torch.tensor(label)\n\n        img = torch.tensor(\n            np.reshape(\n                img, (256, 256, 3))).to(torch.float32).permute(2, 0, 1)\n\n        if self.image_size != 256:\n            img = self.resize_image(img)\n\n        img = self.normalize_image(img)\n\n        return img.float(), label.float()\n\n    def __len__(self):\n        \"\"\"\n        Returns the total number of items in the dataset.\n\n        Returns:\n            int: The number of items in the dataset.\n        \"\"\"\n        return len(self.dataframe)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:12:57.582925Z","iopub.execute_input":"2023-08-08T18:12:57.583276Z","iopub.status.idle":"2023-08-08T18:12:57.595374Z","shell.execute_reply.started":"2023-08-08T18:12:57.583247Z","shell.execute_reply":"2023-08-08T18:12:57.594449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- [X] Lightning Module: ","metadata":{}},{"cell_type":"code","source":"seg_models = {\n    \"Unet\": smp.Unet,\n    \"Unet++\": smp.UnetPlusPlus,\n    \"MAnet\": smp.MAnet,\n    \"Linknet\": smp.Linknet,\n    \"FPN\": smp.FPN,\n    \"PSPNet\": smp.PSPNet,\n    \"PAN\": smp.PAN,\n    \"DeepLabV3\": smp.DeepLabV3,\n    \"DeepLabV3+\": smp.DeepLabV3Plus,\n}\n\nclass SLightningModule(pl.LightningModule):\n    \"\"\"\n    A PyTorch Lightning module for segmentation tasks \n    using various models.\n\n    Args:\n        config (dict): A dictionary containing configuration \n        parameters.\n    \"\"\"\n\n    def __init__(self, config):\n        super().__init__()\n        self.config = config\n        self.model = seg_models[config[\"seg_model\"]](\n            encoder_name=config[\"encoder_name\"],\n            encoder_weights=\"imagenet\",\n            in_channels=3,\n            classes=1,\n            activation=None,\n        )\n        self.loss_module = smp.losses.DiceLoss(\n            mode=\"binary\", smooth=config[\"loss_smooth\"])\n        self.val_step_outputs = []\n        self.val_step_labels = []\n    \n    def forward(self, batch):\n        \"\"\"\n        Forward propagation of the model.\n\n        Args:\n            batch (torch.Tensor): Input batch of images.\n\n        Returns:\n            torch.Tensor: Predicted segmentation masks.\n        \"\"\"\n        preds = self.model(batch)\n        return preds\n\n    def configure_optimizers(self):\n        \"\"\"\n        Configures the optimizer and learning rate scheduler.\n\n        Returns:\n            dict: Dictionary containing the optimizer and \n             learning rate scheduler.\n        \"\"\"\n        optimizer = AdamW(\n            self.parameters(),\n            **self.config[\"optimizer_params\"])\n        scheduler_name = self.config[\"scheduler\"][\"name\"]\n        scheduler_params = self.config[\"scheduler\"][\"params\"]\n\n        if scheduler_name == \"CosineAnnealingLR\":\n            scheduler = CosineAnnealingLR(\n                optimizer, \n                **scheduler_params[\"CosineAnnealingLR\"])\n            lr_scheduler_dict = {\n                \"scheduler\": scheduler, \n                \"interval\": \"step\"}\n        elif scheduler_name == \"ReduceLROnPlateau\":\n            scheduler = ReduceLROnPlateau(\n                optimizer, \n                **scheduler_params[\"ReduceLROnPlateau\"])\n            lr_scheduler_dict = {\n                \"scheduler\": scheduler, \n                \"monitor\": \"val_loss\"}\n        opt_config = {\n            \"optimizer\": optimizer,\n            \"lr_scheduler\": lr_scheduler_dict}\n        return opt_config\n\n\n\n    def training_step(self, batch, batch_idx):\n        \"\"\"\n        Training step for the model.\n\n        Args:\n            batch (tuple): Input batch of images and labels.\n            batch_idx (int): Index of the current batch.\n\n        Returns:\n            torch.Tensor: The computed loss for the current step.\n        \"\"\"\n        imgs, labels = batch\n        preds = self.model(imgs)\n        if self.config[\"image_size\"] != 256:\n            preds = torch.nn.functional.interpolate(\n                preds, size=256, mode='bilinear')\n        loss = self.loss_module(preds, labels)\n        self.log(\"train_loss\", \n                 loss, \n                 on_step=True, \n                 on_epoch=True, \n                 prog_bar=True, \n                 batch_size=16)\n\n        for param_group in self.trainer.optimizers[0].param_groups:\n            lr = param_group[\"lr\"]\n        self.log(\"lr\",\n                 lr,\n                 on_step=True,\n                 on_epoch=False,\n                 prog_bar=True)\n\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        \"\"\"\n        Validation step for the model.\n\n        Args:\n            batch (tuple): Input batch of images and labels.\n            batch_idx (int): Index of the current batch.\n        \"\"\"\n        imgs, labels = batch\n        preds = self.model(imgs)\n        if self.config[\"image_size\"] != 256:\n            preds = torch.nn.functional.interpolate(\n                preds,\n                size=256,\n                mode='bilinear')\n        loss = self.loss_module(preds,labels)\n        self.log(\"val_loss\",\n                 loss, on_step=False,\n                 on_epoch=True, prog_bar=True)\n        self.val_step_outputs.append(preds)\n        self.val_step_labels.append(labels)\n\n\n    def on_validation_epoch_end(self):\n        \"\"\"\n        Called at the end of each validation epoch to \n         compute and log validation metrics.\n        \"\"\"\n        all_preds = torch.cat(self.val_step_outputs)\n        all_labels = torch.cat(self.val_step_labels)\n        all_preds = torch.sigmoid(all_preds)\n        self.val_step_outputs.clear()\n        self.val_step_labels.clear()\n        val_dice = dice(all_preds, all_labels.long())\n        self.log(\"val_dice\", val_dice,\n                 on_step=False, on_epoch=True,\n                 prog_bar=True)\n        if self.trainer.global_rank == 0:\n            print(f\"\\nEpoch: {self.current_epoch}\",\n                  flush=True)\n            \n","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-08T18:13:00.960763Z","iopub.execute_input":"2023-08-08T18:13:00.961113Z","iopub.status.idle":"2023-08-08T18:13:00.982352Z","shell.execute_reply.started":"2023-08-08T18:13:00.961085Z","shell.execute_reply":"2023-08-08T18:13:00.981312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(\"config.yaml\", \"r\") as file_obj:\n    config = yaml.safe_load(file_obj)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:12:43.291843Z","iopub.execute_input":"2023-08-08T18:12:43.292223Z","iopub.status.idle":"2023-08-08T18:12:43.306193Z","shell.execute_reply.started":"2023-08-08T18:12:43.292192Z","shell.execute_reply":"2023-08-08T18:12:43.305079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"All model settings are loaded from config file, and using pytorch lightning objects. \nMaintaining this code is easier as parameters and hyperparameteres are independant of the rest of the code. ","metadata":{}},{"cell_type":"code","source":"    \ndata_paths = {\n    \"contrails\": os.path.join(\n        config[\"data_path\"], \"contrails/\"),\n    \"train_df\": os.path.join(\n        config[\"data_path\"], \"train_df.csv\"),\n    \"valid_df\": os.path.join(\n        config[\"data_path\"], \"valid_df.csv\"),\n}\n\n\ntrain_df = pd.read_csv(data_paths[\"train_df\"])\nvalid_df = pd.read_csv(data_paths[\"valid_df\"])\n\ntrain_df[\"path\"] = data_paths[\"contrails\"] + train_df[\"record_id\"].astype(str) + \".npy\"\nvalid_df[\"path\"] = data_paths[\"contrails\"] + valid_df[\"record_id\"].astype(str) + \".npy\"\n\n\ndataset_train = ContrailsDataset(\n    train_df,config[\"model\"][\"image_size\"],train=True)\ndataset_validation = ContrailsDataset(\n    valid_df, config[\"model\"][\"image_size\"], train=False)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:13:11.753023Z","iopub.execute_input":"2023-08-08T18:13:11.75365Z","iopub.status.idle":"2023-08-08T18:13:11.821184Z","shell.execute_reply.started":"2023-08-08T18:13:11.753616Z","shell.execute_reply":"2023-08-08T18:13:11.820243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- [X] DataLoader: once created DataSet objects, Dataloader receives them as parameter for the contructor.","metadata":{"execution":{"iopub.status.busy":"2023-08-08T05:22:35.603876Z","iopub.execute_input":"2023-08-08T05:22:35.604241Z","iopub.status.idle":"2023-08-08T05:22:35.610876Z","shell.execute_reply.started":"2023-08-08T05:22:35.604194Z","shell.execute_reply":"2023-08-08T05:22:35.609532Z"}}},{"cell_type":"code","source":"data_loader_train = DataLoader(\n    dataset_train,\n    batch_size=config[\"train_batch_size\"],\n    shuffle=True,\n    num_workers=config[\"workers\"],\n)\ndata_loader_train = DataLoader(\n    dataset_train,\n    batch_size=config[\"train_batch_size\"],\n    shuffle=True,\n    num_workers=config[\"workers\"],\n)\ndata_loader_validation = DataLoader(\n    dataset_validation,\n    batch_size=config[\"valid_batch_size\"],\n    shuffle=False,\n    num_workers=config[\"workers\"],\n)\ncheckpoint_callback = ModelCheckpoint(\n    save_weights_only=True,\n    monitor=\"val_dice\",\n    dirpath=config[\"output_dir\"],\n    mode=\"max\",\n    filename=\"model\",\n    save_top_k=1,\n    verbose=1,\n)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:13:21.924753Z","iopub.execute_input":"2023-08-08T18:13:21.925109Z","iopub.status.idle":"2023-08-08T18:13:21.942082Z","shell.execute_reply.started":"2023-08-08T18:13:21.925079Z","shell.execute_reply":"2023-08-08T18:13:21.941043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Model settings and parameters. come from config file.","metadata":{}},{"cell_type":"code","source":"progress_bar_callback = TQDMProgressBar(\n    refresh_rate=config[\"progress_bar_refresh_rate\"]\n)\n\nearly_stop_callback = EarlyStopping(**config[\"early_stop\"])\n\ntrainer = pl.Trainer(\n    callbacks=[checkpoint_callback, early_stop_callback,\n               progress_bar_callback],**config[\"trainer\"],\n)\n\nconfig[\"model\"][\"scheduler\"][\"params\"][\"CosineAnnealingLR\"][\"T_max\"] *= len(data_loader_train)/config[\"trainer\"][\"devices\"]","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:13:24.190532Z","iopub.execute_input":"2023-08-08T18:13:24.190892Z","iopub.status.idle":"2023-08-08T18:13:24.761848Z","shell.execute_reply.started":"2023-08-08T18:13:24.190862Z","shell.execute_reply":"2023-08-08T18:13:24.760865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Model is ready for **training** ","metadata":{}},{"cell_type":"code","source":"model = SLightningModule(config[\"model\"])\n\ntrainer.fit(model, data_loader_train, data_loader_validation)","metadata":{"execution":{"iopub.status.busy":"2023-08-07T01:23:34.227189Z","iopub.execute_input":"2023-08-07T01:23:34.227539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Inference (Submission)","metadata":{}},{"cell_type":"code","source":"batch_size = config[\"valid_batch_size\"]\ndevice    = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ntesting_data =  {\n    'data' : '/kaggle/input/google-research-identify-contrails-reduce-global-warming',\n    'data_test' : '/kaggle/input/google-research-identify-contrails-reduce-global-warming/test/'\n}    ","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:13:38.954879Z","iopub.execute_input":"2023-08-08T18:13:38.955225Z","iopub.status.idle":"2023-08-08T18:13:38.961358Z","shell.execute_reply.started":"2023-08-08T18:13:38.955196Z","shell.execute_reply":"2023-08-08T18:13:38.960501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"filenames = os.listdir(testing_data['data_test'])\ntest_df = pd.DataFrame(filenames, columns=['record_id'])\ntest_df['path'] = testing_data['data_test'] + test_df['record_id'].astype(str)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:14:04.893722Z","iopub.execute_input":"2023-08-08T18:14:04.894075Z","iopub.status.idle":"2023-08-08T18:14:04.903967Z","shell.execute_reply.started":"2023-08-08T18:14:04.894047Z","shell.execute_reply":"2023-08-08T18:14:04.903077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GContrailsDataset(torch.utils.data.Dataset):\n    \"\"\"\n    Custom dataset class for contrails data.\n\n    Args:\n        df (pandas.DataFrame): DataFrame containing paths \n        to contrails data.\n        image_size (int): Size to which the images will \n         be resized (default: 256).\n        train (bool): Flag to indicate if the dataset is \n         for training (default: True).\n    \"\"\"\n\n    _T11_BOUNDS = (243, 303)\n    _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n    _TDIFF_BOUNDS = (-4, 2)\n\n    def __init__(self, df, image_size=256, train=True):\n        self.df = df\n        self.train = train\n        self.df_idx = pd.DataFrame(\n            {'idx': os.listdir(testing_data['data_test'])})\n        self.normalize_image = T.Normalize(\n            (0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n        self.image_size = image_size\n        if image_size != 256:\n            self.resize_image = T.transforms.Resize(image_size)\n    \n    def read_record(self, directory):\n        \"\"\"\n        Reads contrails data from a directory and returns it as a dictionary.\n\n        Args:\n            directory (str): Path to the directory containing contrails data.\n\n        Returns:\n            dict: A dictionary containing the contrails data.\n        \"\"\"\n        record_data = {}\n        for x in [\"band_11\", \"band_14\", \"band_15\"]:\n            record_data[x] = np.load(os.path.join(directory, x + \".npy\"))\n        return record_data\n\n    def normalize_range(self, data, bounds):\n        \"\"\"\n        Maps data to the range [0, 1].\n\n        Args:\n            data (numpy.ndarray): Data to be normalized.\n            bounds (tuple): Lower and upper bounds of the desired range.\n\n        Returns:\n            numpy.ndarray: Normalized data.\n        \"\"\"\n        return (data - bounds[0]) / (bounds[1] - bounds[0])\n    \n    def get_false_color(self, record_data):\n        \"\"\"\n        Generates a false-color image from contrails data.\n\n        Args:\n            record_data (dict): A dictionary containing contrails data.\n\n        Returns:\n            numpy.ndarray: False-color image.\n        \"\"\"\n        N_TIMES_BEFORE = 4\n        r = self.normalize_range(\n            record_data[\"band_15\"] - record_data[\"band_14\"],\n            self._TDIFF_BOUNDS)\n        g = self.normalize_range(\n            record_data[\"band_14\"] - record_data[\"band_11\"],\n            self._CLOUD_TOP_TDIFF_BOUNDS)\n        b = self.normalize_range(\n            record_data[\"band_14\"], self._T11_BOUNDS)\n        false_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n        img = false_color[..., N_TIMES_BEFORE]\n        return img\n    \n    def __getitem__(self, index):\n        \"\"\"\n        Retrieves a sample and its label from the dataset.\n\n        Args:\n            index (int): Index of the sample to be retrieved.\n\n        Returns:\n            torch.Tensor: The sample as a tensor.\n            torch.Tensor: The label as a tensor.\n        \"\"\"\n        row = self.df.iloc[index]\n        con_path = row.path\n        data = self.read_record(con_path)    \n        \n        img = self.get_false_color(data)\n        \n        img = torch.tensor(img).to(torch.float32).permute(2, 0, 1)\n        \n        if self.image_size != 256:\n            img = self.resize_image(img)\n        \n        img = self.normalize_image(img)\n        \n        image_id = int(self.df_idx.iloc[index]['idx'])\n            \n        return img, torch.tensor(image_id)\n    \n    def __len__(self):\n        \"\"\"\n        Returns the total number of samples in the dataset.\n\n        Returns:\n            int: The number of samples in the dataset.\n        \"\"\"\n        return len(self.df)\n","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:15:22.934547Z","iopub.execute_input":"2023-08-08T18:15:22.934962Z","iopub.status.idle":"2023-08-08T18:15:22.951572Z","shell.execute_reply.started":"2023-08-08T18:15:22.934932Z","shell.execute_reply":"2023-08-08T18:15:22.950422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = GContrailsDataset(\n        test_df,\n        config[\"model\"][\"image_size\"],\n        train = False\n    )\n \ntest_dataloader = DataLoader(\n    test_dataset, batch_size=batch_size,\n    num_workers = 1)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:15:27.719144Z","iopub.execute_input":"2023-08-08T18:15:27.719527Z","iopub.status.idle":"2023-08-08T18:15:27.728233Z","shell.execute_reply.started":"2023-08-08T18:15:27.719497Z","shell.execute_reply":"2023-08-08T18:15:27.727163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LightningModule(pl.LightningModule):\n\n    def __init__(self):\n        super().__init__()\n        self.model = smp.Unet(encoder_name=\"timm-resnest26d\",\n                              encoder_weights=None,\n                              in_channels=3,\n                              classes=1,\n                              activation=None,\n                              )\n\n    def forward(self, batch):\n        return self.model(batch)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:15:37.363676Z","iopub.execute_input":"2023-08-08T18:15:37.364025Z","iopub.status.idle":"2023-08-08T18:15:37.370654Z","shell.execute_reply.started":"2023-08-08T18:15:37.363996Z","shell.execute_reply":"2023-08-08T18:15:37.369334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = LightningModule().load_from_checkpoint(\n    \"/kaggle/working/models/model.ckpt\")\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(device)\nmodel.eval()\nmodel.zero_grad()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-08T18:15:44.265436Z","iopub.execute_input":"2023-08-08T18:15:44.265841Z","iopub.status.idle":"2023-08-08T18:15:50.931361Z","shell.execute_reply.started":"2023-08-08T18:15:44.26581Z","shell.execute_reply":"2023-08-08T18:15:50.930424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_encode(x, fg_val=1):\n    \"\"\"\n    Args:\n        x:  numpy array of shape (height, width), \n        1 - mask,\n        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\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","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:15:54.307102Z","iopub.execute_input":"2023-08-08T18:15:54.30748Z","iopub.status.idle":"2023-08-08T18:15:54.315201Z","shell.execute_reply.started":"2023-08-08T18:15:54.307444Z","shell.execute_reply":"2023-08-08T18:15:54.314169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_list = []\nMASK_THRESHOLD = 0.45\nfor i, data in enumerate(test_dataloader):\n    images, image_id = data\n    images = images.to(device)\n    \n    with torch.no_grad():\n        predicted_mask = model(images[:, :, :, :])\n    \n    if config[\"model\"][\"image_size\"] != 256:\n        predicted_mask = torch.nn.functional.interpolate(\n            predicted_mask, size=256, mode='bilinear')\n    \n    predicted_mask = torch.sigmoid(predicted_mask).cpu().detach().numpy()\n    \n    predicted_mask_with_threshold = (predicted_mask[:, 0, :, :] > MASK_THRESHOLD).astype(int)\n    \n    for img_num, current_mask in enumerate(predicted_mask_with_threshold):\n        current_image_id = image_id[img_num].item()\n        encoded_pixels = list_to_string(rle_encode(current_mask))\n        submission_list.append({'record_id': int(current_image_id), 'encoded_pixels': encoded_pixels})\n\nsubmission = pd.DataFrame(submission_list)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:15:59.087067Z","iopub.execute_input":"2023-08-08T18:15:59.087428Z","iopub.status.idle":"2023-08-08T18:18:41.631365Z","shell.execute_reply.started":"2023-08-08T18:15:59.087379Z","shell.execute_reply":"2023-08-08T18:18:41.630068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(submission)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:20:42.426536Z","iopub.execute_input":"2023-08-08T18:20:42.427159Z","iopub.status.idle":"2023-08-08T18:20:42.438139Z","shell.execute_reply.started":"2023-08-08T18:20:42.427119Z","shell.execute_reply":"2023-08-08T18:20:42.437065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T18:20:46.372838Z","iopub.execute_input":"2023-08-08T18:20:46.373867Z","iopub.status.idle":"2023-08-08T18:20:46.403872Z","shell.execute_reply.started":"2023-08-08T18:20:46.373826Z","shell.execute_reply":"2023-08-08T18:20:46.402979Z"},"trusted":true},"execution_count":null,"outputs":[]}]}