{"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":"# Requirements","metadata":{}},{"cell_type":"code","source":"!pip install -q --no-index /kaggle/input/contrail-packages/munch-4.0.0-py2.py3-none-any.whl\n!cp -r /kaggle/input/contrail-packages/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4 /tmp/\n!pip install -q --no-index /tmp/pretrainedmodels-0.7.4\n!cp -r /kaggle/input/contrail-packages/efficientnet_pytorch-0.7.1/efficientnet_pytorch-0.7.1 /tmp/\n!pip install -q --no-index /tmp/efficientnet_pytorch-0.7.1\n!pip install -q --no-index /kaggle/input/contrail-packages/segmentation_models_pytorch-0.3.3-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2023-08-09T16:39:24.815877Z","iopub.execute_input":"2023-08-09T16:39:24.816283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare csv files","metadata":{}},{"cell_type":"code","source":"DEBUG = False\n\nimport glob\nimport pandas as pd\n\nif DEBUG:\n    test_paths = glob.glob('/kaggle/input/google-research-identify-contrails-reduce-global-warming/validation/*')\nelse:\n    test_paths = glob.glob('/kaggle/input/google-research-identify-contrails-reduce-global-warming/test/*')\n\ndf = pd.DataFrame(dict(path=test_paths))\nn = len(df) // 2\ndf.iloc[:n].to_csv('/tmp/test_part1.csv', index=False)\ndf.iloc[n:].to_csv('/tmp/test_part2.csv', index=False)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"%%writefile predict.py\n\nimport argparse\nimport itertools\nimport os\nimport os.path as osp\n\n\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\nimport segmentation_models_pytorch as smp\nfrom torch.utils.data import Dataset, DataLoader\n\n\n_T11_BOUNDS = (243, 303)\n_CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n_TDIFF_BOUNDS = (-4, 2)\n\n\ndef normalize_range(data, bounds):\n    \"\"\"Maps data to the range [0, 1].\"\"\"\n    return (data - bounds[0]) / (bounds[1] - bounds[0])\n\n\ndef get_transform_pair(hflip, k):\n    def pre(img):\n        if hflip:\n            img = img.flip([-1])\n        img = img.rot90(k, [-2, -1])\n        return img\n\n    def post(img):\n        img = img.rot90(-k, [-2, -1])\n        if hflip:\n            img = img.flip([-1])\n        return img\n\n    return pre, post\n\n\nTTA_COMBINATIONS = [\n    get_transform_pair(*args) for args in itertools.product([False, True], [0, 1, 2, 3])\n]\n\n\nclass SegModel(nn.Module):\n    def __init__(\n        self,\n        encoder_name,\n        encoder_weights,\n        extra_args=dict(),\n    ):\n        super().__init__()\n        self.register_buffer(\n            \"pixel_mean\",\n            torch.FloatTensor([123.675, 116.28, 103.53]).reshape(1, 3, 1, 1),\n        )\n        self.register_buffer(\n            \"pixel_std\",\n            torch.FloatTensor([58.395, 57.12, 57.375]).reshape(1, 3, 1, 1),\n        )\n        self.model = smp.Unet(\n            encoder_name=encoder_name,\n            encoder_weights=encoder_weights,\n            in_channels=3,\n            classes=1,\n            decoder_use_batchnorm=True,\n            **extra_args,\n        )\n        self.loss_fn = smp.losses.DiceLoss(\n            smp.losses.BINARY_MODE, from_logits=True, smooth=1.0\n        )\n        self.norm_eval = False\n\n    def forward(self, img):\n        img = img.float() * 255\n        img = (img - self.pixel_mean) / self.pixel_std\n        mask = self.model(img)\n        return mask\n\n    @torch.no_grad()\n    def predict(self, img):\n        # out = self.forward(img)\n        with torch.cuda.amp.autocast(dtype=torch.float16):\n            out = self.forward(img)\n        out = F.interpolate(\n            out,\n            (256, 256),\n            mode=\"bilinear\",\n            align_corners=False,\n        )\n        pred = out.sigmoid()\n        return pred\n\n    @torch.no_grad()\n    def predict_tta8(self, img):\n        pred = torch.zeros(\n            img.size(0), 1, 256, 256, dtype=torch.float32, device=img.device\n        )\n        for pre, post in TTA_COMBINATIONS:\n            out = self.predict(pre(img))\n            out = post(out)\n            pred += out\n        pred /= len(TTA_COMBINATIONS)\n        return pred\n\n\nclass ContrailDataset(Dataset):\n    def __init__(\n        self,\n        csv_file,\n        img_root,\n    ):\n        self.df = pd.read_csv(csv_file)\n        self.img_root = img_root\n\n        self.img_affine_matrix = np.array(\n            [[2.0, 0.0, 1.5], [0.0, 2.0, 1.5]], dtype=np.float64\n        )\n        self.img_size = (512, 512)\n\n        print(f\"Loaded {len(self)} samples.\")\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        data_root = self.df.iloc[index]['path']\n        record_id = osp.basename(data_root)\n\n        with open(osp.join(data_root, \"band_11.npy\"), \"rb\") as f:\n            band11 = np.load(f)[..., 4]\n        with open(osp.join(data_root, \"band_14.npy\"), \"rb\") as f:\n            band14 = np.load(f)[..., 4]\n        with open(osp.join(data_root, \"band_15.npy\"), \"rb\") as f:\n            band15 = np.load(f)[..., 4]\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        img = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n        img = cv2.warpAffine(\n            img,\n            self.img_affine_matrix,\n            self.img_size,\n            flags=cv2.INTER_LINEAR,\n            borderMode=cv2.BORDER_CONSTANT,\n            borderValue=0,\n        )\n\n        img = torch.from_numpy(img).permute(2, 0, 1)\n\n        return dict(record_id=record_id, img=img)\n\n    def collate_fn(self, samples):\n        return dict(\n            record_id=[_[\"record_id\"] for _ in samples],\n            img=torch.stack([_[\"img\"] for _ in samples]),\n        )\n\n\ndef main():\n    parser = argparse.ArgumentParser()\n    parser.add_argument(\"csv_file\")\n    parser.add_argument(\"encoder_name\")\n    parser.add_argument(\"checkpoint\")\n    parser.add_argument(\"out_dir\")\n    parser.add_argument(\"--tta\", default=False, action=\"store_true\")\n    args = parser.parse_args()\n    print(args)\n\n    os.makedirs(args.out_dir, exist_ok=True)\n\n    ds = ContrailDataset(\n        args.csv_file,\n        \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/test\",\n    )\n    model = SegModel(encoder_name=args.encoder_name, encoder_weights=None)\n    ckpt = torch.load(args.checkpoint, \"cpu\")\n    model.load_state_dict(ckpt, strict=True)\n\n    model.eval()\n    model.cuda()\n\n    dl = DataLoader(ds, batch_size=16, collate_fn=ds.collate_fn)\n    from tqdm import tqdm\n\n    for data in tqdm(dl):\n        img = data[\"img\"].cuda()\n        if args.tta:\n            pred = model.predict_tta8(img)\n        else:\n            pred = model.predict(img)\n        pred = pred.data.cpu().numpy()\n        for record_id, img_pred in zip(data[\"record_id\"], pred):\n            np.save(osp.join(args.out_dir, f\"{record_id}.npy\"), img_pred)\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!(CUDA_VISIBLE_DEVICES=0 python predict.py /tmp/test_part1.csv timm-efficientnet-l2 /kaggle/input/contrail-checkpoints/l2.pth /tmp/l2/ --tta \\\n  & CUDA_VISIBLE_DEVICES=1 python predict.py /tmp/test_part2.csv timm-efficientnet-l2 /kaggle/input/contrail-checkpoints/l2.pth /tmp/l2/ --tta \\\n  & wait)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!(CUDA_VISIBLE_DEVICES=0 python predict.py /tmp/test_part1.csv timm-resnest269e /kaggle/input/contrail-checkpoints/s269.pth /tmp/s269/ --tta \\\n  & CUDA_VISIBLE_DEVICES=1 python predict.py /tmp/test_part2.csv timm-resnest269e /kaggle/input/contrail-checkpoints/s269.pth /tmp/s269/ --tta \\\n  & wait)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!(CUDA_VISIBLE_DEVICES=0 python predict.py /tmp/test_part1.csv tu-maxvit_base_tf_512 /kaggle/input/contrail-checkpoints/maxvitb.pth /tmp/maxvitb/ --tta \\\n  & CUDA_VISIBLE_DEVICES=1 python predict.py /tmp/test_part2.csv tu-maxvit_base_tf_512 /kaggle/input/contrail-checkpoints/maxvitb.pth /tmp/maxvitb/ --tta \\\n  & wait)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!(CUDA_VISIBLE_DEVICES=0 python predict.py /tmp/test_part1.csv tu-tf_efficientnetv2_l /kaggle/input/contrail-checkpoints/v2l.pth /tmp/v2l/ --tta \\\n  & CUDA_VISIBLE_DEVICES=1 python predict.py /tmp/test_part2.csv tu-tf_efficientnetv2_l /kaggle/input/contrail-checkpoints/v2l.pth /tmp/v2l/ --tta \\\n  & wait)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!(CUDA_VISIBLE_DEVICES=0 python predict.py /tmp/test_part1.csv tu-tf_efficientnetv2_xl /kaggle/input/contrail-checkpoints/v2xl.pth /tmp/v2xl/ --tta \\\n  & CUDA_VISIBLE_DEVICES=1 python predict.py /tmp/test_part2.csv tu-tf_efficientnetv2_xl /kaggle/input/contrail-checkpoints/v2xl.pth /tmp/v2xl/ --tta \\\n  & wait)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make submission","metadata":{}},{"cell_type":"code","source":"%%writefile make_submission.py\n\n\nimport os.path as osp\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\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\n\n\ndf = pd.concat([\n    pd.read_csv('/tmp/test_part1.csv'),\n    pd.read_csv('/tmp/test_part2.csv'),\n], axis=0)\nmodel_ids = ['l2', 's269', 'maxvitb', 'v2l', 'v2xl']\nweights = [10, 5, 3, 1, 1]\nthr = 0.5\nrecord_ids = [osp.basename(path) for path in df['path']]\n\nrles = []\nfor rid in tqdm(record_ids):\n    prs = [np.load(f'/tmp/{model_id}/{rid}.npy') for model_id in model_ids]\n    pr = np.average(prs, axis=0, weights=weights)\n    pr_binary = (pr > thr)\n    rle = list_to_string(rle_encode(pr_binary))\n    rles.append(rle)\n    \n\nsub = pd.DataFrame(dict(record_id=record_ids, encoded_pixels=rles))\nsub.to_csv('submission.csv', index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python make_submission.py","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('submission.csv')\nprint(df.shape)\ndf.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}