{"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":"## Import Libraries","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nimport os\nimport random\nimport math\nfrom collections import defaultdict\nimport cv2\nimport skimage\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport torch\nfrom torch import nn\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nfrom transformers import get_cosine_schedule_with_warmup\n\ntorch.__version__","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-19T10:54:00.358451Z","iopub.execute_input":"2023-06-19T10:54:00.360333Z","iopub.status.idle":"2023-06-19T10:54:14.010896Z","shell.execute_reply.started":"2023-06-19T10:54:00.360306Z","shell.execute_reply":"2023-06-19T10:54:14.01004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.auto import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:54:38.008402Z","iopub.execute_input":"2023-06-19T10:54:38.0091Z","iopub.status.idle":"2023-06-19T10:54:38.013703Z","shell.execute_reply.started":"2023-06-19T10:54:38.009063Z","shell.execute_reply":"2023-06-19T10:54:38.012831Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install segmentation-models-pytorch\nimport segmentation_models_pytorch as smp","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-19T10:54:45.193949Z","iopub.execute_input":"2023-06-19T10:54:45.194381Z","iopub.status.idle":"2023-06-19T10:55:04.411501Z","shell.execute_reply.started":"2023-06-19T10:54:45.194349Z","shell.execute_reply":"2023-06-19T10:55:04.410487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Preparation","metadata":{}},{"cell_type":"code","source":"class Config:\n    train = True\n    thr = 0.02\n    num_epochs = 10\n    num_classes = 1\n    batch_size = 32\n    seed = 42\n    \n    encoder = 'efficientnet-b0'\n    pretrained = True\n    weights = 'imagenet'\n    classes = ['contrail']\n    activation = nn.ReLU()\n    in_chans = 3\n    \n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    image_size = 256\n    warmup = 0\n    lr = 3e-4\n    \nclass Paths:\n    data_root = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\n    contrails = '/kaggle/input/contrails-images-ash-color/contrails/'\n    train_path = '/kaggle/input/contrails-images-ash-color/train_df.csv'\n    valid_path = '/kaggle/input/contrails-images-ash-color/valid_df.csv'","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:58:02.495075Z","iopub.execute_input":"2023-06-19T10:58:02.495478Z","iopub.status.idle":"2023-06-19T10:58:02.503786Z","shell.execute_reply.started":"2023-06-19T10:58:02.495444Z","shell.execute_reply":"2023-06-19T10:58:02.502776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=1234):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    \n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:55:54.136269Z","iopub.execute_input":"2023-06-19T10:55:54.136879Z","iopub.status.idle":"2023-06-19T10:55:54.142682Z","shell.execute_reply.started":"2023-06-19T10:55:54.136845Z","shell.execute_reply":"2023-06-19T10:55:54.141775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set_seed(9)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:55:55.699944Z","iopub.execute_input":"2023-06-19T10:55:55.700839Z","iopub.status.idle":"2023-06-19T10:55:55.710945Z","shell.execute_reply.started":"2023-06-19T10:55:55.700807Z","shell.execute_reply":"2023-06-19T10:55:55.709954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Import dataframes\ntrain_df = pd.read_csv(Paths.train_path)\nvalid_df = pd.read_csv(Paths.valid_path)\n\ntrain_df['path'] = Paths.contrails + train_df['record_id'].astype(str) + '.npy'\nvalid_df['path'] = Paths.contrails + valid_df['record_id'].astype(str) + '.npy'\n\ntrain_df.shape, valid_df.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:55:57.204444Z","iopub.execute_input":"2023-06-19T10:55:57.204817Z","iopub.status.idle":"2023-06-19T10:55:57.273858Z","shell.execute_reply.started":"2023-06-19T10:55:57.204787Z","shell.execute_reply":"2023-06-19T10:55:57.272828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, train=True):\n        \n        self.df = df\n        self.trn = train\n    \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        con_path = row.path\n        con = np.load(str(con_path))\n        \n        img = con[..., :-1]\n        label = con[..., -1]\n        \n        img = torch.tensor(img)\n        label = torch.tensor(label)\n        \n        img = img.permute(2, 0, 1)\n            \n        return img.float(), label.float()\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:55:59.813529Z","iopub.execute_input":"2023-06-19T10:55:59.813927Z","iopub.status.idle":"2023-06-19T10:55:59.821995Z","shell.execute_reply.started":"2023-06-19T10:55:59.813879Z","shell.execute_reply":"2023-06-19T10:55:59.820749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = ContrailsDataset(train_df, train=True)\nvalid_ds = ContrailsDataset(valid_df, train=False)\n\ntrain_dl = DataLoader(train_ds, batch_size=Config.batch_size , shuffle=True, num_workers = 2)    \nvalid_dl = DataLoader(valid_ds, batch_size=Config.batch_size, num_workers = 2)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:56:01.27788Z","iopub.execute_input":"2023-06-19T10:56:01.278274Z","iopub.status.idle":"2023-06-19T10:56:01.28479Z","shell.execute_reply.started":"2023-06-19T10:56:01.278243Z","shell.execute_reply":"2023-06-19T10:56:01.283796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, label = next(iter(train_dl))\nimg.shape, label.shape\nimg, label = next(iter(valid_dl))\nimg.shape, label.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:56:02.696465Z","iopub.execute_input":"2023-06-19T10:56:02.696825Z","iopub.status.idle":"2023-06-19T10:56:04.808044Z","shell.execute_reply.started":"2023-06-19T10:56:02.696793Z","shell.execute_reply":"2023-06-19T10:56:04.806948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def dice_coef(y_true, y_pred, thr=0.5, epsilon=0.001):\n    y_true = y_true.flatten()\n    y_pred = (y_pred>thr).astype(np.float32).flatten()\n    inter = (y_true*y_pred).sum()\n    den = y_true.sum() + y_pred.sum()\n    dice = ((2*inter+epsilon)/(den+epsilon))\n    return dice","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:56:07.013854Z","iopub.execute_input":"2023-06-19T10:56:07.014316Z","iopub.status.idle":"2023-06-19T10:56:07.022133Z","shell.execute_reply.started":"2023-06-19T10:56:07.014279Z","shell.execute_reply":"2023-06-19T10:56:07.021164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, cfg):\n        super(UNet, self).__init__()\n\n        self.cfg = cfg\n        self.training = True\n\n        # Create the U-Net model using the provided configuration\n        self.model = smp.Unet(\n            encoder_name=cfg.encoder,\n            encoder_weights=cfg.weights,\n            decoder_use_batchnorm=True,\n            classes=len(cfg.classes),\n            activation=None,  # Remove the activation argument from here\n        )\n\n        self.activation = cfg.activation  # Store the activation function separately\n\n        self.loss_fn = smp.losses.DiceLoss(mode='binary')  # Define the Dice loss function\n\n    def forward(self, imgs, targets):\n        x = imgs\n        y = targets\n\n        logits = self.model(x)  # Forward pass through the U-Net model\n        logits = self.activation(logits)  # Apply the activation function to the logits\n        loss = self.loss_fn(logits, y)  # Calculate the loss\n\n        return {\n            \"loss\": loss,  # Return the loss value\n            \"logits\": logits,  # Return the activated logits\n            \"logits_raw\": logits,  # Return the raw logits\n            \"target\": y  # Return the target values\n        }","metadata":{"execution":{"iopub.status.busy":"2023-06-19T11:00:56.171261Z","iopub.execute_input":"2023-06-19T11:00:56.171632Z","iopub.status.idle":"2023-06-19T11:00:56.179958Z","shell.execute_reply.started":"2023-06-19T11:00:56.171603Z","shell.execute_reply":"2023-06-19T11:00:56.179056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_step(model, dataloader, optimizer, device):\n    \n    model.train()\n    \n    train_losses = []\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    \n    for step, (X, y) in pbar:\n        \n        X, y = X.to(device), y.to(device)\n        torch.set_grad_enabled(True)\n        \n        output_dict = model(X, y)\n        loss = output_dict[\"loss\"]\n        train_losses.append(loss.item())\n        \n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n        \n        if scheduler is not None:\n            scheduler.step()\n    \n    train_loss = np.sum(train_losses)\n    \n    return train_loss","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:56:11.704878Z","iopub.execute_input":"2023-06-19T10:56:11.705693Z","iopub.status.idle":"2023-06-19T10:56:11.713457Z","shell.execute_reply.started":"2023-06-19T10:56:11.705662Z","shell.execute_reply":"2023-06-19T10:56:11.71233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_step(model, dataloader, device):\n    \n    model.eval()\n    torch.set_grad_enabled(False)\n    \n    val_data = defaultdict(list)\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid')\n    for step, (X, y) in pbar: \n        X, y = X.to(device), y.to(device)\n\n        output = model(X, y)\n        for key, val in output.items():\n            val_data[key] += [output[key]]\n\n    for key, val in output.items():\n        value = val_data[key]\n        if len(value[0].shape) == 0:\n            val_data[key] = torch.stack(value)\n        else:\n            val_data[key] = torch.cat(value, dim=0).cpu().detach().numpy()\n    \n    val_losses = val_data[\"loss\"].cpu().numpy()\n    val_loss = np.sum(val_losses)\n    \n    val_dice = dice_coef(val_data['target'], val_data['logits'])\n    \n    return val_loss, val_dice","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:56:13.273653Z","iopub.execute_input":"2023-06-19T10:56:13.274041Z","iopub.status.idle":"2023-06-19T10:56:13.283601Z","shell.execute_reply.started":"2023-06-19T10:56:13.274009Z","shell.execute_reply":"2023-06-19T10:56:13.282543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, train_dataloader, test_dataloader, optimizer, epochs, device):\n    results = {'train_loss': [],\n              'val_loss': [],\n              'val_dice': []}\n    for epoch in range(epochs):\n        \n        set_seed(Config.seed + epoch)\n        print(\"EPOCH:\", epoch)\n        \n        train_loss = train_step(model,\n                              train_dataloader,\n                              optimizer,\n                              device)\n        val_loss, val_dice = test_step(model,\n                            test_dataloader,\n                            device)\n        \n        train_loss = train_loss / len(train_ds)\n        val_loss = val_loss / len(valid_ds)\n        \n        print(f'Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | Val Dice: {val_dice:.4f}')\n        print(f\"Learning rate: {optimizer.param_groups[0]['lr']}\")\n        \n        results['train_loss'].append(train_loss)\n        results['val_loss'].append(val_loss)\n        results['val_dice'].append(val_dice)\n        \n        PATH = f\"epoch-{epoch}.pth\"\n        torch.save(model.state_dict(), PATH)\n\n    return results","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:56:15.794429Z","iopub.execute_input":"2023-06-19T10:56:15.794808Z","iopub.status.idle":"2023-06-19T10:56:15.802889Z","shell.execute_reply.started":"2023-06-19T10:56:15.794776Z","shell.execute_reply":"2023-06-19T10:56:15.801779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_optimizer(lr, params):\n    \n    model_optimizer = torch.optim.Adam(\n            filter(lambda p: p.requires_grad, params), \n            lr=lr,\n            weight_decay=0)\n    \n    return model_optimizer","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:56:17.50694Z","iopub.execute_input":"2023-06-19T10:56:17.507318Z","iopub.status.idle":"2023-06-19T10:56:17.514714Z","shell.execute_reply.started":"2023-06-19T10:56:17.507289Z","shell.execute_reply":"2023-06-19T10:56:17.5116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_scheduler(cfg, optimizer, total_steps):\n    scheduler = get_cosine_schedule_with_warmup(\n        optimizer,\n        num_warmup_steps= cfg.warmup * (total_steps // cfg.batch_size),\n        num_training_steps= cfg.num_epochs * (total_steps // cfg.batch_size)\n    )\n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2023-06-19T10:56:20.772641Z","iopub.execute_input":"2023-06-19T10:56:20.773009Z","iopub.status.idle":"2023-06-19T10:56:20.778545Z","shell.execute_reply.started":"2023-06-19T10:56:20.772976Z","shell.execute_reply":"2023-06-19T10:56:20.777605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_EPOCHS = Config.num_epochs\nmodel = UNet(Config).to(Config.device)\n\ntotal_steps = len(train_ds)\noptimizer = get_optimizer(lr=Config.lr, params=model.parameters())\nscheduler = get_scheduler(Config, optimizer, total_steps)\n\nfrom timeit import default_timer as timer\nstart_time = timer()\nmodel_results = train(model, train_dl, valid_dl, optimizer, NUM_EPOCHS, Config.device)\nend_time = timer()\n\nprint(f'Total Training Time: {end_time-start_time:.3f} seconds')","metadata":{"execution":{"iopub.status.busy":"2023-06-19T11:01:00.753945Z","iopub.execute_input":"2023-06-19T11:01:00.754588Z","iopub.status.idle":"2023-06-19T11:01:57.734547Z","shell.execute_reply.started":"2023-06-19T11:01:00.754556Z","shell.execute_reply":"2023-06-19T11:01:57.732611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"test_path = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/test\"\nfilenames = os.listdir(test_path)\ntest_df = pd.DataFrame(filenames, columns=['record_id'])\ntest_df['path'] = test_path + \"/\" + test_df['record_id'].astype(str)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:26:12.726403Z","iopub.execute_input":"2023-06-19T09:26:12.726968Z","iopub.status.idle":"2023-06-19T09:26:12.757593Z","shell.execute_reply.started":"2023-06-19T09:26:12.726923Z","shell.execute_reply":"2023-06-19T09:26:12.756499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:26:20.748023Z","iopub.execute_input":"2023-06-19T09:26:20.748638Z","iopub.status.idle":"2023-06-19T09:26:20.7657Z","shell.execute_reply.started":"2023-06-19T09:26:20.748602Z","shell.execute_reply":"2023-06-19T09:26:20.764586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, train=True):\n        \n        self.df = df\n        self.trn = train\n    \n    def read_record(self, directory):\n        record_data = {}\n        for x in [\n            \"band_11\", \n            \"band_14\", \n            \"band_15\"\n        ]:\n\n            record_data[x] = np.load(os.path.join(directory, x + \".npy\"))\n\n        return record_data\n\n    def normalize_range(self, data, bounds):\n        \"\"\"Maps 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        _T11_BOUNDS = (243, 303)\n        _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n        _TDIFF_BOUNDS = (-4, 2)\n        \n        N_TIMES_BEFORE = 4\n\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        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)    \n        \n        img = self.get_false_color(data)\n        \n        img = torch.tensor(img)\n        img = img.permute(2, 0, 1)\n            \n        return img.float()\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:39:00.628926Z","iopub.execute_input":"2023-06-19T09:39:00.63014Z","iopub.status.idle":"2023-06-19T09:39:00.642674Z","shell.execute_reply.started":"2023-06-19T09:39:00.630088Z","shell.execute_reply":"2023-06-19T09:39:00.64153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = ContrailsDataset(test_df,train=False) \ntest_dl = DataLoader(test_ds, batch_size=Config.batch_size, num_workers = 2)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:39:02.786813Z","iopub.execute_input":"2023-06-19T09:39:02.787373Z","iopub.status.idle":"2023-06-19T09:39:02.793439Z","shell.execute_reply.started":"2023-06-19T09:39:02.787337Z","shell.execute_reply":"2023-06-19T09:39:02.792373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, cfg):\n        super(UNet, self).__init__()\n        \n        self.cfg = cfg\n        self.training = True\n        \n        self.model = smp.Unet(\n            encoder_name=cfg.encoder, \n            encoder_weights=cfg.weights, \n            decoder_use_batchnorm=True,\n            classes=len(cfg.classes), \n            activation=cfg.activation,\n        )\n        \n        self.loss_fn = smp.losses.DiceLoss(mode='binary')\n    \n    def forward(self, imgs):\n        \n        x = imgs\n        logits = self.model(x)\n        \n        return {\"logits\": logits.sigmoid()}","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:49:55.117537Z","iopub.execute_input":"2023-06-19T09:49:55.120282Z","iopub.status.idle":"2023-06-19T09:49:55.129987Z","shell.execute_reply.started":"2023-06-19T09:49:55.120243Z","shell.execute_reply":"2023-06-19T09:49:55.129045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_ckpt = '/kaggle/working/epoch-9.pth'","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:50:21.763265Z","iopub.execute_input":"2023-06-19T09:50:21.764547Z","iopub.status.idle":"2023-06-19T09:50:21.769346Z","shell.execute_reply.started":"2023-06-19T09:50:21.764496Z","shell.execute_reply":"2023-06-19T09:50:21.768307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet(Config).to(Config.device)\nmodel.load_state_dict(torch.load(model_ckpt, map_location=torch.device('cuda')))","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:50:25.662246Z","iopub.execute_input":"2023-06-19T09:50:25.662629Z","iopub.status.idle":"2023-06-19T09:50:26.155938Z","shell.execute_reply.started":"2023-06-19T09:50:25.662596Z","shell.execute_reply":"2023-06-19T09:50:26.154812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\ntorch.set_grad_enabled(False)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:50:37.149247Z","iopub.execute_input":"2023-06-19T09:50:37.150517Z","iopub.status.idle":"2023-06-19T09:50:37.161735Z","shell.execute_reply.started":"2023-06-19T09:50:37.150481Z","shell.execute_reply":"2023-06-19T09:50:37.159816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_data = defaultdict(list)\npbar = tqdm(enumerate(test_dl), total=len(test_dl), desc='Test')\nfor step, X in pbar: \n    X = X.to(Config.device)\n\n    output = model(X)\n    for key, val in output.items():\n        val_data[key] += [output[key]]\n\nfor key, val in output.items():\n    value = val_data[key]\n    if len(value[0].shape) == 0:\n        val_data[key] = torch.stack(value)\n    else:\n        val_data[key] = torch.cat(value, dim=0).cpu().detach().numpy()","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:50:42.1494Z","iopub.execute_input":"2023-06-19T09:50:42.149787Z","iopub.status.idle":"2023-06-19T09:50:43.285574Z","shell.execute_reply.started":"2023-06-19T09:50:42.149755Z","shell.execute_reply":"2023-06-19T09:50:43.284397Z"},"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), 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\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-06-19T09:52:02.653899Z","iopub.execute_input":"2023-06-19T09:52:02.654764Z","iopub.status.idle":"2023-06-19T09:52:02.66229Z","shell.execute_reply.started":"2023-06-19T09:52:02.654727Z","shell.execute_reply":"2023-06-19T09:52:02.661368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv('/kaggle/input/google-research-identify-contrails-reduce-global-warming/sample_submission.csv', index_col='record_id')\n\nfor i, pred in enumerate(val_data['logits']):\n    rec = test_df['record_id'][i]\n    mask = (pred[0]>Config.thr).astype(np.float32)\n    submission.loc[int(rec), 'encoded_pixels'] = list_to_string(rle_encode(mask))\n\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:54:03.091845Z","iopub.execute_input":"2023-06-19T09:54:03.093135Z","iopub.status.idle":"2023-06-19T09:54:03.123454Z","shell.execute_reply.started":"2023-06-19T09:54:03.093088Z","shell.execute_reply":"2023-06-19T09:54:03.121435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-06-19T09:54:38.935705Z","iopub.execute_input":"2023-06-19T09:54:38.936146Z","iopub.status.idle":"2023-06-19T09:54:38.946402Z","shell.execute_reply.started":"2023-06-19T09:54:38.936111Z","shell.execute_reply":"2023-06-19T09:54:38.945333Z"},"trusted":true},"execution_count":null,"outputs":[]}]}