{"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":"code","source":"## import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-16T20:00:21.627459Z","iopub.execute_input":"2023-08-16T20:00:21.62786Z","iopub.status.idle":"2023-08-16T20:00:21.633385Z","shell.execute_reply.started":"2023-08-16T20:00:21.627831Z","shell.execute_reply":"2023-08-16T20:00:21.632242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/googlewarminglibs/einops-0.6.1-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2023-08-16T20:00:21.635694Z","iopub.execute_input":"2023-08-16T20:00:21.63613Z","iopub.status.idle":"2023-08-16T20:00:52.622546Z","shell.execute_reply.started":"2023-08-16T20:00:21.636098Z","shell.execute_reply":"2023-08-16T20:00:52.621079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/googlewarminglibs/')\n","metadata":{"execution":{"iopub.status.busy":"2023-08-16T20:00:52.625152Z","iopub.execute_input":"2023-08-16T20:00:52.625932Z","iopub.status.idle":"2023-08-16T20:00:52.631611Z","shell.execute_reply.started":"2023-08-16T20:00:52.625892Z","shell.execute_reply":"2023-08-16T20:00:52.630523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport re\nfrom dataclasses import dataclass\nfrom typing import Dict\nfrom typing import List\n\nimport albumentations\nimport cv2\nimport numpy as np\nimport pydicom\nimport tifffile\nimport torch\nimport torch.hub\nfrom albumentations import ReplayCompose\nfrom skimage import measure\nfrom torch.functional import Tensor\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-08-16T20:00:52.634813Z","iopub.execute_input":"2023-08-16T20:00:52.635246Z","iopub.status.idle":"2023-08-16T20:00:52.64507Z","shell.execute_reply.started":"2023-08-16T20:00:52.635219Z","shell.execute_reply":"2023-08-16T20:00:52.643791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"resize_augs =  albumentations.ReplayCompose([\n            albumentations.Resize(768, 768),\n            ])\n\n\n_T11_BOUNDS = (243, 303)\n_CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n_TDIFF_BOUNDS = (-4, 2)\n\ndef normalize_range(data, bounds):\n    \"\"\"Maps data to the range [0, 1].\"\"\"\n    return (data - bounds[0]) / (bounds[1] - bounds[0])\n\nclass DatasetSeq(Dataset):\n    def __init__(\n            self,\n            dataset_dir: str,\n            transforms: albumentations.Compose,\n    ):\n        self.transforms = transforms\n        self.dataset_dir = dataset_dir\n        self.ids = sorted(os.listdir(self.dataset_dir))\n\n    def __getitem__(self, i):\n        return self.getitem(i)\n \n\n    def getitem(self, i):\n        file_id = self.ids[i]\n        band11 = np.load(os.path.join(self.dataset_dir, file_id, 'band_11.npy'))\n        band14 = np.load(os.path.join(self.dataset_dir, file_id, 'band_14.npy'))\n        band15 = np.load(os.path.join(self.dataset_dir, file_id, 'band_15.npy'))\n\n        r = normalize_range(band15 - band14, _TDIFF_BOUNDS)\n        g = normalize_range(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\n        b = normalize_range(band14, _T11_BOUNDS)\n\n        imgs = np.array([r, g, b])\n\n        imgs = np.transpose(imgs, (3, 1, 2, 0))\n        \n        replay = None\n        image_crops = []\n        for i in range(len(imgs)):\n            image = imgs[i]\n            if replay is None:\n                sample = self.transforms(image=image)\n                replay = sample[\"replay\"]\n            else:\n                sample = ReplayCompose.replay(replay, image=image)\n            image_ = sample[\"image\"]\n            image_crops.append(image_)\n        images = np.array(image_crops)\n        sample = {}\n        sample['file_id'] = file_id\n        sample['image'] = torch.from_numpy(np.moveaxis(images, -1, 1)).float()\n\n        return sample\n\n\n    def __len__(self):\n        return len(self.ids)\n","metadata":{"execution":{"iopub.status.busy":"2023-08-16T20:00:52.647037Z","iopub.execute_input":"2023-08-16T20:00:52.647396Z","iopub.status.idle":"2023-08-16T20:00:52.663Z","shell.execute_reply.started":"2023-08-16T20:00:52.647364Z","shell.execute_reply":"2023-08-16T20:00:52.661968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_checkpoint(model, checkpoint_path, strict=False, verbose=True):\n    if verbose:\n        print(\"=> loading checkpoint '{}'\".format(checkpoint_path))\n    checkpoint = torch.load(checkpoint_path, map_location='cpu')\n    if 'state_dict' in checkpoint:\n        state_dict = checkpoint['state_dict']\n        state_dict = {re.sub(\"^module.\", \"\", k): w for k, w in state_dict.items()}\n        orig_state_dict = model.state_dict()\n        mismatched_keys = []\n        for k, v in state_dict.items():\n            ori_size = orig_state_dict[k].size() if k in orig_state_dict else None\n            if v.size() != ori_size:\n                if verbose:\n                    print(\"SKIPPING!!! Shape of {} changed from {} to {}\".format(k, v.size(), ori_size))\n                mismatched_keys.append(k)\n        for k in mismatched_keys:\n            del state_dict[k]\n        model.load_state_dict(state_dict, strict=strict)\n        del state_dict\n        del orig_state_dict\n        print(\"=> loaded checkpoint '{}' (epoch {})\"\n              .format(checkpoint_path, checkpoint['epoch']))\n    else:\n        model.load_state_dict(checkpoint)\n    del checkpoint\n\n\ndef load_model(conf: Dict, checkpoint: str):\n    model = conf[\"network\"](**conf[\"encoder_params\"])\n    model = model.cuda()\n    load_checkpoint(model, checkpoint)\n    return model.eval()","metadata":{"execution":{"iopub.status.busy":"2023-08-16T20:00:52.664642Z","iopub.execute_input":"2023-08-16T20:00:52.665049Z","iopub.status.idle":"2023-08-16T20:00:52.678086Z","shell.execute_reply.started":"2023-08-16T20:00:52.66499Z","shell.execute_reply":"2023-08-16T20:00:52.677215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\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\n\n\nall_preds = []\nall_targets = []\ndef process_segmentations(models: List[torch.nn.Module], weights: List[float], test_dataset_dir: str, threshold: float) -> pd.DataFrame:\n    test_dataset = DatasetSeq(dataset_dir=test_dataset_dir, transforms=resize_augs)\n    sampler = None\n    test_loader = DataLoader(\n        test_dataset, batch_size=1, sampler=sampler, shuffle=False, num_workers=1, pin_memory=False\n    )\n    data = []\n    for sample in tqdm(test_loader):\n        image = sample[\"image\"]\n        file_id = sample[\"file_id\"][0]\n        img = image.cuda().float()\n        all_masks = []\n        with torch.no_grad():\n            with torch.cuda.amp.autocast():\n                for i, model in enumerate(models):\n                    w = weights[i]\n                    mask = model(img)[\"mask\"].sigmoid().cpu().float()[0][0]\n                    mask += torch.flip(model(torch.flip(img, dims=(4,)))[\"mask\"].sigmoid().cpu().float(), dims=(3,))[0][0]\n                    mask += torch.rot90(model(torch.rot90(img, k=1, dims=(3, 4)))[\"mask\"].sigmoid().cpu().float(), k=-1, dims=(2, 3))[0][0]\n                    mask += torch.rot90(model(torch.rot90(img, k=-1, dims=(3, 4)))[\"mask\"].sigmoid().cpu().float(), k=1, dims=(2, 3))[0][0]\n                    mask /= 4\n                    all_masks.append(w * mask)\n        preds = sum(all_masks) # / len(all_masks)\n        preds = cv2.resize(preds.numpy().astype(np.float32), dsize=(256, 256))\n        preds = (preds > threshold).astype(np.uint8)\n        data.append([file_id, list_to_string(rle_encode(preds))])\n    return pd.DataFrame(data, columns=['record_id', 'encoded_pixels'])","metadata":{"execution":{"iopub.status.busy":"2023-08-16T20:00:52.679892Z","iopub.execute_input":"2023-08-16T20:00:52.680248Z","iopub.status.idle":"2023-08-16T20:00:52.699186Z","shell.execute_reply.started":"2023-08-16T20:00:52.680218Z","shell.execute_reply":"2023-08-16T20:00:52.696931Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import zoo\nconfig_seg  = {\n  \"network\": zoo.unet.TimmUnetPure,\n  \"encoder_params\": {\n    \"encoder\": \"tf_efficientnetv2_l_in21k\",\n    \"drop_path_rate\": 0.0,\n    \"in_chans\": 3,\n    \"num_classes\": 1,\n    \"pretrained\": False\n  }\n}\nv2l = load_model(config_seg, \"/kaggle/input/warmingfinal/swa_5_best_final_768_TimmUnetPure_tf_efficientnetv2_l_in21k_0.pth\")\n\nconfig_seg  = {\n  \"network\": zoo.unet.TimmUnetPure,\n  \"encoder_params\": {\n    \"encoder\": \"tf_efficientnet_l2_ns\",\n    \"drop_path_rate\": 0.0,\n    \"in_chans\": 3,\n    \"num_classes\": 1,\n    \"pretrained\": False\n  }\n}\nl2ns = load_model(config_seg, \"/kaggle/input/warmingfinal/swa_5_best_final_768_TimmUnetPure_tf_efficientnet_l2_ns_0.pth\")\nconfig_seg  = {\n  \"network\": zoo.unet.TimmUnetPure,\n  \"encoder_params\": {\n    \"encoder\": \"tf_efficientnetv2_xl_in21k\",\n    \"drop_path_rate\": 0.0,\n    \"in_chans\": 3,\n    \"num_classes\": 1,\n    \"pretrained\": False\n  }\n}\nv2xl = load_model(config_seg, \"/kaggle/input/warmingfinal/swa_5_best_final_768_TimmUnetPure_tf_efficientnetv2_xl_in21k_0.pth\")\n\nconfig_seg  = {\n  \"network\": zoo.unet.TimmUnetPure,\n  \"encoder_params\": {\n    \"encoder\": \"maxvit_base_tf_512.in21k_ft_in1k\",\n    \"drop_path_rate\": 0.0,\n    \"in_chans\": 3,\n    \"num_classes\": 1,\n    \"pretrained\": False,\n    \"img_size\": 768\n  }\n}\nmaxvit = load_model(config_seg, \"/kaggle/input/warmingfinal/swa_5_best_final_768_TimmUnetPure_maxvit_base_tf_512.in21k_ft_in1k_0.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-08-16T20:00:52.700475Z","iopub.execute_input":"2023-08-16T20:00:52.700845Z","iopub.status.idle":"2023-08-16T20:01:12.149415Z","shell.execute_reply.started":"2023-08-16T20:00:52.700814Z","shell.execute_reply":"2023-08-16T20:01:12.148399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset_dir = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/test\"\nthreshold = 0.355\npreds_df = process_segmentations([l2ns, v2l, maxvit, v2xl], [0.45, 0.15, 0.35, 0.05], test_dataset_dir, threshold)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-08-16T20:01:12.153363Z","iopub.execute_input":"2023-08-16T20:01:12.154188Z","iopub.status.idle":"2023-08-16T20:01:18.496274Z","shell.execute_reply.started":"2023-08-16T20:01:12.154159Z","shell.execute_reply":"2023-08-16T20:01:18.495104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv('/kaggle/input/google-research-identify-contrails-reduce-global-warming/sample_submission.csv', dtype={\"record_id\": str, \"encoded_pixels\": str})\ndel sub_df['encoded_pixels']\n\nsub_df = sub_df.merge(preds_df, on=\"record_id\")\nsub_df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-08-16T20:01:18.49813Z","iopub.execute_input":"2023-08-16T20:01:18.498913Z","iopub.status.idle":"2023-08-16T20:01:18.513997Z","shell.execute_reply.started":"2023-08-16T20:01:18.498863Z","shell.execute_reply":"2023-08-16T20:01:18.512993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df","metadata":{"execution":{"iopub.status.busy":"2023-08-16T20:01:18.515482Z","iopub.execute_input":"2023-08-16T20:01:18.51648Z","iopub.status.idle":"2023-08-16T20:01:18.52635Z","shell.execute_reply.started":"2023-08-16T20:01:18.516449Z","shell.execute_reply":"2023-08-16T20:01:18.525205Z"},"trusted":true},"execution_count":null,"outputs":[]}]}