{"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":{"papermill":{"duration":0.011008,"end_time":"2023-07-23T07:34:09.076324","exception":false,"start_time":"2023-07-23T07:34:09.065316","status":"completed"},"tags":[]}},{"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.optim as optim\nfrom torch import Tensor\nfrom torch.utils.data import TensorDataset, DataLoader\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","metadata":{"papermill":{"duration":36.388559,"end_time":"2023-07-23T07:34:45.475635","exception":false,"start_time":"2023-07-23T07:34:09.087076","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:09.155187Z","iopub.execute_input":"2023-07-23T12:49:09.15566Z","iopub.status.idle":"2023-07-23T12:49:20.957861Z","shell.execute_reply.started":"2023-07-23T12:49:09.155618Z","shell.execute_reply":"2023-07-23T12:49:20.95655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# load data","metadata":{"papermill":{"duration":0.010927,"end_time":"2023-07-23T07:34:45.497547","exception":false,"start_time":"2023-07-23T07:34:45.48662","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":0.019518,"end_time":"2023-07-23T07:34:45.527614","exception":false,"start_time":"2023-07-23T07:34:45.508096","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:20.960827Z","iopub.execute_input":"2023-07-23T12:49:20.961203Z","iopub.status.idle":"2023-07-23T12:49:20.966125Z","shell.execute_reply.started":"2023-07-23T12:49:20.961165Z","shell.execute_reply":"2023-07-23T12:49:20.965202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"configs = Namespace(\n    base_dir= Path(BASE_DIR),\n    batch_size= 16,\n    num_worers= 2,\n    shuffle= True,\n    epochs= 10,\n    lr=1e-3,\n    device= torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"),\n)","metadata":{"papermill":{"duration":0.070777,"end_time":"2023-07-23T07:34:45.608682","exception":false,"start_time":"2023-07-23T07:34:45.537905","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:20.967696Z","iopub.execute_input":"2023-07-23T12:49:20.968026Z","iopub.status.idle":"2023-07-23T12:49:21.008149Z","shell.execute_reply.started":"2023-07-23T12:49:20.967995Z","shell.execute_reply":"2023-07-23T12:49:21.007149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"papermill":{"duration":0.31094,"end_time":"2023-07-23T07:34:45.93001","exception":false,"start_time":"2023-07-23T07:34:45.61907","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:21.012402Z","iopub.execute_input":"2023-07-23T12:49:21.013241Z","iopub.status.idle":"2023-07-23T12:49:21.341373Z","shell.execute_reply.started":"2023-07-23T12:49:21.013209Z","shell.execute_reply":"2023-07-23T12:49:21.340298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.head()","metadata":{"papermill":{"duration":0.029456,"end_time":"2023-07-23T07:34:45.970653","exception":false,"start_time":"2023-07-23T07:34:45.941197","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:21.344587Z","iopub.execute_input":"2023-07-23T12:49:21.344893Z","iopub.status.idle":"2023-07-23T12:49:21.36085Z","shell.execute_reply.started":"2023-07-23T12:49:21.344867Z","shell.execute_reply":"2023-07-23T12:49:21.359955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, train=True):\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\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            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":{"papermill":{"duration":0.028622,"end_time":"2023-07-23T07:34:46.010706","exception":false,"start_time":"2023-07-23T07:34:45.982084","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:21.363942Z","iopub.execute_input":"2023-07-23T12:49:21.364594Z","iopub.status.idle":"2023-07-23T12:49:21.380547Z","shell.execute_reply.started":"2023-07-23T12:49:21.364562Z","shell.execute_reply":"2023-07-23T12:49:21.379621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = ContrailsDataset(\n        train_df,\n        train = True\n    )\n\ntrain_dl = DataLoader(train_ds, batch_size=configs.batch_size, num_workers = configs.num_worers)\n\nval_ds = ContrailsDataset(\n        val_df,\n        train = True\n    )\n\nval_dl = DataLoader(val_ds, batch_size=configs.batch_size, num_workers = configs.num_worers)","metadata":{"papermill":{"duration":0.019813,"end_time":"2023-07-23T07:34:46.041526","exception":false,"start_time":"2023-07-23T07:34:46.021713","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:21.381863Z","iopub.execute_input":"2023-07-23T12:49:21.382593Z","iopub.status.idle":"2023-07-23T12:49:21.394272Z","shell.execute_reply.started":"2023-07-23T12:49:21.382557Z","shell.execute_reply":"2023-07-23T12:49:21.393369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fetch the one batch from the test dataloader\nfor xb, yb in train_dl:\n    print(xb.shape)\n    print(yb.shape)\n    # check the device\n    print(xb.device)\n    break","metadata":{"papermill":{"duration":2.963765,"end_time":"2023-07-23T07:34:49.016131","exception":false,"start_time":"2023-07-23T07:34:46.052366","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:21.397358Z","iopub.execute_input":"2023-07-23T12:49:21.398549Z","iopub.status.idle":"2023-07-23T12:49:24.765307Z","shell.execute_reply.started":"2023-07-23T12:49:21.398518Z","shell.execute_reply":"2023-07-23T12:49:24.764076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# delete the variables to free up the memory\ndel train_df, val_df, train_ds, val_ds","metadata":{"papermill":{"duration":0.02015,"end_time":"2023-07-23T07:34:49.047974","exception":false,"start_time":"2023-07-23T07:34:49.027824","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:24.766908Z","iopub.execute_input":"2023-07-23T12:49:24.767293Z","iopub.status.idle":"2023-07-23T12:49:24.774243Z","shell.execute_reply.started":"2023-07-23T12:49:24.767253Z","shell.execute_reply":"2023-07-23T12:49:24.772216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Architecture","metadata":{"papermill":{"duration":0.010347,"end_time":"2023-07-23T07:34:49.09746","exception":false,"start_time":"2023-07-23T07:34:49.087113","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n\n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, 3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, 3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n    \nclass Down(nn.Module):\n    \n    def __init__(self, in_channels, out_channels):\n        super(Down, self).__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool2d(2),\n            DoubleConv(in_channels, out_channels)\n        )\n\n    def forward(self, x):\n        return self.maxpool_conv(x)\n    \nclass Up(nn.Module):\n        \n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super(Up, self).__init__()\n        \n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        else:\n            self.up = nn.ConvTranspose2d(in_channels // 2, in_channels // 2, 2, stride=2)\n        \n        self.conv = DoubleConv(in_channels, out_channels)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        \n        # input is CHW\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n        \n        x1 = nn.functional.pad(x1, [diffX // 2, diffX - diffX // 2,\n                                    diffY // 2, diffY - diffY // 2])\n        \n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n    \nclass UNet(nn.Module):\n    def __init__(self, n_channels, n_classes, bilinear=True):\n        super(UNet, self).__init__()\n        self.n_channels = n_channels\n        self.n_classes = n_classes\n        self.bilinear = bilinear\n        \n        self.inc = DoubleConv(n_channels, 64)\n        self.down1 = Down(64, 128)\n        self.down2 = Down(128, 256)\n        self.down3 = Down(256, 512)\n        self.down4 = Down(512, 512)\n        self.up1 = Up(1024, 256, bilinear)\n        self.up2 = Up(512, 128, bilinear)\n        self.up3 = Up(256, 64, bilinear)\n        self.up4 = Up(128, 64, bilinear)\n        self.outc = nn.Conv2d(64, n_classes, 1)\n        \n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        x = self.up1(x5, x4)\n        x = self.up2(x, x3)\n        x = self.up3(x, x2)\n        x = self.up4(x, x1)\n        logits = self.outc(x)\n        return logits","metadata":{"papermill":{"duration":0.036497,"end_time":"2023-07-23T07:34:49.144533","exception":false,"start_time":"2023-07-23T07:34:49.108036","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:24.778521Z","iopub.execute_input":"2023-07-23T12:49:24.779159Z","iopub.status.idle":"2023-07-23T12:49:24.79909Z","shell.execute_reply.started":"2023-07-23T12:49:24.779125Z","shell.execute_reply":"2023-07-23T12:49:24.798142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(UNet(n_channels=3, n_classes=1).to(configs.device), input_size=(3, 256, 256))","metadata":{"papermill":{"duration":6.633612,"end_time":"2023-07-23T07:34:55.823321","exception":false,"start_time":"2023-07-23T07:34:49.189709","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:24.80071Z","iopub.execute_input":"2023-07-23T12:49:24.801419Z","iopub.status.idle":"2023-07-23T12:49:32.923867Z","shell.execute_reply.started":"2023-07-23T12:49:24.801387Z","shell.execute_reply":"2023-07-23T12:49:32.922761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{"papermill":{"duration":0.011194,"end_time":"2023-07-23T07:34:55.845628","exception":false,"start_time":"2023-07-23T07:34:55.834434","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":0.021294,"end_time":"2023-07-23T07:34:55.878166","exception":false,"start_time":"2023-07-23T07:34:55.856872","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:32.925539Z","iopub.execute_input":"2023-07-23T12:49:32.92593Z","iopub.status.idle":"2023-07-23T12:49:32.93351Z","shell.execute_reply.started":"2023-07-23T12:49:32.925895Z","shell.execute_reply":"2023-07-23T12:49:32.932464Z"},"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 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                self.model.train()\n                if i % 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            # 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                torch.save(self.model.state_dict(), \"model_checkpoint_e\" + str(e) + \".pt\")\n                with torch.no_grad():\n                    self.model.eval()\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":{"papermill":{"duration":0.030772,"end_time":"2023-07-23T07:34:55.919566","exception":false,"start_time":"2023-07-23T07:34:55.888794","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:49:32.935116Z","iopub.execute_input":"2023-07-23T12:49:32.93574Z","iopub.status.idle":"2023-07-23T12:49:32.955476Z","shell.execute_reply.started":"2023-07-23T12:49:32.935693Z","shell.execute_reply":"2023-07-23T12:49:32.954146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = False\n\nif train:\n    model = UNet(n_channels=3, n_classes=1)\n    model.to(configs.device)\n\n    criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor(100))\n    optimizer = optim.Adam(model.parameters(), lr=0.01)\n    lr_scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, 0.70)\n\n\n    trainer = MyTrainer(model, optimizer, criterion, lr_scheduler)\n    trainer.fit(train_dl, val_dl, epochs=configs.epochs)\nelse:\n    model = UNet(n_channels=3, n_classes=1)\n    model.load_state_dict(torch.load(os.path.join(\"/kaggle/input/contrails-trained-model/model_checkpoint_e7.pt\")))\n    model.eval()\n    model.to(configs.device)","metadata":{"papermill":{"duration":10097.237352,"end_time":"2023-07-23T10:23:13.167709","exception":false,"start_time":"2023-07-23T07:34:55.930357","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:06.475225Z","iopub.execute_input":"2023-07-23T12:51:06.475666Z","iopub.status.idle":"2023-07-23T12:51:07.242674Z","shell.execute_reply.started":"2023-07-23T12:51:06.475628Z","shell.execute_reply":"2023-07-23T12:51:07.241667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Progress","metadata":{"papermill":{"duration":0.01908,"end_time":"2023-07-23T10:23:13.208737","exception":false,"start_time":"2023-07-23T10:23:13.189657","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":0.574558,"end_time":"2023-07-23T10:23:13.801992","exception":false,"start_time":"2023-07-23T10:23:13.227434","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:10.138191Z","iopub.execute_input":"2023-07-23T12:51:10.138553Z","iopub.status.idle":"2023-07-23T12:51:10.144022Z","shell.execute_reply.started":"2023-07-23T12:51:10.138522Z","shell.execute_reply":"2023-07-23T12:51:10.142949Z"},"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":{"papermill":{"duration":0.321984,"end_time":"2023-07-23T10:23:14.148206","exception":false,"start_time":"2023-07-23T10:23:13.826222","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:11.09698Z","iopub.execute_input":"2023-07-23T12:51:11.09734Z","iopub.status.idle":"2023-07-23T12:51:11.104255Z","shell.execute_reply.started":"2023-07-23T12:51:11.097309Z","shell.execute_reply":"2023-07-23T12:51:11.103247Z"},"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":{"papermill":{"duration":0.356774,"end_time":"2023-07-23T10:23:14.524944","exception":false,"start_time":"2023-07-23T10:23:14.16817","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:11.91276Z","iopub.execute_input":"2023-07-23T12:51:11.913334Z","iopub.status.idle":"2023-07-23T12:51:11.920576Z","shell.execute_reply.started":"2023-07-23T12:51:11.9133Z","shell.execute_reply":"2023-07-23T12:51:11.919661Z"},"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":{"papermill":{"duration":0.332341,"end_time":"2023-07-23T10:23:14.877573","exception":false,"start_time":"2023-07-23T10:23:14.545232","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:12.240191Z","iopub.execute_input":"2023-07-23T12:51:12.241006Z","iopub.status.idle":"2023-07-23T12:51:12.24682Z","shell.execute_reply.started":"2023-07-23T12:51:12.240965Z","shell.execute_reply":"2023-07-23T12:51:12.245525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimum Threshold","metadata":{"papermill":{"duration":0.020636,"end_time":"2023-07-23T10:23:14.919072","exception":false,"start_time":"2023-07-23T10:23:14.898436","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":0.328821,"end_time":"2023-07-23T10:23:15.268136","exception":false,"start_time":"2023-07-23T10:23:14.939315","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:12.900905Z","iopub.execute_input":"2023-07-23T12:51:12.901966Z","iopub.status.idle":"2023-07-23T12:51:12.912615Z","shell.execute_reply.started":"2023-07-23T12:51:12.901924Z","shell.execute_reply":"2023-07-23T12:51:12.911672Z"},"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":{"papermill":{"duration":84.819053,"end_time":"2023-07-23T10:24:40.108135","exception":false,"start_time":"2023-07-23T10:23:15.289082","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:13.114474Z","iopub.execute_input":"2023-07-23T12:51:13.114812Z","iopub.status.idle":"2023-07-23T12:51:13.120937Z","shell.execute_reply.started":"2023-07-23T12:51:13.114782Z","shell.execute_reply":"2023-07-23T12:51:13.119925Z"},"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.97\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.98","metadata":{"papermill":{"duration":305.481414,"end_time":"2023-07-23T10:29:45.610719","exception":false,"start_time":"2023-07-23T10:24:40.129305","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:13.911351Z","iopub.execute_input":"2023-07-23T12:51:13.911756Z","iopub.status.idle":"2023-07-23T12:51:13.918955Z","shell.execute_reply.started":"2023-07-23T12:51:13.9117Z","shell.execute_reply":"2023-07-23T12:51:13.917884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    df_threshold_data.tail(), df_threshold_data.shape","metadata":{"papermill":{"duration":0.036704,"end_time":"2023-07-23T10:29:45.667906","exception":false,"start_time":"2023-07-23T10:29:45.631202","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:14.482259Z","iopub.execute_input":"2023-07-23T12:51:14.482767Z","iopub.status.idle":"2023-07-23T12:51:14.487702Z","shell.execute_reply.started":"2023-07-23T12:51:14.482734Z","shell.execute_reply":"2023-07-23T12:51:14.486669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plt.figure(figsize=(10,10))\n# sns.lineplot(data=df_threshold_data, x='Threshold', y='Dice Score')\n# plt.axhline(y=best_dice_score, color='green')\n# plt.axvline(x=optim_threshold, color='green')\n# plt.text(-0.02, best_dice_score * 0.96, f'{best_dice_score:.3f}', va='center', ha='left', color='green')\n# plt.text(optim_threshold - 0.01, 0.02, f'{optim_threshold}', va='center', ha='right', color='green')\n# plt.ylim(bottom=0)\n# plt.title('Threshold vs Dice Score')\n# plt.show()","metadata":{"papermill":{"duration":0.028037,"end_time":"2023-07-23T10:29:45.717182","exception":false,"start_time":"2023-07-23T10:29:45.689145","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:18.988519Z","iopub.execute_input":"2023-07-23T12:51:18.988914Z","iopub.status.idle":"2023-07-23T12:51:18.993497Z","shell.execute_reply.started":"2023-07-23T12:51:18.988881Z","shell.execute_reply":"2023-07-23T12:51:18.992545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sigmoid(x):\n    return 1 / (1 + np.exp(-x))\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":{"papermill":{"duration":0.030433,"end_time":"2023-07-23T10:29:45.767819","exception":false,"start_time":"2023-07-23T10:29:45.737386","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:19.71819Z","iopub.execute_input":"2023-07-23T12:51:19.720022Z","iopub.status.idle":"2023-07-23T12:51:29.32744Z","shell.execute_reply.started":"2023-07-23T12:51:19.719982Z","shell.execute_reply":"2023-07-23T12:51:29.326346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{"papermill":{"duration":0.020181,"end_time":"2023-07-23T10:29:45.808458","exception":false,"start_time":"2023-07-23T10:29:45.788277","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":0.038504,"end_time":"2023-07-23T10:29:45.867411","exception":false,"start_time":"2023-07-23T10:29:45.828907","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:29.329272Z","iopub.execute_input":"2023-07-23T12:51:29.329595Z","iopub.status.idle":"2023-07-23T12:51:29.339526Z","shell.execute_reply.started":"2023-07-23T12:51:29.329566Z","shell.execute_reply":"2023-07-23T12:51:29.338513Z"},"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":{"papermill":{"duration":0.031021,"end_time":"2023-07-23T10:29:45.919395","exception":false,"start_time":"2023-07-23T10:29:45.888374","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:29.340968Z","iopub.execute_input":"2023-07-23T12:51:29.341578Z","iopub.status.idle":"2023-07-23T12:51:29.351543Z","shell.execute_reply.started":"2023-07-23T12:51:29.341545Z","shell.execute_reply":"2023-07-23T12:51:29.350657Z"},"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":{"papermill":{"duration":0.152671,"end_time":"2023-07-23T10:29:46.092653","exception":false,"start_time":"2023-07-23T10:29:45.939982","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:29.35404Z","iopub.execute_input":"2023-07-23T12:51:29.35451Z","iopub.status.idle":"2023-07-23T12:51:29.377484Z","shell.execute_reply.started":"2023-07-23T12:51:29.354478Z","shell.execute_reply":"2023-07-23T12:51:29.376613Z"},"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":{"papermill":{"duration":0.405453,"end_time":"2023-07-23T10:29:46.518762","exception":false,"start_time":"2023-07-23T10:29:46.113309","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:29.378879Z","iopub.execute_input":"2023-07-23T12:51:29.379181Z","iopub.status.idle":"2023-07-23T12:51:29.695548Z","shell.execute_reply.started":"2023-07-23T12:51:29.379151Z","shell.execute_reply":"2023-07-23T12:51:29.694288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.head()","metadata":{"papermill":{"duration":0.03507,"end_time":"2023-07-23T10:29:46.574916","exception":false,"start_time":"2023-07-23T10:29:46.539846","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:29.697289Z","iopub.execute_input":"2023-07-23T12:51:29.698218Z","iopub.status.idle":"2023-07-23T12:51:29.709496Z","shell.execute_reply.started":"2023-07-23T12:51:29.698185Z","shell.execute_reply":"2023-07-23T12:51:29.708499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"papermill":{"duration":0.077368,"end_time":"2023-07-23T10:29:46.687268","exception":false,"start_time":"2023-07-23T10:29:46.6099","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-07-23T12:51:29.711195Z","iopub.execute_input":"2023-07-23T12:51:29.711986Z","iopub.status.idle":"2023-07-23T12:51:29.722321Z","shell.execute_reply.started":"2023-07-23T12:51:29.711954Z","shell.execute_reply":"2023-07-23T12:51:29.721391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}