{"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":"# Simple Unet Baseline (Infer)\n\nThis is the inference part of the two part Unet Baseline for this competition.\n#### Training Notebook: [Simple Unet Baseline (Train)][1]. \n\n### Please upvote if you find this useful.\n\n[1]: https://www.kaggle.com/code/shashwatraman/simple-unet-pytorch-baseline-train","metadata":{}},{"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\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\n\nimport torch\nfrom torch import nn\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nimport torch.nn.functional as F\n\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nfrom transformers import get_cosine_schedule_with_warmup\nfrom tqdm.auto import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-08T18:54:23.412984Z","iopub.execute_input":"2023-07-08T18:54:23.413329Z","iopub.status.idle":"2023-07-08T18:54:23.420103Z","shell.execute_reply.started":"2023-07-08T18:54:23.413302Z","shell.execute_reply":"2023-07-08T18:54:23.419038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"../input/pretrained-models-pytorch\")\nsys.path.append(\"../input/efficientnet-pytorch\")\nsys.path.append(\"/kaggle/input/smp-github/segmentation_models.pytorch-master\")\nimport segmentation_models_pytorch as smp\n\nprint(f\"Segmentation Models version: {smp.__version__}\")","metadata":{"execution":{"iopub.status.busy":"2023-07-08T18:54:23.573246Z","iopub.execute_input":"2023-07-08T18:54:23.573602Z","iopub.status.idle":"2023-07-08T18:54:23.582515Z","shell.execute_reply.started":"2023-07-08T18:54:23.573573Z","shell.execute_reply":"2023-07-08T18:54:23.581303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    batch_size = 32\n    seed = 42\n    thr = 0.01\n    \n    encoder = 'efficientnet-b3'\n    pretrained = False\n    weights = None\n    classes = ['contrail']\n    activation = None\n    in_chans = 3\n    \n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    image_size = 256\n    \n    model_ckpt = '/kaggle/input/unet-model/epoch-29.pth'\n    \nclass Paths:\n    data = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\n    data_root = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/test/'","metadata":{"execution":{"iopub.status.busy":"2023-07-08T18:54:23.772695Z","iopub.execute_input":"2023-07-08T18:54:23.773711Z","iopub.status.idle":"2023-07-08T18:54:23.780318Z","shell.execute_reply.started":"2023-07-08T18:54:23.773678Z","shell.execute_reply":"2023-07-08T18:54:23.779315Z"},"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-07-08T18:54:23.986809Z","iopub.execute_input":"2023-07-08T18:54:23.98716Z","iopub.status.idle":"2023-07-08T18:54:23.993156Z","shell.execute_reply.started":"2023-07-08T18:54:23.987133Z","shell.execute_reply":"2023-07-08T18:54:23.991831Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set_seed(Config.seed)","metadata":{"execution":{"iopub.status.busy":"2023-07-08T18:54:24.170211Z","iopub.execute_input":"2023-07-08T18:54:24.170569Z","iopub.status.idle":"2023-07-08T18:54:24.176621Z","shell.execute_reply.started":"2023-07-08T18:54:24.17054Z","shell.execute_reply":"2023-07-08T18:54:24.175623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Preparation","metadata":{}},{"cell_type":"code","source":"filenames = os.listdir(Paths.data_root)\ntest_df = pd.DataFrame(filenames, columns=['record_id'])\n\ntest_df['path'] = Paths.data_root + test_df['record_id'].astype(str)","metadata":{"execution":{"iopub.status.busy":"2023-07-08T18:54:24.494899Z","iopub.execute_input":"2023-07-08T18:54:24.495246Z","iopub.status.idle":"2023-07-08T18:54:24.503313Z","shell.execute_reply.started":"2023-07-08T18:54:24.495218Z","shell.execute_reply":"2023-07-08T18:54:24.502126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-07-08T18:54:24.710226Z","iopub.execute_input":"2023-07-08T18:54:24.710906Z","iopub.status.idle":"2023-07-08T18:54:24.722307Z","shell.execute_reply.started":"2023-07-08T18:54:24.710859Z","shell.execute_reply":"2023-07-08T18:54:24.721152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform_size = A.Compose([\n    A.Resize(Config.image_size, Config.image_size, interpolation=cv2.INTER_LANCZOS4, always_apply=True)\n])","metadata":{"execution":{"iopub.status.busy":"2023-07-08T18:54:24.816764Z","iopub.execute_input":"2023-07-08T18:54:24.819055Z","iopub.status.idle":"2023-07-08T18:54:24.824539Z","shell.execute_reply.started":"2023-07-08T18:54:24.819013Z","shell.execute_reply":"2023-07-08T18:54:24.823401Z"},"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        if Config.image_size != 256:\n            img = transform_size(image=img)[\"image\"]\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-07-08T18:54:24.977392Z","iopub.execute_input":"2023-07-08T18:54:24.979973Z","iopub.status.idle":"2023-07-08T18:54:24.991738Z","shell.execute_reply.started":"2023-07-08T18:54:24.979938Z","shell.execute_reply":"2023-07-08T18:54:24.990591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = ContrailsDataset(\n        test_df,\n        train = False\n    )\n \ntest_dl = DataLoader(test_ds, batch_size=Config.batch_size, num_workers = 2)","metadata":{"execution":{"iopub.status.busy":"2023-07-08T18:54:25.221096Z","iopub.execute_input":"2023-07-08T18:54:25.22204Z","iopub.status.idle":"2023-07-08T18:54:25.228036Z","shell.execute_reply.started":"2023-07-08T18:54:25.222Z","shell.execute_reply":"2023-07-08T18:54:25.226911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, cfg):\n        super(UNet, self).__init__()\n        \n        self.cfg = cfg\n        self.training = True\n        \n        self.model = smp.Unet(\n            encoder_name=cfg.encoder, \n            encoder_weights=cfg.weights, \n            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        if Config.image_size != 256:\n            logits = F.interpolate(logits, size=(256, 256), mode='nearest-exact')\n        \n        return {\"logits\": logits.sigmoid()}","metadata":{"execution":{"iopub.status.busy":"2023-07-08T18:54:25.494583Z","iopub.execute_input":"2023-07-08T18:54:25.495295Z","iopub.status.idle":"2023-07-08T18:54:25.508669Z","shell.execute_reply.started":"2023-07-08T18:54:25.495257Z","shell.execute_reply":"2023-07-08T18:54:25.507547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet(Config).to(Config.device)\nmodel.load_state_dict(torch.load(Config.model_ckpt, map_location=torch.device('cuda')))","metadata":{"execution":{"iopub.status.busy":"2023-07-08T18:54:25.673218Z","iopub.execute_input":"2023-07-08T18:54:25.674105Z","iopub.status.idle":"2023-07-08T18:54:31.767482Z","shell.execute_reply.started":"2023-07-08T18:54:25.674068Z","shell.execute_reply":"2023-07-08T18:54:31.766532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\ntorch.set_grad_enabled(False)\n\nval_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-07-08T18:54:31.769498Z","iopub.execute_input":"2023-07-08T18:54:31.769859Z","iopub.status.idle":"2023-07-08T18:54:37.452062Z","shell.execute_reply.started":"2023-07-08T18:54:31.769825Z","shell.execute_reply":"2023-07-08T18:54:37.45106Z"},"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-07-08T18:54:37.453565Z","iopub.execute_input":"2023-07-08T18:54:37.455042Z","iopub.status.idle":"2023-07-08T18:54:37.463515Z","shell.execute_reply.started":"2023-07-08T18:54:37.455002Z","shell.execute_reply":"2023-07-08T18:54:37.462167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv(Paths.data + '/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-07-08T18:54:37.466253Z","iopub.execute_input":"2023-07-08T18:54:37.466631Z","iopub.status.idle":"2023-07-08T18:54:37.499938Z","shell.execute_reply.started":"2023-07-08T18:54:37.466598Z","shell.execute_reply":"2023-07-08T18:54:37.498929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-07-08T18:54:37.501263Z","iopub.execute_input":"2023-07-08T18:54:37.502177Z","iopub.status.idle":"2023-07-08T18:54:37.511083Z","shell.execute_reply.started":"2023-07-08T18:54:37.502144Z","shell.execute_reply":"2023-07-08T18:54:37.509838Z"},"trusted":true},"execution_count":null,"outputs":[]}]}