{"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":"### Training part","metadata":{}},{"cell_type":"code","source":"import torch\nimport sys\nimport gc\nimport os\nimport glob\nimport yaml\nimport pandas as pd\nimport numpy as np\nimport torch.nn as n\nsys.path.append(\"../input/pretrained-models-pytorch\")\nsys.path.append(\"../input/efficientnet-pytorch\")\nsys.path.append(\"../input/smp-github/segmentation_models.pytorch-master\")\nsys.path.append(\"../input/timm-pretrained-resnest/resnest/\")","metadata":{"execution":{"iopub.status.busy":"2023-08-09T00:57:05.006164Z","iopub.execute_input":"2023-08-09T00:57:05.006856Z","iopub.status.idle":"2023-08-09T00:57:09.847904Z","shell.execute_reply.started":"2023-08-09T00:57:05.006825Z","shell.execute_reply":"2023-08-09T00:57:09.846589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pytorch_lightning as pl\nfrom torch.utils.data import DataLoader\nfrom pytorch_lightning.loggers import CSVLogger\nimport pytorch_lightning as pl\nimport torchvision.transforms as T\nfrom torchmetrics.functional import dice\nfrom torch.utils.data import Dataset, DataLoader\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, TQDMProgressBar\nimport segmentation_models_pytorch as smp\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\ntorch.set_float32_matmul_precision(\"medium\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-09T00:57:09.850362Z","iopub.execute_input":"2023-08-09T00:57:09.851059Z","iopub.status.idle":"2023-08-09T00:57:25.184267Z","shell.execute_reply.started":"2023-08-09T00:57:09.851016Z","shell.execute_reply":"2023-08-09T00:57:25.182187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile config.yaml\n\ndata_path: \"/kaggle/input/contrails-images-ash-color\"\noutput_dir: \"models\"\n\nfolds:\n    n_splits: 4\n    random_state: 42\ntrain_folds: [0, 1, 2, 3]\n    \nseed: 42\n\ntrain_bs: 48\nvalid_bs: 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: 20\n    min_epochs: 20\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: 512","metadata":{"execution":{"iopub.status.busy":"2023-08-09T00:57:25.186694Z","iopub.execute_input":"2023-08-09T00:57:25.187367Z","iopub.status.idle":"2023-08-09T00:57:25.194178Z","shell.execute_reply.started":"2023-08-09T00:57:25.187332Z","shell.execute_reply":"2023-08-09T00:57:25.193233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seg_models = {\n    \"Unet\": smp.Unet,\n}","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-09T00:57:25.197104Z","iopub.execute_input":"2023-08-09T00:57:25.197472Z","iopub.status.idle":"2023-08-09T00:57:25.207573Z","shell.execute_reply.started":"2023-08-09T00:57:25.197435Z","shell.execute_reply":"2023-08-09T00:57:25.206573Z"},"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)\n\npl.seed_everything(config[\"seed\"])\ngc.enable()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-09T00:57:25.20889Z","iopub.execute_input":"2023-08-09T00:57:25.210731Z","iopub.status.idle":"2023-08-09T00:57:25.228793Z","shell.execute_reply.started":"2023-08-09T00:57:25.210697Z","shell.execute_reply":"2023-08-09T00:57:25.227857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32\nnum_workers = 1\nTHR = 0.47\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndata = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\ndata_root = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/test/'\nsubmission = pd.read_csv(os.path.join(data, 'sample_submission.csv'), index_col='record_id')","metadata":{"execution":{"iopub.status.busy":"2023-08-09T00:57:25.231418Z","iopub.execute_input":"2023-08-09T00:57:25.232058Z","iopub.status.idle":"2023-08-09T00:57:25.279927Z","shell.execute_reply.started":"2023-08-09T00:57:25.232018Z","shell.execute_reply":"2023-08-09T00:57:25.277949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"filenames = os.listdir(data_root)\ntest_df = pd.DataFrame(filenames, columns=['record_id'])\ntest_df['path'] = data_root + test_df['record_id'].astype(str)","metadata":{"execution":{"iopub.status.busy":"2023-08-09T00:57:25.285098Z","iopub.execute_input":"2023-08-09T00:57:25.28606Z","iopub.status.idle":"2023-08-09T00:57:25.297579Z","shell.execute_reply.started":"2023-08-09T00:57:25.286018Z","shell.execute_reply":"2023-08-09T00:57:25.295789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, image_size=256, train=True):\n        \n        self.df = df\n        self.trn = train\n        self.df_idx: pd.DataFrame = pd.DataFrame({'idx': os.listdir(f'/kaggle/input/google-research-identify-contrails-reduce-global-warming/test')})\n        self.normalize_image = T.Normalize((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        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(np.reshape(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        image_id = int(self.df_idx.iloc[index]['idx'])\n            \n        return img.float(), torch.tensor(image_id)\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-08-09T00:57:25.302015Z","iopub.execute_input":"2023-08-09T00:57:25.302833Z","iopub.status.idle":"2023-08-09T00:57:25.331893Z","shell.execute_reply.started":"2023-08-09T00:57:25.302796Z","shell.execute_reply":"2023-08-09T00:57:25.330834Z"},"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-08-09T00:57:25.337989Z","iopub.execute_input":"2023-08-09T00:57:25.341317Z","iopub.status.idle":"2023-08-09T00:57:25.356732Z","shell.execute_reply.started":"2023-08-09T00:57:25.341279Z","shell.execute_reply":"2023-08-09T00:57:25.355549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LightningModule(pl.LightningModule):\n\n    def __init__(self, config):\n        super().__init__()\n        self.model = smp.Unet(encoder_name=config[\"encoder_name\"],\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-09T00:57:25.366527Z","iopub.execute_input":"2023-08-09T00:57:25.367145Z","iopub.status.idle":"2023-08-09T00:57:25.374055Z","shell.execute_reply.started":"2023-08-09T00:57:25.367112Z","shell.execute_reply":"2023-08-09T00:57:25.373168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_PATH = \"/kaggle/input/test512pesudo8/\"","metadata":{"execution":{"iopub.status.busy":"2023-08-09T00:57:25.375342Z","iopub.execute_input":"2023-08-09T00:57:25.376753Z","iopub.status.idle":"2023-08-09T00:57:25.385007Z","shell.execute_reply.started":"2023-08-09T00:57:25.376721Z","shell.execute_reply":"2023-08-09T00:57:25.384139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = ContrailsDataset(\n        test_df,\n        config[\"model\"][\"image_size\"],\n        train = False\n    )\n \ntest_dl = DataLoader(test_ds, batch_size=batch_size, num_workers = num_workers)","metadata":{"execution":{"iopub.status.busy":"2023-08-09T00:57:25.388167Z","iopub.execute_input":"2023-08-09T00:57:25.389604Z","iopub.status.idle":"2023-08-09T00:57:25.397169Z","shell.execute_reply.started":"2023-08-09T00:57:25.389566Z","shell.execute_reply":"2023-08-09T00:57:25.396207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.enable()\n\nall_preds = {}\n\nfor i, model_path in enumerate(glob.glob(MODEL_PATH + '*.ckpt')):\n    print(model_path)\n    model = LightningModule(config[\"model\"]).load_from_checkpoint(model_path, config=config[\"model\"])\n    model.to(device)\n    model.eval()\n\n    model_preds = {}\n    \n    for _, data in enumerate(test_dl):\n        images, image_id = data\n    \n        images = images.to(device)\n        \n        with torch.no_grad():\n            predicted_mask = model(images[:, :, :, :])\n        if config[\"model\"][\"image_size\"] != 256:\n            predicted_mask = torch.nn.functional.interpolate(predicted_mask, size=256, mode='bilinear')\n        predicted_mask = torch.sigmoid(predicted_mask).cpu().detach().numpy()\n                \n        for img_num in range(0, images.shape[0]):\n            current_mask = predicted_mask[img_num, :, :, :]\n            current_image_id = image_id[img_num].item()\n            model_preds[current_image_id] = current_mask\n    all_preds[f\"f{i}\"] = model_preds\n    \n    del model    \n    torch.cuda.empty_cache()\n    gc.collect() ","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-09T00:57:25.398829Z","iopub.execute_input":"2023-08-09T00:57:25.399186Z","iopub.status.idle":"2023-08-09T00:57:25.420567Z","shell.execute_reply.started":"2023-08-09T00:57:25.399155Z","shell.execute_reply":"2023-08-09T00:57:25.419496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2","metadata":{"execution":{"iopub.status.busy":"2023-08-09T00:57:25.422015Z","iopub.execute_input":"2023-08-09T00:57:25.422704Z","iopub.status.idle":"2023-08-09T00:57:25.612357Z","shell.execute_reply.started":"2023-08-09T00:57:25.422669Z","shell.execute_reply":"2023-08-09T00:57:25.611409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for index in submission.index.tolist():\n    for i in range(len(glob.glob(MODEL_PATH + '*.ckpt'))):\n        if i == 0:\n            predicted_mask = all_preds[f\"f{i}\"][index]\n        else:\n            predicted_mask += all_preds[f\"f{i}\"][index]\n        #kernel = np.ones(shape=(3, 3), dtype=np.uint8)\n        #predicted_mask = cv2.dilate(predicted_mask.astype(np.uint8), kernel, 4)  \n    predicted_mask = predicted_mask / len(glob.glob(MODEL_PATH + '*.ckpt'))\n    predicted_mask_with_threshold = np.zeros((256, 256))\n    predicted_mask_with_threshold[predicted_mask[0, :, :] < THR] = 0\n    predicted_mask_with_threshold[predicted_mask[0, :, :] > THR] = 1\n    submission.loc[int(index), 'encoded_pixels'] = list_to_string(rle_encode(predicted_mask_with_threshold))","metadata":{"execution":{"iopub.status.busy":"2023-08-09T00:57:25.614693Z","iopub.execute_input":"2023-08-09T00:57:25.615373Z","iopub.status.idle":"2023-08-09T00:57:26.513993Z","shell.execute_reply.started":"2023-08-09T00:57:25.615338Z","shell.execute_reply":"2023-08-09T00:57:26.512492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-08-09T00:57:26.515098Z","iopub.status.idle":"2023-08-09T00:57:26.516933Z","shell.execute_reply.started":"2023-08-09T00:57:26.516674Z","shell.execute_reply":"2023-08-09T00:57:26.516699Z"},"trusted":true},"execution_count":null,"outputs":[]}]}