{"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":"# %env TORCH_CUDNN_V8_API_DISABLED=1\n%env CUDA_MODULE_LOADING=LAZY","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:56:59.853908Z","iopub.execute_input":"2023-08-08T06:56:59.854476Z","iopub.status.idle":"2023-08-08T06:56:59.868557Z","shell.execute_reply.started":"2023-08-08T06:56:59.854441Z","shell.execute_reply":"2023-08-08T06:56:59.867446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%env CUDA_MODULE_LOADING","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:56:59.870404Z","iopub.execute_input":"2023-08-08T06:56:59.870974Z","iopub.status.idle":"2023-08-08T06:56:59.884812Z","shell.execute_reply.started":"2023-08-08T06:56:59.87094Z","shell.execute_reply":"2023-08-08T06:56:59.88367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvcc -V","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:56:59.886098Z","iopub.execute_input":"2023-08-08T06:56:59.886951Z","iopub.status.idle":"2023-08-08T06:57:00.873415Z","shell.execute_reply.started":"2023-08-08T06:56:59.886878Z","shell.execute_reply":"2023-08-08T06:57:00.872177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python -m torch.utils.collect_env","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:57:00.879563Z","iopub.execute_input":"2023-08-08T06:57:00.880253Z","iopub.status.idle":"2023-08-08T06:57:55.754009Z","shell.execute_reply.started":"2023-08-08T06:57:00.880214Z","shell.execute_reply":"2023-08-08T06:57:55.752782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\nimport numpy as np \nimport pandas as pd \nfrom tqdm.auto import trange\nfrom glob import glob\n\n\nimport argparse\nimport datetime\nimport math\nimport os\nimport warnings\nfrom functools import partial\nfrom typing import Callable, List, Tuple\n\nimport albumentations as albu\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\nimport timm\nimport torch\nfrom albumentations.pytorch import ToTensorV2\nfrom pytorch_lightning import LightningDataModule, callbacks\nfrom pytorch_lightning.loggers import WandbLogger\nfrom pytorch_lightning.utilities import rank_zero_info\nfrom scipy.optimize import minimize\nfrom timm.utils import ModelEmaV2\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.optim import AdamW\nfrom torch.utils.data import DataLoader, Dataset\nfrom transformers import get_cosine_schedule_with_warmup\nimport shutil\nimport gc\n\nfrom tqdm.auto import tqdm\nimport matplotlib.pyplot as plt\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-08T06:57:55.756148Z","iopub.execute_input":"2023-08-08T06:57:55.756528Z","iopub.status.idle":"2023-08-08T06:58:10.710597Z","shell.execute_reply.started":"2023-08-08T06:57:55.756493Z","shell.execute_reply":"2023-08-08T06:58:10.709674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\n# sys.path.append('../input/pytorch-image-models/pytorch-image-models')\nsys.path.append(\"../input/pretrained-models-pytorch\")\nsys.path.append(\"../input/efficientnet-pytorch\")\nsys.path.append(\"../input/segmentation-models-pytorch/segmentation_models_pytorch\")\nsys.path.append(\"../input/pytorch-pfn-extras/pytorch-pfn-extras\")\n\nimport os\nimport gc\nimport random\nimport shutil\nimport typing as tp\nfrom pathlib import Path\n\nimport yaml\nimport numpy as np\nimport pandas as pd\n\n\nfrom tqdm.notebook import tqdm, trange\nfrom joblib import Parallel, delayed\n\nimport cv2\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom torch.utils import data\n\nimport pytorch_pfn_extras as ppe\nfrom pytorch_pfn_extras.training import extensions as exts\nfrom pytorch_pfn_extras.training import triggers as trgrs\nfrom pytorch_pfn_extras.config import Config\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:10.712151Z","iopub.execute_input":"2023-08-08T06:58:10.712494Z","iopub.status.idle":"2023-08-08T06:58:14.17571Z","shell.execute_reply.started":"2023-08-08T06:58:10.712458Z","shell.execute_reply":"2023-08-08T06:58:14.174716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preprocessing (img to npy)","metadata":{}},{"cell_type":"code","source":"ash_dataset_dir = \"/kaggle/working/tmp/ash_dataset\"\n\ndata_types = {\"record_id\": str}\ntest_df = pd.read_csv(\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/sample_submission.csv\",dtype=data_types)\nrecords = test_df.record_id.to_numpy()\n\ndef get_image(record_id):\n    with open(os.path.join(\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/test\", record_id, \"band_11.npy\"), \"rb\") as f:\n        band11 = np.load(f)\n    with open(os.path.join(\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/test\", record_id, \"band_14.npy\"), \"rb\") as f:\n        band14 = np.load(f)\n    with open(os.path.join(\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/test\", record_id, \"band_15.npy\"), \"rb\") as f:\n        band15 = np.load(f)\n\n    _T11_BOUNDS = (243, 303)\n    _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n    _TDIFF_BOUNDS = (-4, 2)\n\n    def normalize_range(data, bounds):\n        \"\"\"Maps data to the range [0, 1].\"\"\"\n        return (data - bounds[0]) / (bounds[1] - bounds[0])\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    false_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n    return false_color.astype(np.float16)\n\nash_data_path = []\nfor record_id in records:\n    img = get_image(record_id)\n    out_dir = os.path.join(ash_dataset_dir, \"test\", record_id)\n    os.makedirs(out_dir, exist_ok=True)\n    ash_data_path.append(os.path.abspath(out_dir))\n    np.save(os.path.join(out_dir, \"img\"), img)\ntest_df[\"ash_data_path\"] = ash_data_path\n# test_df.to_csv(\"/kaggle/working/tmp/test.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:14.180074Z","iopub.execute_input":"2023-08-08T06:58:14.180373Z","iopub.status.idle":"2023-08-08T06:58:14.470636Z","shell.execute_reply.started":"2023-08-08T06:58:14.180347Z","shell.execute_reply":"2023-08-08T06:58:14.469667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data processing","metadata":{}},{"cell_type":"code","source":"img_mean = (\n    np.concatenate([np.asarray([55.936405, 140.56143, 143.563]) for _ in range(8)])\n    / 255.0\n)\nimg_std = (\n    np.concatenate([np.asarray([30.419374, 37.97873, 48.56115]) for _ in range(8)])\n    / 255.0\n)\n\ndef get_transforms(train: bool = False) -> Callable:\n    return albu.Compose(\n        [\n            albu.Normalize(img_mean, img_std),\n            ToTensorV2(transpose_mask=True),\n        ]\n    )\n\n\nclass ContrailsDataset(Dataset):\n    def __init__(\n        self,\n        df: pd.DataFrame,\n        mode: str = \"train\",  # \"train\" | \"valid\" | \"test\"\n    ):\n        self.df = df\n        self.mode = mode\n        self.train = mode == \"train\"\n        self.transforms = get_transforms(self.train)\n\n    def __len__(self) -> int:\n        return len(self.df)\n\n    def np_load(self, file) -> np.ndarray:\n        if type(file) == str:\n            file = open(file, \"rb\")\n        header = file.read(128)\n        if not header:\n            return None\n        descr = str(header[19:25], \"utf-8\").replace(\"'\", \"\").replace(\" \", \"\")\n        shape = tuple(\n            int(num)\n            for num in str(header[60:120], \"utf-8\")\n            .replace(\", }\", \"\")\n            .replace(\"(\", \"\")\n            .replace(\")\", \"\")\n            .split(\",\")\n        )\n        datasize = np.lib.format.descr_to_dtype(descr).itemsize\n        for dimension in shape:\n            datasize *= dimension\n        return np.ndarray(shape, dtype=descr, buffer=file.read(datasize))\n\n    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]:\n        image = (\n            self.np_load(\n                os.path.join(self.df.ash_data_path.iloc[idx], \"img.npy\")\n            ).astype(np.float32)\n            * 255\n        )  # (h, w, c, l)\n        h, w, c, d = image.shape\n        image = image.reshape(h, w, c * d)\n        aug = self.transforms(image=image)\n        image = aug[\"image\"]  # (c * d, h, w)\n        return image\n\n\nclass ContrailsDataModule(LightningDataModule):\n    def __init__(\n        self,\n        train_df: pd.DataFrame,\n        valid_df: pd.DataFrame,\n        num_workers: int = 4,\n        batch_size: int = 16,\n    ):\n        super().__init__()\n\n        self._num_workers = num_workers\n        self._batch_size = batch_size\n        self.train_df = train_df\n        self.valid_df = valid_df\n        self.save_hyperparameters(\n            \"num_workers\",\n            \"batch_size\",\n        )\n\n    def create_dataset(self, mode: str = \"train\") -> ContrailsDataset:\n        if mode == \"train\":\n            return ContrailsDataset(\n                df=self.train_df,\n                mode=mode,\n            )\n        else:\n            return ContrailsDataset(\n                df=self.valid_df,\n                mode=mode,\n            )\n\n    def __dataloader(self, mode: str = \"train\") -> DataLoader:\n        \"\"\"Train/validation loaders.\"\"\"\n        dataset = self.create_dataset(mode)\n        return DataLoader(\n            dataset=dataset,\n            batch_size=self._batch_size,\n            num_workers=self._num_workers,\n            shuffle=(mode == \"train\"),\n            drop_last=(mode == \"train\"),\n            pin_memory=True,\n        )\n\n    def train_dataloader(self) -> DataLoader:\n        return self.__dataloader(mode=\"train\")\n\n    def val_dataloader(self) -> DataLoader:\n        return self.__dataloader(mode=\"valid\")\n\n    def test_dataloader(self) -> DataLoader:\n        return self.__dataloader(mode=\"test\")\n\n    @staticmethod\n    def add_model_specific_args(\n        parent_parser: argparse.ArgumentParser,\n    ) -> argparse.ArgumentParser:\n        parser = parent_parser.add_argument_group(\"ContrailsDataModule\")\n        parser.add_argument(\n            \"--num_workers\",\n            default=4,\n            type=int,\n            metavar=\"W\",\n            help=\"number of CPU workers\",\n            dest=\"num_workers\",\n        )\n        parser.add_argument(\n            \"--batch_size\",\n            default=16,\n            type=int,\n            metavar=\"BS\",\n            help=\"number of sample in a batch\",\n            dest=\"batch_size\",\n        )\n        return parent_parser","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-08-08T06:58:14.471962Z","iopub.execute_input":"2023-08-08T06:58:14.472317Z","iopub.status.idle":"2023-08-08T06:58:14.498666Z","shell.execute_reply.started":"2023-08-08T06:58:14.472275Z","shell.execute_reply":"2023-08-08T06:58:14.497785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT = Path.cwd().parent\nINPUT = ROOT / \"input\"\nOUTPUT = ROOT / \"output\"\nDATA = INPUT / \"google-research-identify-contrails-reduce-global-warming\"\nTRAIN = DATA / \"train\"\nVARID = DATA / \"valid\" \nTEST = DATA / \"test\"\n\nTMP = ROOT / \"tmp\"\nTMP.mkdir(exist_ok=True)\nTMP_TEST = TMP / \"test\"\nTMP_TEST.mkdir(exist_ok=True)\n\n# # Dataset\n\nclass GRICRGWTestDataset(data.Dataset):\n    \"\"\"\"\"\"\n    \n    def __init__(self, image_paths, transform):\n        \"\"\"\"\"\"\n        self.image_paths = image_paths\n        self.transform = transform\n        \n    def __len__(self):\n        \"\"\"\"\"\"\n        return len(self.image_paths)\n    \n    def __getitem__(self, index):\n        \"\"\"\"\"\"\n        image_path = self.image_paths[index]\n        image = np.load(image_path)\n        image = self._apply_transform(image)\n        \n        return {\"data\": image}\n        \n    def _apply_transform(self, image: np.ndarray):\n        \"\"\"\"\"\" \n        image = self.transform(image=image)[\"image\"]\n        return image\n\n\n# # Other functions\n\ndef load_yaml_file(path: str):\n    \"\"\"Load YAML setting file.\"\"\"\n    with open(path) as f:\n        settings = yaml.safe_load(f)\n    return settings\n\ndef get_array_module(x: tp.Union[np.ndarray, torch.Tensor]):\n    \"\"\"\"\"\"\n    if isinstance(x, torch.Tensor):\n        return torch\n    else:\n        return np\n\ndef sigmoid(x: tp.Union[np.ndarray, torch.Tensor]):\n    \"\"\"\"\"\"\n    xp = get_array_module(x)\n    return 1 / (1 + xp.exp(-x))\n\n\ndef to_device(\n    tensors: tp.Union[tp.Tuple[torch.Tensor], tp.Dict[str, torch.Tensor]],\n    device: torch.device, *args, **kwargs\n):\n    \"\"\"\"\"\"\n    if isinstance(tensors, tuple):\n        return (t.to(device, *args, **kwargs) for t in tensors)\n    elif isinstance(tensors, dict):\n        return {\n            k: t.to(device, *args, **kwargs) for k, t in tensors.items()}\n    else:\n        return tensors.to(device, *args, **kwargs)\n    \n    \ndef save_tdiff_plus_all_band_image(from_dir: Path):\n    to_dir = TMP_TEST / from_dir.name\n    to_dir.mkdir(exist_ok=True)\n    \n    img_list = []\n    for b in range(8, 17):\n        img = np.load(from_dir / f\"band_{b:0>2}.npy\")[..., 4]  # shape: (256, 256)\n        img_list.append(img)\n        \n    img_list = [\n        img_list[15 - 8] - img_list[14 - 8],  # shape: (256, 256)\n        img_list[14 - 8] - img_list[11 - 8],  # shape: (256, 256)\n    ] + img_list\n    \n    tdiff_plus_all_band = np.stack(img_list, axis=2)  # shape: (256, 256, 11)\n    np.save(to_dir / \"tdiff_plus_all_band.npy\", tdiff_plus_all_band)\n    \nfor from_dir in sorted(TEST.iterdir()):\n    save_tdiff_plus_all_band_image(from_dir)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:14.500212Z","iopub.execute_input":"2023-08-08T06:58:14.500626Z","iopub.status.idle":"2023-08-08T06:58:14.893553Z","shell.execute_reply.started":"2023-08-08T06:58:14.500589Z","shell.execute_reply":"2023-08-08T06:58:14.8926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Modeling (Decorder)","metadata":{}},{"cell_type":"code","source":"import math\n\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn\nfrom torch.autograd import Variable\n\n\nclass SCSEModule(nn.Module):\n    def __init__(self, ch, re=16):\n        super().__init__()\n        self.cSE = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(ch, ch // re, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(ch // re, ch, 1),\n            nn.Sigmoid(),\n        )\n        self.sSE = nn.Sequential(nn.Conv2d(ch, ch, 1), nn.Sigmoid())\n\n    def forward(self, x):\n        return x * self.cSE(x) + x * self.sSE(x)\n\n\nclass Conv2dReLU(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        out_channels,\n        kernel_size,\n        padding=0,\n        stride=1,\n        use_batchnorm=True,\n        **batchnorm_params\n    ):\n        super().__init__()\n\n        layers = [\n            nn.Conv2d(\n                in_channels,\n                out_channels,\n                kernel_size,\n                stride=stride,\n                padding=padding,\n                bias=not (use_batchnorm),\n            ),\n            nn.ReLU(inplace=True),\n        ]\n\n        if use_batchnorm:\n            layers.insert(1, nn.BatchNorm2d(out_channels, **batchnorm_params))\n\n        self.block = nn.Sequential(*layers)\n\n    def forward(self, x):\n        return self.block(x)\n\n\nclass Flatten(nn.Module):\n    \"\"\"\n    Simple class for flattening layer.\n    \"\"\"\n\n    def forward(self, x):\n        return x.view(x.size()[0], -1)\n\n\nclass BasicConv(nn.Module):\n    def __init__(\n        self,\n        in_planes,\n        out_planes,\n        kernel_size,\n        stride=1,\n        padding=0,\n        dilation=1,\n        groups=1,\n        relu=True,\n        bn=True,\n        bias=False,\n    ):\n        super(BasicConv, self).__init__()\n        self.out_channels = out_planes\n        self.conv = nn.Conv2d(\n            in_planes,\n            out_planes,\n            kernel_size=kernel_size,\n            stride=stride,\n            padding=padding,\n            dilation=dilation,\n            groups=groups,\n            bias=bias,\n        )\n        self.bn = (\n            nn.BatchNorm2d(out_planes, eps=1e-5, momentum=0.01, affine=True)\n            if bn\n            else None\n        )\n        self.relu = nn.ReLU() if relu else None\n\n    def forward(self, x):\n        x = self.conv(x)\n        if self.bn is not None:\n            x = self.bn(x)\n        if self.relu is not None:\n            x = self.relu(x)\n        return x\n\n\nclass ChannelGate(nn.Module):\n    def __init__(self, gate_channels, reduction_ratio=16, pool_types=[\"avg\", \"max\"]):\n        super(ChannelGate, self).__init__()\n        self.gate_channels = gate_channels\n        self.mlp = nn.Sequential(\n            Flatten(),\n            nn.Linear(gate_channels, gate_channels // reduction_ratio),\n            nn.ReLU(),\n            nn.Linear(gate_channels // reduction_ratio, gate_channels),\n        )\n        self.pool_types = pool_types\n\n    def forward(self, x):\n        channel_att_sum = None\n        for pool_type in self.pool_types:\n            if pool_type == \"avg\":\n                avg_pool = F.avg_pool2d(\n                    x, (x.size(2), x.size(3)), stride=(x.size(2), x.size(3))\n                )\n                channel_att_raw = self.mlp(avg_pool)\n            elif pool_type == \"max\":\n                max_pool = F.max_pool2d(\n                    x, (x.size(2), x.size(3)), stride=(x.size(2), x.size(3))\n                )\n                channel_att_raw = self.mlp(max_pool)\n            elif pool_type == \"lp\":\n                lp_pool = F.lp_pool2d(\n                    x, 2, (x.size(2), x.size(3)), stride=(x.size(2), x.size(3))\n                )\n                channel_att_raw = self.mlp(lp_pool)\n            elif pool_type == \"lse\":\n                # LSE pool only\n                lse_pool = logsumexp_2d(x)\n                channel_att_raw = self.mlp(lse_pool)\n\n            if channel_att_sum is None:\n                channel_att_sum = channel_att_raw\n            else:\n                channel_att_sum = channel_att_sum + channel_att_raw\n\n        scale = F.sigmoid(channel_att_sum).unsqueeze(2).unsqueeze(3).expand_as(x)\n        return x * scale\n\n\ndef logsumexp_2d(tensor):\n    tensor_flatten = tensor.view(tensor.size(0), tensor.size(1), -1)\n    s, _ = torch.max(tensor_flatten, dim=2, keepdim=True)\n    outputs = s + (tensor_flatten - s).exp().sum(dim=2, keepdim=True).log()\n    return outputs\n\n\nclass ChannelPool(nn.Module):\n    def forward(self, x):\n        return torch.cat(\n            (torch.max(x, 1)[0].unsqueeze(1), torch.mean(x, 1).unsqueeze(1)), dim=1\n        )\n\n\nclass SpatialGate(nn.Module):\n    def __init__(self):\n        super(SpatialGate, self).__init__()\n        kernel_size = 7\n        self.compress = ChannelPool()\n        self.spatial = BasicConv(\n            2, 1, kernel_size, stride=1, padding=(kernel_size - 1) // 2, relu=False\n        )\n\n    def forward(self, x):\n        x_compress = self.compress(x)\n        x_out = self.spatial(x_compress)\n        scale = F.sigmoid(x_out)  # broadcasting\n        return x * scale\n\n\nclass CBAM(nn.Module):\n    def __init__(\n        self,\n        gate_channels,\n        reduction_ratio=4,\n        pool_types=[\"avg\", \"max\"],\n        no_spatial=False,\n    ):\n        super(CBAM, self).__init__()\n        self.ChannelGate = ChannelGate(gate_channels, reduction_ratio, pool_types)\n        self.no_spatial = no_spatial\n        if not no_spatial:\n            self.SpatialGate = SpatialGate()\n\n    def forward(self, x):\n        x_out = self.ChannelGate(x)\n        if not self.no_spatial:\n            x_out = self.SpatialGate(x_out)\n        return x_out\n\n\ndef get_sinusoid_encoding_table(n_position, d_hid, padding_idx=None):\n    \"\"\"Sinusoid position encoding table\"\"\"\n\n    def cal_angle(position, hid_idx):\n        return position / np.power(10000, 2 * (hid_idx // 2) / d_hid)\n\n    def get_posi_angle_vec(position):\n        return [cal_angle(position, hid_j) for hid_j in range(d_hid)]\n\n    sinusoid_table = np.array(\n        [get_posi_angle_vec(pos_i) for pos_i in range(n_position)]\n    )\n\n    sinusoid_table[:, 0::2] = np.sin(sinusoid_table[:, 0::2])  # dim 2i\n    sinusoid_table[:, 1::2] = np.cos(sinusoid_table[:, 1::2])  # dim 2i+1\n\n    if padding_idx is not None:\n        # zero vector for padding dimension\n        sinusoid_table[padding_idx] = 0.0\n\n    return sinusoid_table\n\n\ndef get_sinusoid_encoding_table_2d(H, W, d_hid):\n    \"\"\"Sinusoid position encoding table\"\"\"\n    n_position = H * W\n    sinusoid_table = get_sinusoid_encoding_table(n_position, d_hid)\n    sinusoid_table = sinusoid_table.reshape(H, W, d_hid)\n    return sinusoid_table\n\n\nclass CBAMModule(nn.Module):\n    def __init__(\n        self, channels, reduction=4, attention_kernel_size=3, position_encode=False\n    ):\n        super(CBAMModule, self).__init__()\n        self.position_encode = position_encode\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n        self.fc1 = nn.Conv2d(channels, channels // reduction, kernel_size=1, padding=0)\n        self.relu = nn.ReLU(inplace=True)\n        self.fc2 = nn.Conv2d(channels // reduction, channels, kernel_size=1, padding=0)\n        self.sigmoid_channel = nn.Sigmoid()\n        if self.position_encode:\n            k = 3\n        else:\n            k = 2\n        self.conv_after_concat = nn.Conv2d(\n            k,\n            1,\n            kernel_size=attention_kernel_size,\n            stride=1,\n            padding=attention_kernel_size // 2,\n        )\n        self.sigmoid_spatial = nn.Sigmoid()\n        self.position_encoded = None\n\n    def forward(self, x):\n        # Channel attention module\n        module_input = x\n        avg = self.avg_pool(x)\n        mx = self.max_pool(x)\n        avg = self.fc1(avg)\n        mx = self.fc1(mx)\n        avg = self.relu(avg)\n        mx = self.relu(mx)\n        avg = self.fc2(avg)\n        mx = self.fc2(mx)\n        x = avg + mx\n        x = self.sigmoid_channel(x)\n        # Spatial attention module\n        x = module_input * x\n        module_input = x\n        b, c, h, w = x.size()\n        if self.position_encode:\n            if self.position_encoded is None:\n                pos_enc = get_sinusoid_encoding_table(h, w)\n                pos_enc = Variable(torch.FloatTensor(pos_enc), requires_grad=False)\n                if x.is_cuda:\n                    pos_enc = pos_enc.cuda()\n                self.position_encoded = pos_enc\n        avg = torch.mean(x, 1, True)\n        mx, _ = torch.max(x, 1, True)\n        if self.position_encode:\n            pos_enc = self.position_encoded\n            pos_enc = pos_enc.view(1, 1, h, w).repeat(b, 1, 1, 1)\n            x = torch.cat((avg, mx, pos_enc), 1)\n        else:\n            x = torch.cat((avg, mx), 1)\n        x = self.conv_after_concat(x)\n        x = self.sigmoid_spatial(x)\n        x = module_input * x\n        return x\n\n\nclass AdaptiveConcatPool2d(nn.Module):\n    def __init__(self, sz=None):\n        super().__init__()\n        sz = sz or (1, 1)\n        self.ap = nn.AdaptiveAvgPool2d(sz)\n        self.mp = nn.AdaptiveMaxPool2d(sz)\n\n    def forward(self, x):\n        return torch.cat([self.mp(x), self.ap(x)], 1)\n\n\nclass CenterBlock(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        out_channels,\n        use_batchnorm=True,\n    ):\n        super().__init__()\n        self.block = nn.Sequential(\n            Conv2dReLU(\n                in_channels, out_channels, kernel_size=1, use_batchnorm=use_batchnorm\n            ),\n            Conv2dReLU(\n                out_channels, out_channels, kernel_size=1, use_batchnorm=use_batchnorm\n            ),\n        )\n\n    def forward(self, x):\n        return self.block(x)\n\n\nclass FPA(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        \"\"\"\n        Feature Pyramid Attention\n        https://github.com/JaveyWang/Pyramid-Attention-Networks-pytorch/blob/master/networks.py\n        :type channels: int\n        \"\"\"\n        super(FPA, self).__init__()\n        channels_mid = int(in_channels / 4)\n\n        self.channels_cond = in_channels\n\n        # Master branch\n        self.conv_master = nn.Conv2d(\n            self.channels_cond, out_channels, kernel_size=1, bias=False\n        )\n        self.bn_master = nn.BatchNorm2d(out_channels)\n\n        # Global pooling branch\n        self.conv_gpb = nn.Conv2d(\n            self.channels_cond, out_channels, kernel_size=1, bias=False\n        )\n        self.bn_gpb = nn.BatchNorm2d(out_channels)\n\n        # C333 because of the shape of last feature maps is (16, 16).\n        self.conv7x7_1 = nn.Conv2d(\n            self.channels_cond,\n            channels_mid,\n            kernel_size=(7, 7),\n            stride=2,\n            padding=3,\n            bias=False,\n        )\n        self.bn1_1 = nn.BatchNorm2d(channels_mid)\n        self.conv5x5_1 = nn.Conv2d(\n            channels_mid,\n            channels_mid,\n            kernel_size=(5, 5),\n            stride=2,\n            padding=2,\n            bias=False,\n        )\n        self.bn2_1 = nn.BatchNorm2d(channels_mid)\n        self.conv3x3_1 = nn.Conv2d(\n            channels_mid,\n            channels_mid,\n            kernel_size=(3, 3),\n            stride=2,\n            padding=1,\n            bias=False,\n        )\n        self.bn3_1 = nn.BatchNorm2d(channels_mid)\n\n        self.conv7x7_2 = nn.Conv2d(\n            channels_mid,\n            out_channels,\n            kernel_size=(7, 7),\n            stride=1,\n            padding=3,\n            bias=False,\n        )\n        self.bn1_2 = nn.BatchNorm2d(out_channels)\n        self.conv5x5_2 = nn.Conv2d(\n            channels_mid,\n            out_channels,\n            kernel_size=(5, 5),\n            stride=1,\n            padding=2,\n            bias=False,\n        )\n        self.bn2_2 = nn.BatchNorm2d(out_channels)\n        self.conv3x3_2 = nn.Conv2d(\n            channels_mid,\n            out_channels,\n            kernel_size=(3, 3),\n            stride=1,\n            padding=1,\n            bias=False,\n        )\n        self.bn3_2 = nn.BatchNorm2d(out_channels)\n\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, x):\n        # Master branch\n        h, w = x.size(2), x.size(3)\n        x_master = self.conv_master(x)\n        x_master = self.bn_master(x_master)\n\n        # Global pooling branch\n        x_gpb = nn.AvgPool2d(x.shape[2:])(x).view(x.shape[0], self.channels_cond, 1, 1)\n        x_gpb = self.conv_gpb(x_gpb)\n        x_gpb = self.bn_gpb(x_gpb)\n\n        # Branch 1\n        x1_1 = self.conv7x7_1(x)\n        x1_1 = self.bn1_1(x1_1)\n        x1_1 = self.relu(x1_1)\n        x1_2 = self.conv7x7_2(x1_1)\n        x1_2 = self.bn1_2(x1_2)\n        x1_2 = self.relu(x1_2)\n\n        # Branch 2\n        x2_1 = self.conv5x5_1(x1_1)\n        x2_1 = self.bn2_1(x2_1)\n        x2_1 = self.relu(x2_1)\n        x2_2 = self.conv5x5_2(x2_1)\n        x2_2 = self.bn2_2(x2_2)\n        x2_2 = self.relu(x2_2)\n\n        # Branch 3\n        x3_1 = self.conv3x3_1(x2_1)\n        x3_1 = self.bn3_1(x3_1)\n        x3_1 = self.relu(x3_1)\n        x3_2 = self.conv3x3_2(x3_1)\n        x3_2 = self.bn3_2(x3_2)\n        x3_2 = self.relu(x3_2)\n\n        # Merge branch 1 and\n        x3_upsample = nn.Upsample(size=(h // 4, w // 4), mode=\"nearest\")(x3_2)\n        x2_merge = x2_2 + x3_upsample\n        x2_upsample = nn.Upsample(size=(h // 2, w // 2), mode=\"nearest\")(x2_merge)\n        x1_merge = x1_2 + x2_upsample\n        x_master = x_master * nn.Upsample(size=(h, w), mode=\"nearest\")(x1_merge)\n\n        out = x_master + x_gpb\n\n        return out\n\n\nclass _ASPPModule(nn.Module):\n    def __init__(self, inplanes, planes, kernel_size, padding, dilation):\n        super(_ASPPModule, self).__init__()\n        planes = int(planes)\n        self.atrous_conv = nn.Conv2d(\n            inplanes,\n            int(planes),\n            kernel_size=kernel_size,\n            stride=1,\n            padding=padding,\n            dilation=dilation,\n            bias=False,\n        )\n        self.bn = nn.BatchNorm2d(planes)\n        self.relu = nn.ReLU()\n\n        self._init_weight()\n\n    def forward(self, x):\n        x = self.atrous_conv(x)\n        x = self.bn(x)\n\n        return self.relu(x)\n\n    def _init_weight(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                torch.nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\n\nclass ASPP(nn.Module):\n    def __init__(self, inplanes=512, mid_c=256, dilations=[1, 6, 12, 18]):\n        super(ASPP, self).__init__()\n        self.aspp1 = _ASPPModule(inplanes, mid_c, 1, padding=0, dilation=dilations[0])\n        self.aspp2 = _ASPPModule(\n            inplanes, mid_c, 3, padding=dilations[1], dilation=dilations[1]\n        )\n        self.aspp3 = _ASPPModule(\n            inplanes, mid_c, 3, padding=dilations[2], dilation=dilations[2]\n        )\n        self.aspp4 = _ASPPModule(\n            inplanes, mid_c, 3, padding=dilations[3], dilation=dilations[3]\n        )\n        mid_c = int(mid_c)\n        self.global_avg_pool = nn.Sequential(\n            nn.AdaptiveAvgPool2d((1, 1)),\n            nn.Conv2d(inplanes, mid_c, 1, stride=1, bias=False),\n            nn.BatchNorm2d(mid_c),\n            nn.ReLU(),\n        )\n        self.conv1 = nn.Conv2d(mid_c * 5, mid_c, 1, bias=False)\n        self.bn1 = nn.BatchNorm2d(mid_c)\n        self.relu = nn.ReLU()\n        self.dropout = nn.Dropout(0.5)\n        self._init_weight()\n\n    def forward(self, x):\n        x1 = self.aspp1(x)\n        x2 = self.aspp2(x)\n        x3 = self.aspp3(x)\n        x4 = self.aspp4(x)\n        x5 = self.global_avg_pool(x)\n        x5 = F.interpolate(x5, size=x4.size()[2:], mode=\"nearest\")\n        x = torch.cat((x1, x2, x3, x4, x5), dim=1)\n\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n\n        return self.dropout(x)\n\n    def _init_weight(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                torch.nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\n\nclass DecoderBlock(nn.Module):\n    def __init__(\n        self, in_channels, out_channels, use_batchnorm=True, attention_type=None\n    ):\n        super().__init__()\n        if attention_type is None:\n            self.attention1 = nn.Identity()\n            self.attention2 = nn.Identity()\n        elif attention_type == \"scse\":\n            self.attention1 = SCSEModule(in_channels)\n            self.attention2 = SCSEModule(out_channels)\n        elif attention_type == \"cbam\":\n            self.attention1 = CBAMModule(in_channels)\n            self.attention2 = CBAMModule(out_channels)\n\n        self.block = nn.Sequential(\n            Conv2dReLU(\n                in_channels,\n                out_channels,\n                kernel_size=3,\n                padding=1,\n                use_batchnorm=use_batchnorm,\n            ),\n            Conv2dReLU(\n                out_channels,\n                out_channels,\n                kernel_size=3,\n                padding=1,\n                use_batchnorm=use_batchnorm,\n            ),\n        )\n\n    def forward(self, x):\n        x, skip = x[:-1], x[-1]\n        if skip is not None:\n            x = [\n                F.interpolate(xi, size=skip.shape[-2:], mode=\"nearest\")\n                if skip.shape[-1] > xi.shape[-1]\n                else F.max_pool2d(\n                    xi, math.ceil(xi.shape[-1] / skip.shape[-1]), ceil_mode=True\n                )\n                for xi in x\n            ]\n            x.append(skip)\n            x = torch.cat(x, dim=1)\n            x = self.attention1(x)\n        else:\n            x = [F.interpolate(xi, scale_factor=2, mode=\"nearest\") for xi in x]\n            x = torch.cat(x, dim=1)\n        x = self.block(x)\n        x = self.attention2(x)\n        return x\n\n\nclass UNetHead(nn.Module):\n    __name__ = \"UNetHead\"\n\n    def __init__(\n        self,\n        encoder_channels,\n        decoder_channels=[1024, 512, 256, 128, 64],\n        num_class=1,\n        use_batchnorm=True,\n        center=None,\n        attention_type=None,\n        classification=False,\n        deep_supervision=False,\n    ):\n        super().__init__()\n        encoder_channels = encoder_channels[::-1]\n        decoder_channels = decoder_channels[: len(encoder_channels)]\n        if center == \"fpa\":\n            self.center = FPA(encoder_channels[0], decoder_channels[0])\n        elif center == \"aspp\":\n            self.center = ASPP(\n                encoder_channels[0],\n                decoder_channels[0],\n                dilations=[1, (1, 6), (2, 12), (3, 18)],\n            )\n        else:\n            self.center = CenterBlock(\n                encoder_channels[0], decoder_channels[0], use_batchnorm=use_batchnorm\n            )\n        in_channels = self.compute_channels(encoder_channels[1:], decoder_channels[:-1])\n        layers = []\n        for i in range(len(in_channels)):\n            layers.append(\n                DecoderBlock(\n                    in_channels[i],\n                    decoder_channels[i + 1],\n                    use_batchnorm=use_batchnorm,\n                    attention_type=attention_type,\n                )\n            )\n        for i in range(5 - len(decoder_channels)):\n            layers.append(\n                DecoderBlock(\n                    decoder_channels[-1],\n                    decoder_channels[-1],\n                    use_batchnorm=use_batchnorm,\n                    attention_type=attention_type,\n                )\n            )\n        layers.append(\n            DecoderBlock(\n                decoder_channels[-1],\n                decoder_channels[-1],\n                use_batchnorm=use_batchnorm,\n                attention_type=attention_type,\n            )\n        )\n        self.layers = nn.ModuleList(layers)\n        del layers\n        gc.collect()\n        self.final_conv = nn.Conv2d(decoder_channels[-1], num_class, kernel_size=(1, 1))\n\n        self.classification = classification\n        if self.classification:\n            self.linear_feature = nn.Sequential(\n                nn.Conv2d(encoder_channels[0], 512, kernel_size=1),\n                AdaptiveConcatPool2d(1),\n                Flatten(),\n                nn.ReLU(),\n                nn.Linear(1024, 512),\n                nn.BatchNorm1d(512),\n                nn.ReLU(),\n                nn.Dropout(0.2),\n                nn.Linear(512, num_class),\n            )\n        self.deep_supervision = deep_supervision\n        if self.deep_supervision:\n            layers_ds = []\n            for i in range(len(decoder_channels)):\n                layers_ds.append(\n                    nn.Conv2d(decoder_channels[i], num_class, kernel_size=(1, 1))\n                )\n            layers_ds.append(\n                nn.Conv2d(decoder_channels[i + 1], num_class, kernel_size=(1, 1))\n            )\n            self.layers_ds = nn.ModuleList(layers_ds)\n\n    def compute_channels(self, encoder_channels, decoder_channels):\n        channels = [e + d for e, d in zip(encoder_channels, decoder_channels)]\n        return channels\n\n    def forward(self, x):\n        x = x[::-1]\n        encoder_head = x[0]\n        skips = x[1:]\n        x_o = []\n        x_o.append(self.center(encoder_head))\n        for i in range(len(self.layers)):\n            if i < len(skips):\n                skip = skips[i]\n            else:\n                skip = None\n            x_o.append(self.layers[i]([x_o[-1], skip]))\n        x_final = self.final_conv(x_o[-1])\n        output = [x_final]\n        if self.classification:\n            class_logits = self.linear_feature(encoder_head)\n            output.append(class_logits)\n        if self.deep_supervision:\n            x_ds = []\n            for i in range(len(self.layers_ds)):\n                x_ds.append(self.layers_ds[i](x_o[i]))\n            output.append(x_ds)\n        return output[0] if len(output) == 1 else output\n\n\nclass SeparableConv2d(nn.Module):\n    def __init__(\n        self,\n        inplanes,\n        planes,\n        kernel_size=3,\n        stride=1,\n        padding=1,\n        dilation=1,\n        bias=False,\n        BatchNorm=nn.BatchNorm2d,\n    ):\n        super(SeparableConv2d, self).__init__()\n\n        self.conv1 = nn.Conv2d(\n            inplanes,\n            inplanes,\n            kernel_size,\n            stride,\n            padding,\n            dilation,\n            groups=inplanes,\n            bias=bias,\n        )\n        self.bn = BatchNorm(inplanes)\n        self.pointwise = nn.Conv2d(inplanes, planes, 1, 1, 0, 1, 1, bias=bias)\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.bn(x)\n        x = self.pointwise(x)\n        return x\n\n\nclass JPU(nn.Module):\n    def __init__(self, in_channels, width=512):\n        super(JPU, self).__init__()\n        self.conv5 = nn.Sequential(\n            nn.Conv2d(in_channels[0], width, 3, padding=1, bias=False),\n            nn.BatchNorm2d(width),\n            nn.ReLU(inplace=True),\n        )\n        self.conv4 = nn.Sequential(\n            nn.Conv2d(in_channels[1], width, 3, padding=1, bias=False),\n            nn.BatchNorm2d(width),\n            nn.ReLU(inplace=True),\n        )\n        self.conv3 = nn.Sequential(\n            nn.Conv2d(in_channels[2], width, 3, padding=1, bias=False),\n            nn.BatchNorm2d(width),\n            nn.ReLU(inplace=True),\n        )\n\n        self.dilation1 = nn.Sequential(\n            SeparableConv2d(\n                3 * width, width, kernel_size=3, padding=1, dilation=1, bias=False\n            ),\n            nn.BatchNorm2d(width),\n            nn.ReLU(inplace=True),\n        )\n        self.dilation2 = nn.Sequential(\n            SeparableConv2d(\n                3 * width, width, kernel_size=3, padding=2, dilation=2, bias=False\n            ),\n            nn.BatchNorm2d(width),\n            nn.ReLU(inplace=True),\n        )\n        self.dilation3 = nn.Sequential(\n            SeparableConv2d(\n                3 * width, width, kernel_size=3, padding=4, dilation=4, bias=False\n            ),\n            nn.BatchNorm2d(width),\n            nn.ReLU(inplace=True),\n        )\n        self.dilation4 = nn.Sequential(\n            SeparableConv2d(\n                3 * width, width, kernel_size=3, padding=8, dilation=8, bias=False\n            ),\n            nn.BatchNorm2d(width),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, *inputs):\n        feats = [self.conv5(inputs[0]), self.conv4(inputs[1]), self.conv3(inputs[2])]\n        _, _, h, w = feats[-1].size()\n        feats[-2] = F.interpolate(feats[-2], size=(h, w), mode=\"nearest\")\n        feats[-3] = F.interpolate(feats[-3], size=(h, w), mode=\"nearest\")\n        feat = torch.cat(feats, dim=1)\n        feat = torch.cat(\n            [\n                self.dilation1(feat),\n                self.dilation2(feat),\n                self.dilation3(feat),\n                self.dilation4(feat),\n            ],\n            dim=1,\n        )\n        return feat\n\n\nclass FastFCNImproveHead(nn.Module):\n    __name__ = \"FastFCNImproveHead\"\n\n    def __init__(\n        self,\n        encoder_channels,\n        decoder_channels=(256, 128, 64),\n        num_class=1,\n        use_batchnorm=True,\n        attention_type=None,\n        classification=False,\n        deep_supervision=False,\n    ):\n        super().__init__()\n        encoder_channels = encoder_channels[::-1]\n        self.jpu = JPU(\n            [encoder_channels[0], encoder_channels[1], encoder_channels[2]],\n            decoder_channels[0],\n        )\n        self.aspp = ASPP(\n            decoder_channels[0] * 4,\n            decoder_channels[0],\n            dilations=[1, (1, 4), (2, 8), (3, 12)],\n        )\n        self.decoder1 = DecoderBlock(\n            encoder_channels[3] + decoder_channels[0],\n            decoder_channels[1],\n            use_batchnorm,\n            attention_type,\n        )\n        self.decoder2 = DecoderBlock(\n            encoder_channels[4] + decoder_channels[1],\n            decoder_channels[2],\n            use_batchnorm,\n            attention_type,\n        )\n        self.decoder3 = DecoderBlock(decoder_channels[2], decoder_channels[2])\n        self.final_conv = nn.Conv2d(decoder_channels[2], num_class, kernel_size=(1, 1))\n\n        self.classification = classification\n        if self.classification:\n            self.linear_feature = nn.Sequential(\n                nn.Conv2d(encoder_channels[0], 512, kernel_size=1),\n                AdaptiveConcatPool2d(1),\n                Flatten(),\n                nn.ReLU(),\n                nn.Linear(1024, 512),\n                nn.BatchNorm1d(512),\n                nn.ReLU(),\n                nn.Dropout(0.2),\n                nn.Linear(512, num_class),\n            )\n        self.deep_supervision = deep_supervision\n        if self.deep_supervision:\n            self.layer0_ds = nn.Conv2d(\n                decoder_channels[0], num_class, kernel_size=(1, 1)\n            )\n            self.layer1_ds = nn.Conv2d(\n                decoder_channels[1], num_class, kernel_size=(1, 1)\n            )\n            self.layer2_ds = nn.Conv2d(\n                decoder_channels[2], num_class, kernel_size=(1, 1)\n            )\n\n    def forward(self, x):\n        x = x[::-1]\n        skips = x\n        x_0 = self.jpu(skips[0], skips[1], skips[2])\n        x_0 = self.aspp(x_0)\n        x_1 = self.decoder1([x_0, skips[3]])\n        x_2 = self.decoder2([x_1, skips[4]])\n        x_3 = self.decoder3([x_2, None])\n        x_final = self.final_conv(x_3)\n        output = [x_final]\n        if self.classification:\n            class_refine = self.linear_feature(x[0])\n            output.append(class_refine)\n        if self.deep_supervision:\n            x_ds = []\n            x_ds.append(self.layer0_ds(x_0))\n            x_ds.append(self.layer1_ds(x_1))\n            x_ds.append(self.layer2_ds(x_2))\n            output.append(x_ds)\n        return output[0] if len(output) == 1 else output\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-08-08T06:58:14.897895Z","iopub.execute_input":"2023-08-08T06:58:14.898199Z","iopub.status.idle":"2023-08-08T06:58:15.024502Z","shell.execute_reply.started":"2023-08-08T06:58:14.898173Z","shell.execute_reply":"2023-08-08T06:58:15.02354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Modeling (2.5D)","metadata":{}},{"cell_type":"code","source":"def downsample_conv(\n    in_channels: int,\n    out_channels: int,\n    stride: int = 2,\n):\n    return nn.Sequential(\n        *[\n            nn.Conv3d(\n                in_channels,\n                out_channels,\n                1,\n                stride=(1, stride, stride),\n                padding=0,\n                bias=False,\n            ),\n            nn.BatchNorm3d(out_channels),\n        ]\n    )\n\n\nclass ResidualConv3D(nn.Module):\n    def __init__(\n        self,\n        in_channels: int,\n        mid_channels: int,\n        out_channels: int,\n        stride: int = 2,\n    ):\n        super().__init__()\n\n        self.conv1 = nn.Conv3d(in_channels, mid_channels, kernel_size=1, bias=False)\n        self.bn1 = nn.BatchNorm3d(mid_channels)\n        self.act1 = nn.ReLU(inplace=True)\n\n        self.conv2 = nn.Sequential(\n            nn.Conv3d(mid_channels, mid_channels, kernel_size=1, stride=1, bias=False),\n            nn.Conv3d(\n                mid_channels,\n                mid_channels,\n                kernel_size=3,\n                stride=(1, stride, stride),\n                padding=1,\n                bias=False,\n                groups=mid_channels,\n            ),\n        )\n        self.bn2 = nn.BatchNorm3d(mid_channels)\n        self.act2 = nn.ReLU(inplace=True)\n\n        self.conv3 = nn.Conv3d(mid_channels, out_channels, kernel_size=1, bias=False)\n        self.bn3 = nn.BatchNorm3d(out_channels)\n\n        self.act3 = nn.ReLU(inplace=True)\n        self.downsample = downsample_conv(\n            in_channels,\n            out_channels,\n            stride=stride,\n        )\n        self.stride = stride\n        self.zero_init_last()\n\n    def zero_init_last(self):\n        if getattr(self.bn3, \"weight\", None) is not None:\n            nn.init.zeros_(self.bn3.weight)\n\n    def forward(self, x: torch.Tensor):\n        shortcut = x\n\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.act1(x)\n\n        x = self.conv2(x)\n        x = self.bn2(x)\n        x = self.act2(x)\n\n        x = self.conv3(x)\n        x = self.bn3(x)\n\n        if self.downsample is not None:\n            shortcut = self.downsample(shortcut)\n        x += shortcut\n        x = self.act3(x)\n\n        return x\n    \nclass ContrailsModel2_5D(nn.Module):\n    def __init__(\n        self,\n        model_name: str = \"resnet34\",\n        pretrained: bool = False,\n        drop_path_rate: float = 0,\n        in_chans: int = 3,\n        num_class: int = 1,\n        image_size: int = 256,\n        seq_len: int = 7,\n    ):\n        super().__init__()\n        self.encoder = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            in_chans=in_chans,\n            features_only=True,\n            drop_path_rate=drop_path_rate,\n        )\n        self.seq_len = seq_len\n        self.output_fmt = getattr(self.encoder, \"output_fmt\", \"NHCW\")\n        self.in_chans = in_chans\n        assert image_size % 32 == 0\n        self.image_size = image_size\n        num_features = self.encoder.feature_info.channels()\n        self.conv3d = nn.Sequential(\n            *[\n                ResidualConv3D(\n                    num_features[-1],\n                    num_features[-1] // 4,\n                    num_features[-1],\n                    1,\n                )\n                for _ in range(3)\n            ]\n        )\n        self.head = nn.Sequential(\n            nn.Conv2d(num_features[-1], 512, kernel_size=1),\n            AdaptiveConcatPool2d(1),\n            Flatten(),\n            nn.ReLU(),\n            nn.Linear(1024, 512),\n            nn.BatchNorm1d(512),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(512, num_class),\n        )\n\n    def forward_image_feats(self, img):\n        # img -> (bs, 3 * 8, h, w)\n        img = F.interpolate(\n            img, size=(self.image_size, self.image_size), mode=\"bilinear\"\n        )\n        bs, _, h, w = img.shape\n        img = img.reshape(bs, 3, 8, h, w)\n        assert 4 - self.seq_len // 2 > 0\n        img = img[:, :, 4 - self.seq_len // 2 : 4 - self.seq_len // 2 + self.seq_len]\n        img = img.permute((0, 2, 1, 3, 4)).reshape(bs * self.seq_len, 3, h, w)\n        img_feats = self.encoder(img)\n        if self.output_fmt == \"NHWC\":\n            img_feats = [\n                img_feat.permute(0, 3, 1, 2).contiguous() for img_feat in img_feats\n            ]\n        img_feat = img_feats[-1]\n        _, ch, h, w = img_feat.shape  # (bs * seq_len, ch, h, w)\n        img_feat = img_feat.reshape(bs, self.seq_len, ch, h, w).transpose(\n            1, 2\n        )  # (bs, ch, seq_len, h, w)\n        img_feat = self.conv3d(img_feat)[\n            :, :, self.seq_len // 2\n        ]  # (bs, ch, seq_len, h, w) -> (bs, ch, h, w)\n\n        return img_feat\n\n    def forward_head(self, img_feats):\n        output = self.head(img_feats)\n        return output\n\n    def forward(\n        self,\n        img: torch.Tensor,\n    ):\n        \"\"\"\n        img: (bs, ch, h, w)\n        \"\"\"\n        img_feats = self.forward_image_feats(img)\n        return self.forward_head(img_feats)\n    \nclass ContrailsLightningModel2_5D(pl.LightningModule):\n    def __init__(\n        self,\n        model_name: str = \"resnet34\",\n        pretrained: bool = False,\n        drop_path_rate: float = 0,\n        image_size: int = 256,\n        in_chans: int = 3,\n        num_class: int = 1,\n        seq_len: int = 7,\n        mixup_p: float = 0.0,\n        mixup_alpha: float = 0.5,\n        no_mixup_epochs: int = 0,\n        lr: float = 1e-3,\n        backbone_lr: float = None,\n        disable_compile: bool = False,\n    ) -> None:\n        super().__init__()\n        self.lr = lr\n        self.backbone_lr = backbone_lr if backbone_lr is not None else lr\n        self.__build_model(\n            model_name=model_name,\n            pretrained=pretrained,\n            drop_path_rate=drop_path_rate,\n            in_chans=in_chans,\n            image_size=image_size,\n            num_class=num_class,\n            seq_len=seq_len,\n        )\n        if not disable_compile:\n            self.__compile_model()\n        self.save_hyperparameters()\n\n    def __build_model(\n        self,\n        model_name: str = \"resnet34\",\n        pretrained: bool = False,\n        drop_path_rate: float = 0,\n        in_chans: int = 3,\n        num_class: int = 1,\n        image_size: int = 256,\n        seq_len: int = 7,\n    ):\n        self.model = None\n        self.model_ema = ModelEmaV2(ContrailsModel2_5D(\n            model_name=model_name,\n            pretrained=pretrained,\n            drop_path_rate=drop_path_rate,\n            in_chans=in_chans,\n            num_class=num_class,\n            seq_len=seq_len,\n            image_size=image_size,\n        ), decay=0.998)\n\n    def __compile_model(self):\n        # self.model = torch.compile(self.model)\n        self.model_ema = torch.compile(self.model_ema)\n        \n\nclass ContrailsSegModel2_5D(nn.Module):\n    def __init__(\n        self,\n        model_name: str = \"resnet34\",\n        pretrained: bool = False,\n        drop_path_rate: float = 0,\n        decoder_type: str = \"UNet\",  # UNet or FastFCNImprove\n        center=None,\n        attention_type=None,\n        in_chans: int = 3,\n        num_class: int = 1,\n        image_size: int = 256,\n        seq_len: int = 7,\n    ):\n        super().__init__()\n        self.encoder = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            in_chans=in_chans,\n            features_only=True,\n            drop_path_rate=drop_path_rate,\n        )\n        self.seq_len = seq_len\n        self.output_fmt = getattr(self.encoder, \"output_fmt\", \"NHCW\")\n        self.in_chans = in_chans\n        assert image_size % 32 == 0\n        self.image_size = image_size\n        num_features = self.encoder.feature_info.channels()\n        conv3d = []\n        for ch_3d in num_features:\n            conv3d.append(\n                nn.Sequential(\n                    *[\n                        ResidualConv3D(\n                            ch_3d,\n                            ch_3d // 4,\n                            ch_3d,\n                            1,\n                        )\n                        for _ in range(3)\n                    ]\n                )\n            )\n        self.conv3d = nn.ModuleList(conv3d)\n        del conv3d\n        gc.collect()\n        if decoder_type == \"UNet\":\n            self.head = UNetHead(\n                encoder_channels=num_features,\n                num_class=num_class,\n                center=center,\n                attention_type=attention_type,\n                classification=False,\n                deep_supervision=False,\n            )\n        elif decoder_type == \"FastFCNImprove\":\n            self.head = FastFCNImproveHead(\n                encoder_channels=num_features,\n                num_class=num_class,\n                attention_type=attention_type,\n                classification=False,\n                deep_supervision=False,\n            )\n        else:\n            raise NotImplementedError\n\n    def forward_image_feats(self, img):\n        # img -> (bs, 3 * 8, h, w)\n        img = F.interpolate(\n            img, size=(self.image_size, self.image_size), mode=\"bilinear\"\n        )\n        bs, _, h, w = img.shape\n        img = img.reshape(bs, 3, 8, h, w)\n        assert 4 - self.seq_len // 2 > 0\n        img = img[:, :, 4 - self.seq_len // 2 : 4 - self.seq_len // 2 + self.seq_len]\n        img = img.permute((0, 2, 1, 3, 4)).reshape(bs * self.seq_len, 3, h, w)\n        img_feats = self.encoder(img)\n        if self.output_fmt == \"NHWC\":\n            img_feats = [\n                img_feat.permute(0, 3, 1, 2).contiguous() for img_feat in img_feats\n            ]\n        for i in range(len(img_feats)):\n            img_feat = img_feats[i]\n            _, ch, h, w = img_feat.shape  # (bs * seq_len, ch, h, w)\n            img_feat = img_feat.reshape(bs, self.seq_len, ch, h, w).transpose(\n                1, 2\n            )  # (bs, ch, seq_len, h, w)\n            img_feats[i] = self.conv3d[i](img_feat)[\n                :, :, self.seq_len // 2\n            ]  # (bs, ch, seq_len, h, w) -> (bs, ch, h, w)\n        return img_feats\n\n    def forward_head(self, img_feats):\n        output = self.head(img_feats)\n        return F.interpolate(output, size=(256, 256), mode=\"bilinear\")\n\n    def forward(\n        self,\n        img: torch.Tensor,\n    ):\n        \"\"\"\n        img: (bs, ch, h, w)\n        \"\"\"\n        img_feats = self.forward_image_feats(img)\n        return self.forward_head(img_feats)\n    \n    \nclass ContrailsLightningSegModel2_5D(pl.LightningModule):\n    def __init__(\n        self,\n        model_name: str = \"resnet34\",\n        pretrained: bool = False,\n        drop_path_rate: float = 0,\n        decoder_type: str = \"UNet\",  # UNet or FastFCNImprove\n        center=None,\n        attention_type=None,\n        image_size: int = 256,\n        in_chans: int = 3,\n        num_class: int = 1,\n        seq_len: int = 7,\n        mixup_p: float = 0.0,\n        mixup_alpha: float = 0.5,\n        no_mixup_epochs: int = 0,\n        dice_ratio: float = 0.25,\n        lr: float = 1e-3,\n        backbone_lr: float = None,\n        disable_compile: bool = False,\n    ) -> None:\n        super().__init__()\n        self.__build_model(\n            model_name=model_name,\n            pretrained=pretrained,\n            drop_path_rate=drop_path_rate,\n            decoder_type=decoder_type,\n            center=center,\n            attention_type=attention_type,\n            in_chans=in_chans,\n            num_class=num_class,\n            image_size=image_size,\n            seq_len=seq_len,\n        )\n        if not disable_compile:\n            self.__compile_model()\n        self.save_hyperparameters()\n\n    def __build_model(\n        self,\n        model_name: str = \"resnet34\",\n        pretrained: bool = False,\n        drop_path_rate: float = 0,\n        decoder_type: str = \"UNet\",  # UNet or FastFCNImprove\n        center=None,\n        attention_type=None,\n        in_chans: int = 3,\n        num_class: int = 1,\n        image_size: int = 256,\n        seq_len: int = 7,\n    ):\n        self.model = None\n        self.model_ema = ModelEmaV2(ContrailsSegModel2_5D(\n            model_name=model_name,\n            pretrained=pretrained,\n            drop_path_rate=drop_path_rate,\n            decoder_type=decoder_type,\n            center=center,\n            attention_type=attention_type,\n            in_chans=in_chans,\n            num_class=num_class,\n            image_size=image_size,\n            seq_len=seq_len,\n        ), decay=0.998)\n\n    def __compile_model(self):\n        # self.model = torch.compile(self.model)\n        self.model_ema = torch.compile(self.model_ema)\n        ","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-08-08T06:58:15.026188Z","iopub.execute_input":"2023-08-08T06:58:15.026534Z","iopub.status.idle":"2023-08-08T06:58:15.081239Z","shell.execute_reply.started":"2023-08-08T06:58:15.026502Z","shell.execute_reply":"2023-08-08T06:58:15.080291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# version 21: score: 0.7068770580047657 cls_score: 0.89058039961941 seg_threshold: 0.9931640625000013(0.5787000231213124) cls_threshold: 0.7312500000000005(0.4255064306780771)\n\nfrom dataclasses import dataclass, field\nimport pickle\n\n@dataclass\nclass tattaka_args_class:\n    seed: int = 0\n    num_workers: int = 2\n    batch_size: int = 2\n    cls_model_conf: dict = field(default_factory=dict)\n    seg_model_stage1_top100_conf: dict = field(default_factory=dict)\n    seg_model_stage1_conf: dict = field(default_factory=dict)\n    seg_model_stage2_conf: dict = field(default_factory=dict)\n    cls_threshold_percent: float = 0.5\n    seg_threshold_percent: float = 0.5\n\ntattaka_args = tattaka_args_class(\n    seed=42,\n    num_workers=2,\n    batch_size=2,\n    cls_model_conf={\n#         \"/kaggle/input/contrailseg-weights-exp043/convnext_base_cls_stage1/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnest101e_cls_stage1_320/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043/resnetrs101_cls_stage1/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnetrs50_cls_stage1/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043/swinv2_base_window16_cls_stage1/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/swin_base_patch4_window12_cls_stage1/fold0\": {\"weight\": 1},\n        \n    },\n    seg_model_stage1_top100_conf={\n#         \"/kaggle/input/contrailseg-weights-exp043/convnext_base_unet_cbam_stage1_ep30/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043/resnest101e_fastfcn_stage1_320_ep30/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043/resnetrs101_unet_stage1_ep30/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnetrs200_fastfcn_stage1_384_ep25/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnetrs50_unet_stage1_ep30/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043/swinv2_base_window16_unet_stage1_ep30/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/swin_base_patch4_window12_unet_stage1_ep30/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043-2/resnetrs101_512_unet_stage1_ep30/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/convnext_large_unet_stage1_ep20/fold0\": {\"weight\": 1},\n\n        \"/kaggle/input/contrailseg-weights-exp055/resnest101e_320_fastfcn_stage1_ep30/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp055/convnext_base_512_unet_stage1_ep25/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp055/resnetrs101_384_unet_stage1_ep30/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp055/resnetrs101_512_fastfcn_stage1_ep25/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp055/resnetrs101_512_unet_stage1_ep30/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp055/convnext_large_384_unet_stage1_ep20\": {\"weight\": 1},\n    },\n    seg_model_stage1_conf={\n#         \"/kaggle/input/contrailseg-weights-exp043/convnext_base_unet_cbam_stage1_ep30/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnest101e_fastfcn_stage1_320_ep30/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnetrs101_unet_stage1_ep30/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnetrs200_fastfcn_stage1_384_ep25/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnetrs50_unet_stage1_ep30/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/swinv2_base_window16_unet_stage1_ep30/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043/swin_base_patch4_window12_unet_stage1_ep30/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043-2/resnetrs101_512_unet_stage1_ep30/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043-2/convnext_base_512_unet_stage1_ep25/fold0\": {\"weight\": 1},\n        \n        \n        \"/kaggle/input/contrailseg-weights-exp055/convnext_base_512_unet_stage1_ep25/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp055/resnetrs101_384_unet_stage1_ep30/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp055/resnetrs101_512_fastfcn_stage1_ep25/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp055/resnetrs101_512_unet_stage1_ep30/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp055-2/convnext_base_512_unet_cbam_stage1_ep12/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/convnext_large_unet_stage1_ep20/fold0\": {\"weight\": 1},\n    },\n    seg_model_stage2_conf={\n#         \"/kaggle/input/contrailseg-weights-exp043/convnext_base_unet_cbam_stage2_ep60/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnest101e_fastfcn_stage2_320_ep60/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnetrs101_unet_stage2_ep60/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnetrs200_fastfcn_stage2_384_ep50/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/resnetrs50_unet_stage2_ep60/fold0\": {\"weight\": 1},\n#         \"/kaggle/input/contrailseg-weights-exp043/swinv2_base_window16_unet_stage2_ep60/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043/swin_base_patch4_window12_unet_stage2_ep60/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043-2/resnetrs101_512_unet_stage2_ep60/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp043-2/convnext_base_512_unet_stage2_ep50/fold0\": {\"weight\": 1},\n\n        \"/kaggle/input/contrailseg-weights-exp055/convnext_base_512_unet_stage2_ep50/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp055/resnetrs101_384_unet_stage2_ep60/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp055/resnetrs101_512_fastfcn_stage2_ep50/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp055/resnetrs101_512_unet_stage2_ep60/fold0\": {\"weight\": 1},\n        \"/kaggle/input/contrailseg-weights-exp055-2/convnext_base_512_unet_cbam_stage2_ep30/fold0\": {\"weight\": 1},\n        \n#         \"/kaggle/input/contrailseg-weights-exp043/convnext_large_unet_stage2_ep40/fold0\": {\"weight\": 1}, \n    },\n    cls_threshold_percent=0.7312500000000005,\n    seg_threshold_percent=0.9931640625000013,\n)\n    \n@dataclass\nclass tawara_args_class:\n    TRAINED_MODELS_ALL_TAWARA: list = field(default_factory=list)\n    TRAINED_MODELS_POS_TAWARA: list = field(default_factory=list)\n            \ntawara_args = tawara_args_class(\n    TRAINED_MODELS_ALL_TAWARA = [\n        # # hard label\n#         \"exp052/2_tu-resnet34d_unet_256\",\n#         \"exp053/1_tu-res2net50d_unet_256\",\n#         \"exp053/4_tu-regnetz_c16_unet_256\",\n#         \"exp053/5_tu-tf_efficientnetv2_s_unet_256\",\n#         INPUT / \"gricrgw-weights/exp053/0_tu-resnet34d_unet_384\",\n        INPUT / \"gricrgw-weights/exp053/2_tu-res2net50d_unet_384\",\n        INPUT / \"gricrgw-weights/exp053/13_tu-regnetz_c16_unet_384\",\n        INPUT / \"gricrgw-weights/exp053/16_tu-regnetz_d8_unet_384\",\n        INPUT / \"gricrgw-weights/exp053/17_tu-regnetz_d32_unet_384\",\n        INPUT / \"gricrgw-weights/exp053/18_tu-regnetz_e8_unet_384\",\n        # # soft label\n#         \"exp057/0_tu-resnet34d_unet_256\",\n#         \"exp057/1_tu-res2net50d_unet_256\",\n#         \"exp057/2_tu-regnetz_c16_unet_256\",\n#         \"exp057/3_tu-tf_efficientnetv2_s_unet_256\",\n#         INPUT / \"gricrgw-weights/exp057/4_tu-resnet34d_unet_384\",\n        INPUT / \"gricrgw-weights/exp057/5_tu-res2net50d_unet_384\",\n        INPUT / \"gricrgw-weights/exp057/6_tu-regnetz_c16_unet_384\",\n        INPUT / \"gricrgw-weights/exp057/7_tu-regnetz_d8_unet_384\",\n        INPUT / \"gricrgw-weights/exp057/8_tu-regnetz_d32_unet_384\",\n        INPUT / \"gricrgw-weights/exp057/9_tu-regnetz_e8_unet_384\",\n    ], \n    TRAINED_MODELS_POS_TAWARA = [\n        # # hard label\n#         \"exp056/0_tu-resnet34d_unet_256\",\n#         \"exp056/1_tu-res2net50d_unet_256\",\n#         \"exp056/2_tu-regnetz_c16_unet_256\",\n#         \"exp056/3_tu-tf_efficientnetv2_s_unet_256\",\n#         INPUT / \"gricrgw-weights/exp056/4_tu-resnet34d_unet_384\",\n        INPUT / \"gricrgw-weights/exp056/5_tu-res2net50d_unet_384\",\n        INPUT / \"gricrgw-weights/exp056/6_tu-regnetz_c16_unet_384\",\n        INPUT / \"gricrgw-weights/exp056/7_tu-regnetz_d8_unet_384\",\n        INPUT / \"gricrgw-weights/exp056/8_tu-regnetz_d32_unet_384\",\n        INPUT / \"gricrgw-weights/exp056/9_tu-regnetz_e8_unet_384\",\n    ],\n)\n\n#     SEG_THRESHOLD = 0.4\n#     TEST_CLS_THRESHOLD_PERCENT     = 0.7026\n#     TEST_SEG_THRESHOLD_PERCENT_ALL = 0.9982\n#     TEST_SEG_THRESHOLD_PERCENT_POS = 0.9939","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:15.082797Z","iopub.execute_input":"2023-08-08T06:58:15.083318Z","iopub.status.idle":"2023-08-08T06:58:15.102709Z","shell.execute_reply.started":"2023-08-08T06:58:15.08328Z","shell.execute_reply":"2023-08-08T06:58:15.101898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def stage1_inference_classification_tattaka(dataloader, cls_models_stage1):\n    cls_logits_all = []\n    with torch.no_grad():\n        for batch in tqdm(dataloader):\n            image = batch\n            # image = torch.stack(\n            #     [image, image.flip(2), image.flip(3), image.flip(2).flip(3)]\n            # )\n            image = torch.stack([image])\n            tta_num, bs, c, h, w = image.shape\n            image = image.reshape((tta_num * bs, c, h, w))\n            image = image.half().to(device=device)\n            cls_logits = np.stack(\n                [torch.sigmoid(model(image))\n                .reshape((tta_num, bs, -1))\n                .mean(0)\n                .detach()\n                .cpu()\n                .numpy() for model in cls_models_stage1\n                ]\n            )  # (model_len, bs, 1)\n            cls_logits = cls_logits.sum(0)\n            cls_logits_all.append(cls_logits)\n            \n    cls_logits_all = np.concatenate(cls_logits_all)\n    return cls_logits_all\n\n\ndef stage1_inference_top100_tattaka(dataloader, seg_models_stage1):\n    cls_logits_all = []\n    with torch.no_grad():\n        for batch in tqdm(dataloader):\n            image = batch\n            # image = torch.stack(\n            #     [image, image.flip(2), image.flip(3), image.flip(2).flip(3)]\n            # )\n            image = torch.stack([image])\n            tta_num, bs, c, h, w = image.shape\n            image = image.reshape((tta_num * bs, c, h, w))\n            image = image.half().to(device=device)\n            seg_logits_stage1 = torch.stack([torch.sigmoid(\n                seg_model_stage1(image)\n            ).reshape((tta_num, bs, -1, h, w)) for seg_model_stage1 in seg_models_stage1]) # (model_len, tta_num, bs, 1, h, w)\n            # if use TTA, reverse logits\n            seg_logits_stage1 = seg_logits_stage1.mean(1).detach().cpu().numpy() # (model_len, bs, 1, h, w)\n            seg_topk_stage1 = np.mean(\n                np.sort(seg_logits_stage1.reshape(len(seg_models_stage1), bs, -1), axis=-1)[:, :, -100:],\n                axis=-1,\n                keepdims=True,\n            ) # (model_len, bs, 1)\n\n            cls_logits = seg_topk_stage1.sum(0)\n            cls_logits_all.append(cls_logits)\n            \n    cls_logits_all = np.concatenate(cls_logits_all)\n    return cls_logits_all","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:15.104372Z","iopub.execute_input":"2023-08-08T06:58:15.104752Z","iopub.status.idle":"2023-08-08T06:58:15.120239Z","shell.execute_reply.started":"2023-08-08T06:58:15.104705Z","shell.execute_reply":"2023-08-08T06:58:15.119285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_inference_loop(cfg, model, loader, device):\n    img_size = cfg[\"/globals/img_size\"]\n    pred_list = []\n    # print(img_size)\n    with torch.no_grad():\n        for batch in tqdm(loader):\n            x = to_device(batch[\"data\"], device).half()\n            x = torch.nn.functional.interpolate(x, size=img_size, mode='bilinear')\n            y = model(x)\n            y = torch.nn.functional.interpolate(y, size=256, mode='bilinear')\n            pred_list.append(y.sigmoid().detach().cpu().numpy())\n        \n        pred_arr = np.concatenate(pred_list)\n    del pred_list\n    gc.collect()\n    return pred_arr\n\n\ndef get_test_pred_arr_tawara(\n    test_df: pd.DataFrame, trained_model_paths: list[Path], return_top100mean: bool=False\n):\n    \"\"\"\"\"\"\n    test_image_paths = [\n        TMP_TEST / str(record_id) / \"tdiff_plus_all_band.npy\"\n        for record_id in test_df[\"record_id\"].values\n    ]\n    test_transform = A.Compose([\n        # A.Resize(p=1.0, height=IMG_SIZE, width=IMG_SIZE),\n        A.Normalize(p=1.0,\n            mean=[ -2.7986, 0.9446, 233.6702, 242.2449, 250.7397, 274.4096, 255.5284, 276.5997, 275.3542, 272.5556, 260.4157],\n            std=[   1.2727, 2.1788,   7.0210,   9.1716,  11.3542,  19.6210,  13.1217,  20.7206,  21.1159,  20.5739,  15.8393],\n            max_pixel_value=1,\n        ),\n        ToTensorV2(p=1.0, transpose_mask=True)\n    ])\n    test_dataset = GRICRGWTestDataset(test_image_paths, test_transform)\n    test_loader = data.DataLoader(\n        dataset=test_dataset, batch_size=32, num_workers=2, shuffle=False, drop_last=False)\n    \n    \n    test_pred_arr = np.zeros((len(test_df), 1, 256, 256), dtype=\"float16\")\n    test_pred_arr_top100mean = np.zeros((len(test_df), 1,), dtype=\"float16\")\n    \n    for trained_model in trained_model_paths:\n        trained_model = Path(trained_model)\n        pre_eval = load_yaml_file(trained_model / \"config.yml\")\n        pre_eval[\"model\"][\"encoder_weights\"] = None\n\n        cfg = Config(pre_eval, types={\"Unet\": smp.Unet})\n        model = cfg[\"/model\"]\n        weight = torch.load(trained_model / f\"best_model.pth\", map_location=device)\n        for key in weight.keys():\n            weight[key] = weight[key].half()\n        ###\n        torch.save(weight, \"/kaggle/working/tmp/best_model.pth\")\n        del weight\n        model.load_state_dict(torch.load(\"/kaggle/working/tmp/best_model.pth\", map_location=device))\n        \n        model = model.half().to(device)\n        model.eval()\n        gc.collect()\n        pred_arr = run_inference_loop(cfg, model, test_loader, device)\n        test_pred_arr += pred_arr\n        \n        if return_top100mean:\n            pred_arr_top100mean = np.sort(\n                pred_arr.reshape(pred_arr.shape[0], 256 * 256), axis=1\n            )[:, -100:].mean(axis=1, keepdims=True)\n            test_pred_arr_top100mean += pred_arr_top100mean\n        model = model.to(device=\"cpu\")\n        del model, cfg\n        gc.collect()\n        torch.cuda.empty_cache()\n    del test_loader\n    gc.collect()\n#     test_pred_arr /= len(trained_model_paths)\n#     test_pred_arr_top100mean /= len(trained_model_paths)\n    \n    if return_top100mean:\n        return test_pred_arr, test_pred_arr_top100mean\n    \n    return test_pred_arr","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:59:57.807267Z","iopub.execute_input":"2023-08-08T06:59:57.807681Z","iopub.status.idle":"2023-08-08T06:59:57.825807Z","shell.execute_reply.started":"2023-08-08T06:59:57.807649Z","shell.execute_reply":"2023-08-08T06:59:57.824797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pl.seed_everything(tattaka_args.seed)\nwarnings.simplefilter(\"ignore\")\n\ntest_pred_arr, test_pred_arr_all_top100mean = get_test_pred_arr_tawara(test_df, tawara_args.TRAINED_MODELS_ALL_TAWARA, return_top100mean=True)\n# del tmp\n# gc.collect()\n\ncls_models_stage1 = []\nfor logdir in tattaka_args.cls_model_conf.keys():\n    checkpoint = glob(f\"{logdir}/**/best_dice.ckpt\", recursive=True)[0]\n    weight = torch.load(checkpoint)\n    # to fp16\n    for key in weight[\"state_dict\"].keys():\n        weight[\"state_dict\"][key] = weight[\"state_dict\"][key].half()\n    ###\n    torch.save(weight, \"/kaggle/working/tmp/best_dice.ckpt\")\n    del weight\n    gc.collect()\n    cls_models_stage1.append(ContrailsLightningModel2_5D.load_from_checkpoint(\"/kaggle/working/tmp/best_dice.ckpt\", pretrained=False, strict=False).model_ema.module.eval().half().to(device=device))\nprint(\"cls_model loaded: \", [glob(f\"{logdir_cls}/**/best_dice.ckpt\", recursive=True)[0] for logdir_cls in tattaka_args.cls_model_conf.keys()])\n\ndataloader = ContrailsDataModule(\n    train_df=test_df,\n    valid_df=test_df,\n    num_workers=tattaka_args.num_workers,\n    batch_size=tattaka_args.batch_size,\n).test_dataloader()\ncls_logits_all = stage1_inference_classification_tattaka(dataloader, cls_models_stage1)\ndel cls_models_stage1, dataloader\ngc.collect()\ntorch.cuda.empty_cache()\n\nseg_models_stage1 = []\nfor logdir in tattaka_args.seg_model_stage1_top100_conf.keys():\n    checkpoint = glob(f\"{logdir}/**/best_dice.ckpt\", recursive=True)[0]\n    weight = torch.load(checkpoint)\n    # to fp16\n    for key in weight[\"state_dict\"].keys():\n        weight[\"state_dict\"][key] = weight[\"state_dict\"][key].half()\n    ###\n    torch.save(weight, \"/kaggle/working/tmp/best_dice.ckpt\")\n    del weight\n    gc.collect()\n    seg_models_stage1.append(ContrailsLightningSegModel2_5D.load_from_checkpoint(\"/kaggle/working/tmp/best_dice.ckpt\", pretrained=False, strict=False).model_ema.module.eval().half().to(device=device))\nprint(\"seg_model_stage1 loaded: \", [glob(f\"{logdir_seg1}/**/best_dice.ckpt\", recursive=True)[0] for logdir_seg1 in tattaka_args.seg_model_stage1_top100_conf.keys()])\ndataloader = ContrailsDataModule(\n    train_df=test_df,\n    valid_df=test_df,\n    num_workers=tattaka_args.num_workers,\n    batch_size=tattaka_args.batch_size,\n).test_dataloader()\nseg_top100_all = stage1_inference_top100_tattaka(dataloader, seg_models_stage1) \ndel seg_models_stage1, dataloader\ngc.collect()\ntorch.cuda.empty_cache()\n\ncls_logits_all = (cls_logits_all + seg_top100_all + test_pred_arr_all_top100mean) / (len(tattaka_args.cls_model_conf) + len(tattaka_args.seg_model_stage1_top100_conf) + len(tawara_args.TRAINED_MODELS_ALL_TAWARA)) # (bs, 1)\ncls_threshold = np.quantile(cls_logits_all, tattaka_args.cls_threshold_percent)\ncls_preds_all = cls_logits_all > cls_threshold","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:59:59.321211Z","iopub.execute_input":"2023-08-08T06:59:59.321576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def stage2_inference_tattaka(dataloader, seg_models):\n    seg_logits_all = []\n    with torch.no_grad():\n        for batch in tqdm(dataloader):\n            image = batch\n            image = torch.stack([image])\n            tta_num, bs, c, h, w = image.shape\n            image = image.reshape((tta_num * bs, c, h, w))\n            image = image.half().to(device=device)\n            seg_logits = torch.stack([torch.sigmoid(\n                seg_model(image)\n            ).reshape((tta_num, bs, -1, h, w)) for seg_model in seg_models]) # (model_len, tta_num, bs, 1, h, w)\n            # if use TTA, reverse logits\n            seg_logits = seg_logits.mean(1).detach().cpu().numpy() # (model_len, bs, 1, h, w)\n            seg_logits = seg_logits.sum(0)\n            seg_logits_all.append(seg_logits)\n\n    seg_logits_all_cat = np.concatenate(seg_logits_all)\n    del image, seg_logits, seg_logits_all\n    gc.collect()\n    return seg_logits_all_cat\n","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:23.752144Z","iopub.status.idle":"2023-08-08T06:58:23.753359Z","shell.execute_reply.started":"2023-08-08T06:58:23.753109Z","shell.execute_reply":"2023-08-08T06:58:23.753134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"warnings.simplefilter(\"ignore\")\n\npos_idx = cls_preds_all[:, 0] > 0.5\n\n# test_pred_arr_pos = get_test_pred_arr_tawara(test_df[pos_idx], tawara_args.TRAINED_MODELS_POS_TAWARA + tawara_args.TRAINED_MODELS_ALL_TAWARA)\ntest_pred_arr_pos = test_pred_arr[pos_idx]\n\nseg_models_stage1 = []\nfor logdir in tattaka_args.seg_model_stage1_conf.keys():\n    checkpoint = glob(f\"{logdir}/**/best_dice.ckpt\", recursive=True)[0]\n    weight = torch.load(checkpoint)\n    # to fp16\n    for key in weight[\"state_dict\"].keys():\n        weight[\"state_dict\"][key] = weight[\"state_dict\"][key].half()\n    ###\n    torch.save(weight, \"/kaggle/working/tmp/best_dice.ckpt\")\n    del weight\n    gc.collect()\n    seg_models_stage1.append(ContrailsLightningSegModel2_5D.load_from_checkpoint(\"/kaggle/working/tmp/best_dice.ckpt\", pretrained=False, strict=False).model_ema.module.eval().half().to(device=device))\nprint(\"seg_model_stage1 loaded: \", [glob(f\"{logdir_seg1}/**/best_dice.ckpt\", recursive=True)[0] for logdir_seg1 in tattaka_args.seg_model_stage1_conf.keys()])\n\ndataloader = ContrailsDataModule(\n    train_df=test_df,\n    valid_df=test_df[pos_idx],\n    num_workers=tattaka_args.num_workers,\n    batch_size=tattaka_args.batch_size,\n).test_dataloader()\nseg_logits_all1 = stage2_inference_tattaka(dataloader, seg_models_stage1) \ndel seg_models_stage1, dataloader\ngc.collect()\ntorch.cuda.empty_cache()\n\nseg_models_stage2 = []\nfor logdir in tattaka_args.seg_model_stage2_conf.keys():\n    checkpoint = glob(f\"{logdir}/**/best_dice.ckpt\", recursive=True)[0]\n    weight = torch.load(checkpoint)\n    # to fp16\n    for key in weight[\"state_dict\"].keys():\n        weight[\"state_dict\"][key] = weight[\"state_dict\"][key].half()\n    ###\n    torch.save(weight, \"/kaggle/working/tmp/best_dice.ckpt\")\n    del weight\n    gc.collect()\n    seg_models_stage2.append(ContrailsLightningSegModel2_5D.load_from_checkpoint(\"/kaggle/working/tmp/best_dice.ckpt\", pretrained=False, strict=False).model_ema.module.eval().half().to(device=device))\nprint(\"seg_model_stage2 loaded: \", [glob(f\"{logdir_seg2}/**/best_dice.ckpt\", recursive=True)[0] for logdir_seg2 in tattaka_args.seg_model_stage2_conf.keys()])\n\ndataloader = ContrailsDataModule(\n    train_df=test_df,\n    valid_df=test_df[pos_idx],\n    num_workers=tattaka_args.num_workers,\n    batch_size=tattaka_args.batch_size,\n).test_dataloader()\nseg_logits_all2 = stage2_inference_tattaka(dataloader, seg_models_stage2)\ndel seg_models_stage2, dataloader\ngc.collect()\ntorch.cuda.empty_cache()\n\nseg_logits_all = (seg_logits_all1 + seg_logits_all2 + test_pred_arr_pos) / (len(tattaka_args.seg_model_stage1_conf) + len(tattaka_args.seg_model_stage2_conf) + len(tawara_args.TRAINED_MODELS_ALL_TAWARA) + len(tawara_args.TRAINED_MODELS_POS_TAWARA))\nseg_threshold = np.quantile(\n    seg_logits_all,\n    tattaka_args.seg_threshold_percent,\n)\nseg_preds_all = np.zeros((len(cls_preds_all), 256, 256), dtype=np.bool)\nseg_preds_all[pos_idx] = seg_logits_all[:, 0] > seg_threshold","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:23.754719Z","iopub.status.idle":"2023-08-08T06:58:23.755574Z","shell.execute_reply.started":"2023-08-08T06:58:23.755332Z","shell.execute_reply":"2023-08-08T06:58:23.755355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for record_id, mask in zip(test_df.record_id.to_numpy()[:2], seg_preds_all[:2]):\n    img = get_image(record_id).astype(np.float32)[..., 4]\n    plt.subplot(121)\n    plt.imshow(img)\n    plt.subplot(122)\n    plt.imshow(mask)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:23.757156Z","iopub.status.idle":"2023-08-08T06:58:23.757622Z","shell.execute_reply.started":"2023-08-08T06:58:23.75738Z","shell.execute_reply":"2023-08-08T06:58:23.757403Z"},"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\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\ndef rle_decode(mask_rle, shape=(256, 256)):\n    '''\n    mask_rle: run-length as string formatted (start length)\n              empty predictions need to be encoded with '-'\n    shape: (height, width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n    '''\n\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    if mask_rle != '-': \n        s = mask_rle.split()\n        starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n        starts -= 1\n        ends = starts + lengths\n        for lo, hi in zip(starts, ends):\n            img[lo:hi] = 1\n    return img.reshape(shape, order='F')  # Needed to align to RLE direction","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:23.759391Z","iopub.status.idle":"2023-08-08T06:58:23.759902Z","shell.execute_reply.started":"2023-08-08T06:58:23.759643Z","shell.execute_reply":"2023-08-08T06:58:23.759665Z"},"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')\nfor i, mask in enumerate(seg_preds_all):\n    current_image_id = submission.index.to_numpy()[i]\n    submission.loc[int(current_image_id), 'encoded_pixels'] = list_to_string(list_to_string(rle_encode(mask)))\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:23.76153Z","iopub.status.idle":"2023-08-08T06:58:23.762012Z","shell.execute_reply.started":"2023-08-08T06:58:23.761766Z","shell.execute_reply":"2023-08-08T06:58:23.761788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shutil.rmtree(\"/kaggle/working/tmp\")\nsubmission.to_csv(\"submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:23.763445Z","iopub.status.idle":"2023-08-08T06:58:23.76436Z","shell.execute_reply.started":"2023-08-08T06:58:23.764113Z","shell.execute_reply":"2023-08-08T06:58:23.764142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\n\n# These are the usual ipython objects, including this one you are creating\nipython_vars = ['In', 'Out', 'exit', 'quit', 'get_ipython', 'ipython_vars']\n\n# Get a sorted list of the objects and their sizes\nsorted([(x, sys.getsizeof(globals().get(x))/(1024*1024)) for x in dir() if not x.startswith('_') and x not in sys.modules and x not in ipython_vars], key=lambda x: x[1], reverse=True)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:58:23.765858Z","iopub.status.idle":"2023-08-08T06:58:23.76631Z","shell.execute_reply.started":"2023-08-08T06:58:23.766078Z","shell.execute_reply":"2023-08-08T06:58:23.7661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}