{"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":"!pip uninstall -y timm\nimport sys\nsys.path.append(\"../input/timmmaster/\")\nimport timm\nprint(timm.__version__)","metadata":{"execution":{"iopub.status.busy":"2023-05-25T15:43:47.871376Z","iopub.execute_input":"2023-05-25T15:43:47.87219Z","iopub.status.idle":"2023-05-25T15:43:56.194396Z","shell.execute_reply.started":"2023-05-25T15:43:47.872158Z","shell.execute_reply":"2023-05-25T15:43:56.192942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/mmdetection-offline-lib/mmcv_full-1.3.14-cp37-cp37m-linux_x86_64.whl --no-deps\n!pip install /kaggle/input/mmdetection-offline-lib/pycocotools-2.0.2-cp37-cp37m-linux_x86_64.whl --no-deps\n!pip install /kaggle/input/mmdetection-offline-lib/terminaltables-3.1.0-py3-none-any.whl --no-deps\n!pip install /kaggle/input/mmdetection-offline-lib/pytest_runner-5.3.1-py3-none-any.whl --no-deps\n!pip install /kaggle/input/mmdetection-offline-lib/mmpycocotools-12.0.3-cp37-cp37m-linux_x86_64.whl --no-deps\n!pip install /kaggle/input/mmdetection-offline-lib/terminal-0.4.0-py3-none-any.whl --no-deps\n!pip install /kaggle/input/mmdetection-offline-lib/mmdet-2.17.0-py3-none-any.whl --no-deps\n!pip install /kaggle/input/mmdetection-offline-lib/addict-2.4.0-py3-none-any.whl --no-deps\n!pip install /kaggle/input/mmdetection-offline-lib/yapf-0.31.0-py2.py3-none-any.whl --no-deps","metadata":{"execution":{"iopub.status.busy":"2023-05-25T15:43:56.197167Z","iopub.execute_input":"2023-05-25T15:43:56.197531Z","iopub.status.idle":"2023-05-25T15:47:18.973155Z","shell.execute_reply.started":"2023-05-25T15:43:56.197488Z","shell.execute_reply":"2023-05-25T15:47:18.971917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import argparse\nimport datetime\nimport math\nimport os\nimport warnings\nfrom functools import lru_cache, partial\nfrom glob import glob\nfrom typing import Callable, List, Tuple\n\nimport albumentations as albu\nimport numpy as np\nimport PIL.Image as Image\nimport pytorch_lightning as pl\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 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\nfrom timm.utils import ModelEmaV2\n\nfrom io import StringIO\n\nimport gc\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nPATCH_SIZE = 32\n","metadata":{"execution":{"iopub.status.busy":"2023-05-25T15:47:18.976418Z","iopub.execute_input":"2023-05-25T15:47:18.976862Z","iopub.status.idle":"2023-05-25T15:47:29.913403Z","shell.execute_reply.started":"2023-05-25T15:47:18.976802Z","shell.execute_reply":"2023-05-25T15:47:29.912164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Volume to patch(32x32 size npy)","metadata":{}},{"cell_type":"code","source":"import os\nimport warnings\n\nimport numpy as np\nimport PIL.Image as Image\nfrom tqdm.auto import tqdm, trange\n\nwarnings.simplefilter(\"ignore\")\n\nPREFIX = \"/kaggle/input/vesuvius-challenge-ink-detection/test\"\n\ntest_fragments = sorted(glob(f\"{PREFIX}/*\"))\nfragment_ids = [idx.split(\"/\")[-1] for idx in test_fragments]\nfor data_id in tqdm(fragment_ids):\n    mask = np.array(Image.open(PREFIX + f\"/{data_id}/mask.png\").convert(\"1\"))\n    volume = np.stack(\n        [\n            (np.array(Image.open(filename), dtype=np.float32) / 65535.0).astype(np.float16)\n            for filename in sorted(\n                glob(PREFIX + f\"/{data_id}/surface_volume/*.tif\")\n            )\n        ]\n    )\n    volume_dir = f\"vesuvius_patches_{PATCH_SIZE}/test/{data_id}/surface_volume/\"\n    mask_dir = f\"vesuvius_patches_{PATCH_SIZE}/test/{data_id}/mask/\"\n    os.makedirs(volume_dir, exist_ok=True)\n    os.makedirs(mask_dir, exist_ok=True)\n\n    h, w = volume.shape[-2:]\n    for i in trange(h // PATCH_SIZE, leave=False):\n        for j in range(w // PATCH_SIZE):\n            start_h = i * PATCH_SIZE\n            start_w = j * PATCH_SIZE\n            mask_patch = mask[\n                ..., start_h : start_h + PATCH_SIZE, start_w : start_w + PATCH_SIZE\n            ]\n            if not mask_patch.sum():\n                continue\n            volume_patch = volume[\n                ..., start_h : start_h + PATCH_SIZE, start_w : start_w + PATCH_SIZE\n            ].astype(np.float32)\n            np.save(os.path.join(volume_dir, f\"volume_{i}_{j}\"), volume_patch)\n            np.save(os.path.join(mask_dir, f\"mask_{i}_{j}\"), mask_patch)\n    del volume, mask\n    gc.collect()\n","metadata":{"execution":{"iopub.status.busy":"2023-05-25T15:47:29.916939Z","iopub.execute_input":"2023-05-25T15:47:29.917866Z","iopub.status.idle":"2023-05-25T15:49:53.07368Z","shell.execute_reply.started":"2023-05-25T15:47:29.917817Z","shell.execute_reply":"2023-05-25T15:49:53.072348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define dataloader, preprocess","metadata":{}},{"cell_type":"code","source":"def get_transforms(train: bool = False) -> Callable:\n    if train:\n        return albu.Compose(\n            [\n                albu.Flip(p=0.5),\n                albu.RandomRotate90(p=0.9),\n                albu.ShiftScaleRotate(\n                    shift_limit=0.0625,\n                    scale_limit=0.2,\n                    rotate_limit=15,\n                    p=0.9,\n                ),\n                albu.OneOf(\n                    [\n                        albu.ElasticTransform(p=0.3),\n                        albu.GaussianBlur(p=0.3),\n                        albu.GaussNoise(p=0.3),\n                        albu.OpticalDistortion(p=0.3),\n                        albu.GridDistortion(p=0.1),\n                        albu.PiecewiseAffine(p=0.3),  # IAAPiecewiseAffine\n                    ],\n                    p=0.9,\n                ),\n                albu.RandomBrightnessContrast(\n                    brightness_limit=0.3, contrast_limit=0.3, p=0.3\n                ),\n                ToTensorV2(),\n            ]\n        )\n\n    else:\n        return albu.Compose(\n            [\n                ToTensorV2(),\n            ]\n        )\n\n\nclass PatchDataset(Dataset):\n    def __init__(\n        self,\n        volume_paths: List[str],\n        image_size: Tuple[int, int] = (256, 256),\n        mode: str = \"train\",  # \"train\" | \"valid\" | \"test\"\n        preprocess_in_model: bool = False,\n        start_z: int = 8,\n        end_z: int = -8,\n        shift_z: int = 2,\n    ):\n        self.volume_paths = volume_paths\n        self.image_size = image_size\n        assert (image_size[0] % 32 == 0) and (image_size[1] % 32 == 0)\n        self.mode = mode\n        self.train = mode == \"train\"\n        self.transforms = get_transforms(self.train)\n        self.PATCH_SIZE = 32\n        self.preprocess_in_model = preprocess_in_model\n        self.start_z = start_z\n        self.end_z = end_z\n        self.shift_z = shift_z\n\n    def __len__(self) -> int:\n        if self.mode == \"train\":\n            return 25000\n        elif self.mode == \"valid\":\n            return 24000\n        else:\n            return len(self.volume_paths)\n\n#     @lru_cache(maxsize=64)\n    def np_load(self, path: str) -> np.ndarray:\n        return np.load(path)\n\n    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor, int, int]:\n        if self.train:\n            np_load = np.load\n            idx = np.random.choice(np.arange(len(self.volume_paths)))\n        else:\n            np_load = self.np_load\n        volume = np.zeros((65, *self.image_size), dtype=np.float32)\n        label = np.zeros(self.image_size)\n        volume_lt_path = self.volume_paths[idx]\n        data_prefix = \"/\".join(volume_lt_path.split(\"/\")[:-3])\n        data_source = volume_lt_path.split(\"/\")[-3]\n        y, x = volume_lt_path.split(\"/\")[-1].split(\".\")[-2].split(\"_\")[-2:]\n        x = int(x)\n        y = int(y)\n        for i in range(self.image_size[0] // self.PATCH_SIZE):\n            for j in range(self.image_size[1] // self.PATCH_SIZE):\n                volume_path = os.path.join(\n                    data_prefix,\n                    data_source,\n                    f\"surface_volume/volume_{y + i}_{x + j}.npy\",\n                )\n                label_path = os.path.join(\n                    data_prefix, data_source, f\"label/label_{y + i}_{x + j}.npy\"\n                )\n                if os.path.exists(volume_path):\n                    volume[\n                        :,\n                        i * self.PATCH_SIZE : (i + 1) * self.PATCH_SIZE,\n                        j * self.PATCH_SIZE : (j + 1) * self.PATCH_SIZE,\n                    ] = np_load(volume_path)\n                    if os.path.exists(label_path):\n                        label[\n                            i * self.PATCH_SIZE : (i + 1) * self.PATCH_SIZE,\n                            j * self.PATCH_SIZE : (j + 1) * self.PATCH_SIZE,\n                        ] = np_load(label_path)\n        if not self.preprocess_in_model:\n            if self.train and np.random.rand() < 0.5:\n                shift = np.random.randint(-self.shift_z, self.shift_z + 1)\n            else:\n                shift = 0\n            # TODO: Add test for shift range\n            volume = volume[\n                self.start_z + shift : self.end_z + shift\n            ]  # use middle 49 layer\n        volume = volume.transpose(1, 2, 0)\n        aug = self.transforms(image=volume, mask=label)\n        volume = aug[\"image\"]\n        label = aug[\"mask\"][None, :]\n        if self.train:\n            volume, label = grid_cutout(\n                volume=volume, label=label, max_height=2, max_width=2, prob=0.5\n            )\n        return (\n            volume.half(),\n            label.half(),\n            x,\n            y,\n        )\n\n\nclass InkDetDataModule(LightningDataModule):\n    def __init__(\n        self,\n        train_volume_paths: List[str],\n        valid_volume_paths: List[str],\n        image_size: int = 256,\n        num_workers: int = 4,\n        batch_size: int = 16,\n        preprocess_in_model: bool = False,\n        start_z: int = 8,\n        end_z: int = -8,\n        shift_z: int = 2,\n    ):\n        super().__init__()\n\n        self._num_workers = num_workers\n        self._batch_size = batch_size\n        self.train_volume_paths = train_volume_paths\n        self.valid_volume_paths = valid_volume_paths\n        self.image_size = (image_size, image_size)\n        self.preprocess_in_model = preprocess_in_model\n        self.start_z = start_z\n        self.end_z = end_z\n        self.shift_z = shift_z\n        self.save_hyperparameters(\n            \"num_workers\",\n            \"batch_size\",\n            \"image_size\",\n            \"preprocess_in_model\",\n            \"start_z\",\n            \"end_z\",\n            \"shift_z\",\n        )\n\n    def create_dataset(self, mode: str = \"train\") -> PatchDataset:\n        if mode == \"train\":\n            return PatchDataset(\n                volume_paths=self.train_volume_paths,\n                image_size=self.image_size,\n                mode=mode,\n                preprocess_in_model=self.preprocess_in_model,\n                start_z=self.start_z,\n                end_z=self.end_z,\n                shift_z=self.shift_z,\n            )\n        else:\n            return PatchDataset(\n                volume_paths=self.valid_volume_paths,\n                image_size=self.image_size,\n                mode=mode,\n                preprocess_in_model=self.preprocess_in_model,\n                start_z=self.start_z,\n                end_z=self.end_z,\n                shift_z=self.shift_z,\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            # worker_init_fn=lambda x: np.random.seed(np.random.get_state()[1][0] + x),\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\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-25T15:49:53.075701Z","iopub.execute_input":"2023-05-25T15:49:53.082138Z","iopub.status.idle":"2023-05-25T15:49:53.162087Z","shell.execute_reply.started":"2023-05-25T15:49:53.082072Z","shell.execute_reply":"2023-05-25T15:49:53.157707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define 2.5D_3DCNN Model (timm backbone)","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\n\nclass InkDetModel(nn.Module):\n    def __init__(\n        self,\n        model_name: str = \"resnet34\",\n        pretrained: bool = False,\n        drop_rate: float = 0,\n        drop_path_rate: float = 0,\n        num_3d_layer: int = 3,\n        in_chans: int = 7,\n        preprocess_in_model: bool = False,\n        start_z: int = 8,\n        end_z: int = -8,\n        shift_z: int = 2,\n        num_class: int = 1,\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_rate=drop_rate,\n            drop_path_rate=drop_path_rate,\n        )\n        self.output_fmt = getattr(self.encoder, \"output_fmt\", \"NHCW\")\n        self.in_chans = in_chans\n        num_features = self.encoder.feature_info.channels()[-1]\n        self.conv_proj = nn.Sequential(\n            nn.Conv2d(num_features, 512, 1, stride=1),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n        )\n        self.conv3d = nn.Sequential(\n            *[\n                ResidualConv3D(\n                    512,\n                    512,\n                    512,\n                    1,\n                )\n                for _ in range(num_3d_layer)\n            ]\n        )\n        self.head = nn.Sequential(\n            nn.Conv2d(\n                512 * 2,\n                512,\n                1,\n            ),\n            nn.ReLU(inplace=True),\n            nn.BatchNorm2d(512),\n            nn.Conv2d(\n                512,\n                num_class,\n                1,\n            ),\n        )\n        self.preprocess_in_model = preprocess_in_model\n        self.start_z = start_z\n        self.end_z = end_z\n        self.shift_z = shift_z\n\n    def preprocess(self, img):\n        if self.training and np.random.rand() < 0.5:\n            shift = np.random.randint(-self.shift_z, self.shift_z + 1)\n        else:\n            shift = 0\n        img = img[:, self.start_z + shift : self.end_z + shift]\n        return img\n\n    def forward_image_feats(self, img):\n        if self.preprocess_in_model:\n            img = self.preprocess(img)\n        mean = img.mean(dim=(1, 2, 3), keepdim=True)\n        std = img.std(dim=(1, 2, 3), keepdim=True) + 1e-6\n        img = (img - mean) / std\n        bs, ch, h, w = img.shape\n        assert ch % self.in_chans == 0\n        groups_3d = ch // self.in_chans\n        img = img.reshape((bs, groups_3d, self.in_chans, h, w))\n\n        if self.training:\n            ch_arr = list(range(img.shape[2]))\n            ch_arr = [\n                random.sample(ch_arr, len(ch_arr)) if np.random.rand() < 0.2 else ch_arr\n                for _ in range(img.shape[0])\n            ]\n            for i, ca in enumerate(ch_arr):\n                img[i] = img[i, :, ca]\n        img = img.reshape(bs * groups_3d, self.in_chans, h, w)\n        img_feat = self.encoder(img)[-1]\n        if self.output_fmt == \"NHWC\":\n            img_feat = img_feat.permute(0, 3, 1, 2).contiguous()\n        img_feat = self.conv_proj(img_feat)  # (bs * groups_3d, 512, h, w)\n        _, ch, h, w = img_feat.shape\n        img_feat = img_feat.reshape(bs, groups_3d, ch, h, w).transpose(\n            1, 2\n        )  # (bs, ch, groups_3d, h, w)\n        img_feat = self.conv3d(img_feat)  # (bs, ch, groups_3d, h, w)\n        img_feat = torch.cat([img_feat.mean(2), img_feat.max(2)[0]], 1)\n        return img_feat\n\n    def forward(\n        self,\n        img: torch.Tensor,\n    ):\n        \"\"\"\n        img: (bs, ch, h, w)\n        \"\"\"\n        img_feat = self.forward_image_feats(img)\n        return self.head(img_feat)\n\n\nclass InkDetLightningModel(pl.LightningModule):\n    def __init__(\n        self,\n        valid_fragment_id: str,\n        model_name: str = \"resnet34\",\n        pretrained: bool = False,\n        drop_rate: float = 0,\n        drop_path_rate: float = 0,\n        num_3d_layer: int = 3,\n        in_chans: int = 7,\n        preprocess_in_model: bool = False,\n        start_z: int = 8,\n        end_z: int = -8,\n        shift_z: int = 2,\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    ) -> None:\n        super().__init__()\n        self.__build_model(\n            model_name=model_name,\n            pretrained=pretrained,\n            drop_rate=drop_rate,\n            drop_path_rate=drop_path_rate,\n            num_3d_layer=num_3d_layer,\n            preprocess_in_model=preprocess_in_model,\n            start_z=start_z,\n            end_z=end_z,\n            shift_z=shift_z,\n            in_chans=in_chans,\n        )\n        self.save_hyperparameters()\n\n    def __build_model(\n        self,\n        model_name: str = \"resnet34\",\n        pretrained: bool = False,\n        drop_rate: float = 0,\n        drop_path_rate: float = 0,\n        num_3d_layer: int = 3,\n        in_chans: int = 7,\n        preprocess_in_model: bool = False,\n        start_z: int = 8,\n        end_z: int = -8,\n        shift_z: int = 2,\n    ):\n        self.model = InkDetModel(\n            model_name=model_name,\n            pretrained=pretrained,\n            drop_rate=drop_rate,\n            drop_path_rate=drop_path_rate,\n            num_3d_layer=num_3d_layer,\n            in_chans=in_chans,\n            preprocess_in_model=preprocess_in_model,\n            start_z=start_z,\n            end_z=end_z,\n            shift_z=shift_z,\n            num_class=1,\n        )\n        self.model_ema = ModelEmaV2(self.model, decay=0.99)\n        \n    def predict(self, volume):\n        output = torch.sigmoid(self.model_ema.module(volume))\n        output = F.interpolate(\n            output.float(),\n            scale_factor=32,\n            mode=\"bilinear\",\n            align_corners=True,\n        ).half()\n        return output","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-25T15:49:53.16402Z","iopub.execute_input":"2023-05-25T15:49:53.16475Z","iopub.status.idle":"2023-05-25T15:49:53.2544Z","shell.execute_reply.started":"2023-05-25T15:49:53.164697Z","shell.execute_reply":"2023-05-25T15:49:53.25267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define ResNet3DCSN","metadata":{}},{"cell_type":"code","source":"# Copyright (c) OpenMMLab. All rights reserved.\n\nimport logging\nimport warnings\n\nimport torch.utils.checkpoint as cp\nfrom mmcv.cnn import (\n    ConvModule,\n    NonLocal3d,\n    build_activation_layer,\n    constant_init,\n    kaiming_init,\n)\nfrom mmcv.runner import _load_checkpoint, load_checkpoint\nfrom timm.layers import DropPath\nfrom torch import nn\nfrom torch.nn.modules.batchnorm import _BatchNorm\nfrom torch.nn.modules.utils import _ntuple, _triple\n\n\nclass BasicBlock3d(nn.Module):\n    \"\"\"BasicBlock 3d block for ResNet3D.\n    Args:\n        inplanes (int): Number of channels for the input in first conv3d layer.\n        planes (int): Number of channels produced by some norm/conv3d layers.\n        spatial_stride (int): Spatial stride in the conv3d layer. Default: 1.\n        temporal_stride (int): Temporal stride in the conv3d layer. Default: 1.\n        dilation (int): Spacing between kernel elements. Default: 1.\n        downsample (nn.Module | None): Downsample layer. Default: None.\n        style (str): ``pytorch`` or ``caffe``. If set to \"pytorch\", the\n            stride-two layer is the 3x3 conv layer, otherwise the stride-two\n            layer is the first 1x1 conv layer. Default: 'pytorch'.\n        inflate (bool): Whether to inflate kernel. Default: True.\n        non_local (bool): Determine whether to apply non-local module in this\n            block. Default: False.\n        non_local_cfg (dict): Config for non-local module. Default: ``dict()``.\n        conv_cfg (dict): Config dict for convolution layer.\n            Default: ``dict(type='Conv3d')``.\n        norm_cfg (dict): Config for norm layers. required keys are ``type``,\n            Default: ``dict(type='BN3d')``.\n        act_cfg (dict): Config dict for activation layer.\n            Default: ``dict(type='ReLU')``.\n        with_cp (bool): Use checkpoint or not. Using checkpoint will save some\n            memory while slowing down the training speed. Default: False.\n    \"\"\"\n\n    expansion = 1\n\n    def __init__(\n        self,\n        inplanes,\n        planes,\n        spatial_stride=1,\n        temporal_stride=1,\n        dilation=1,\n        downsample=None,\n        style=\"pytorch\",\n        inflate=True,\n        non_local=False,\n        non_local_cfg=dict(),\n        conv_cfg=dict(type=\"Conv3d\"),\n        norm_cfg=dict(type=\"BN3d\"),\n        act_cfg=dict(type=\"ReLU\"),\n        with_cp=False,\n        drop_path=None,\n        **kwargs,\n    ):\n        super().__init__()\n        assert style in [\"pytorch\", \"caffe\"]\n        # make sure that only ``inflate_style`` is passed into kwargs\n        assert set(kwargs).issubset([\"inflate_style\"])\n\n        self.inplanes = inplanes\n        self.planes = planes\n        self.spatial_stride = spatial_stride\n        self.temporal_stride = temporal_stride\n        self.dilation = dilation\n        self.style = style\n        self.inflate = inflate\n        self.conv_cfg = conv_cfg\n        self.norm_cfg = norm_cfg\n        self.act_cfg = act_cfg\n        self.with_cp = with_cp\n        self.non_local = non_local\n        self.non_local_cfg = non_local_cfg\n\n        self.conv1_stride_s = spatial_stride\n        self.conv2_stride_s = 1\n        self.conv1_stride_t = temporal_stride\n        self.conv2_stride_t = 1\n\n        if self.inflate:\n            conv1_kernel_size = (3, 3, 3)\n            conv1_padding = (1, dilation, dilation)\n            conv2_kernel_size = (3, 3, 3)\n            conv2_padding = (1, 1, 1)\n        else:\n            conv1_kernel_size = (1, 3, 3)\n            conv1_padding = (0, dilation, dilation)\n            conv2_kernel_size = (1, 3, 3)\n            conv2_padding = (0, 1, 1)\n\n        self.conv1 = ConvModule(\n            inplanes,\n            planes,\n            conv1_kernel_size,\n            stride=(self.conv1_stride_t, self.conv1_stride_s, self.conv1_stride_s),\n            padding=conv1_padding,\n            dilation=(1, dilation, dilation),\n            bias=False,\n            conv_cfg=self.conv_cfg,\n            norm_cfg=self.norm_cfg,\n            act_cfg=self.act_cfg,\n        )\n\n        self.conv2 = ConvModule(\n            planes,\n            planes * self.expansion,\n            conv2_kernel_size,\n            stride=(self.conv2_stride_t, self.conv2_stride_s, self.conv2_stride_s),\n            padding=conv2_padding,\n            bias=False,\n            conv_cfg=self.conv_cfg,\n            norm_cfg=self.norm_cfg,\n            act_cfg=None,\n        )\n\n        self.downsample = downsample\n        self.relu = build_activation_layer(self.act_cfg)\n\n        if self.non_local:\n            self.non_local_block = NonLocal3d(\n                self.conv2.norm.num_features, **self.non_local_cfg\n            )\n        self.drop_path = drop_path\n\n    def forward(self, x):\n        \"\"\"Defines the computation performed at every call.\"\"\"\n\n        def _inner_forward(x):\n            \"\"\"Forward wrapper for utilizing checkpoint.\"\"\"\n            identity = x\n\n            out = self.conv1(x)\n            out = self.conv2(out)\n            if self.drop_path is not None:\n                x = self.drop_path(x)\n            if self.downsample is not None:\n                identity = self.downsample(x)\n\n            out = out + identity\n            return out\n\n        if self.with_cp and x.requires_grad:\n            out = cp.checkpoint(_inner_forward, x)\n        else:\n            out = _inner_forward(x)\n        out = self.relu(out)\n\n        if self.non_local:\n            out = self.non_local_block(out)\n\n        return out\n\n\nclass Bottleneck3d(nn.Module):\n    \"\"\"Bottleneck 3d block for ResNet3D.\n    Args:\n        inplanes (int): Number of channels for the input in first conv3d layer.\n        planes (int): Number of channels produced by some norm/conv3d layers.\n        spatial_stride (int): Spatial stride in the conv3d layer. Default: 1.\n        temporal_stride (int): Temporal stride in the conv3d layer. Default: 1.\n        dilation (int): Spacing between kernel elements. Default: 1.\n        downsample (nn.Module | None): Downsample layer. Default: None.\n        style (str): ``pytorch`` or ``caffe``. If set to \"pytorch\", the\n            stride-two layer is the 3x3 conv layer, otherwise the stride-two\n            layer is the first 1x1 conv layer. Default: 'pytorch'.\n        inflate (bool): Whether to inflate kernel. Default: True.\n        inflate_style (str): ``3x1x1`` or ``3x3x3``. which determines the\n            kernel sizes and padding strides for conv1 and conv2 in each block.\n            Default: '3x1x1'.\n        non_local (bool): Determine whether to apply non-local module in this\n            block. Default: False.\n        non_local_cfg (dict): Config for non-local module. Default: ``dict()``.\n        conv_cfg (dict): Config dict for convolution layer.\n            Default: ``dict(type='Conv3d')``.\n        norm_cfg (dict): Config for norm layers. required keys are ``type``,\n            Default: ``dict(type='BN3d')``.\n        act_cfg (dict): Config dict for activation layer.\n            Default: ``dict(type='ReLU')``.\n        with_cp (bool): Use checkpoint or not. Using checkpoint will save some\n            memory while slowing down the training speed. Default: False.\n    \"\"\"\n\n    expansion = 4\n\n    def __init__(\n        self,\n        inplanes,\n        planes,\n        spatial_stride=1,\n        temporal_stride=1,\n        dilation=1,\n        downsample=None,\n        style=\"pytorch\",\n        inflate=True,\n        inflate_style=\"3x1x1\",\n        non_local=False,\n        non_local_cfg=dict(),\n        conv_cfg=dict(type=\"Conv3d\"),\n        norm_cfg=dict(type=\"BN3d\"),\n        act_cfg=dict(type=\"ReLU\"),\n        with_cp=False,\n        drop_path=None,\n    ):\n        super().__init__()\n        assert style in [\"pytorch\", \"caffe\"]\n        assert inflate_style in [\"3x1x1\", \"3x3x3\"]\n\n        self.inplanes = inplanes\n        self.planes = planes\n        self.spatial_stride = spatial_stride\n        self.temporal_stride = temporal_stride\n        self.dilation = dilation\n        self.style = style\n        self.inflate = inflate\n        self.inflate_style = inflate_style\n        self.norm_cfg = norm_cfg\n        self.conv_cfg = conv_cfg\n        self.act_cfg = act_cfg\n        self.with_cp = with_cp\n        self.non_local = non_local\n        self.non_local_cfg = non_local_cfg\n\n        if self.style == \"pytorch\":\n            self.conv1_stride_s = 1\n            self.conv2_stride_s = spatial_stride\n            self.conv1_stride_t = 1\n            self.conv2_stride_t = temporal_stride\n        else:\n            self.conv1_stride_s = spatial_stride\n            self.conv2_stride_s = 1\n            self.conv1_stride_t = temporal_stride\n            self.conv2_stride_t = 1\n\n        if self.inflate:\n            if inflate_style == \"3x1x1\":\n                conv1_kernel_size = (3, 1, 1)\n                conv1_padding = (1, 0, 0)\n                conv2_kernel_size = (1, 3, 3)\n                conv2_padding = (0, dilation, dilation)\n            else:\n                conv1_kernel_size = (1, 1, 1)\n                conv1_padding = (0, 0, 0)\n                conv2_kernel_size = (3, 3, 3)\n                conv2_padding = (1, dilation, dilation)\n        else:\n            conv1_kernel_size = (1, 1, 1)\n            conv1_padding = (0, 0, 0)\n            conv2_kernel_size = (1, 3, 3)\n            conv2_padding = (0, dilation, dilation)\n\n        self.conv1 = ConvModule(\n            inplanes,\n            planes,\n            conv1_kernel_size,\n            stride=(self.conv1_stride_t, self.conv1_stride_s, self.conv1_stride_s),\n            padding=conv1_padding,\n            bias=False,\n            conv_cfg=self.conv_cfg,\n            norm_cfg=self.norm_cfg,\n            act_cfg=self.act_cfg,\n        )\n\n        self.conv2 = ConvModule(\n            planes,\n            planes,\n            conv2_kernel_size,\n            stride=(self.conv2_stride_t, self.conv2_stride_s, self.conv2_stride_s),\n            padding=conv2_padding,\n            dilation=(1, dilation, dilation),\n            bias=False,\n            conv_cfg=self.conv_cfg,\n            norm_cfg=self.norm_cfg,\n            act_cfg=self.act_cfg,\n        )\n\n        self.conv3 = ConvModule(\n            planes,\n            planes * self.expansion,\n            1,\n            bias=False,\n            conv_cfg=self.conv_cfg,\n            norm_cfg=self.norm_cfg,\n            # No activation in the third ConvModule for bottleneck\n            act_cfg=None,\n        )\n\n        self.downsample = downsample\n        self.relu = build_activation_layer(self.act_cfg)\n\n        if self.non_local:\n            self.non_local_block = NonLocal3d(\n                self.conv3.norm.num_features, **self.non_local_cfg\n            )\n        self.drop_path = drop_path\n\n    def forward(self, x):\n        \"\"\"Defines the computation performed at every call.\"\"\"\n\n        def _inner_forward(x):\n            \"\"\"Forward wrapper for utilizing checkpoint.\"\"\"\n            identity = x\n\n            out = self.conv1(x)\n            out = self.conv2(out)\n            out = self.conv3(out)\n\n            if self.drop_path is not None:\n                x = self.drop_path(x)\n\n            if self.downsample is not None:\n                identity = self.downsample(x)\n\n            out = out + identity\n            return out\n\n        if self.with_cp and x.requires_grad:\n            out = cp.checkpoint(_inner_forward, x)\n        else:\n            out = _inner_forward(x)\n        out = self.relu(out)\n\n        if self.non_local:\n            out = self.non_local_block(out)\n\n        return out\n\n\nclass ResNet3d(nn.Module):\n    \"\"\"ResNet 3d backbone.\n    Args:\n        depth (int): Depth of resnet, from {18, 34, 50, 101, 152}.\n        pretrained (str | None): Name of pretrained model.\n        stage_blocks (tuple | None): Set number of stages for each res layer.\n            Default: None.\n        pretrained2d (bool): Whether to load pretrained 2D model.\n            Default: True.\n        in_channels (int): Channel num of input features. Default: 3.\n        base_channels (int): Channel num of stem output features. Default: 64.\n        out_indices (Sequence[int]): Indices of output feature. Default: (3, ).\n        num_stages (int): Resnet stages. Default: 4.\n        spatial_strides (Sequence[int]):\n            Spatial strides of residual blocks of each stage.\n            Default: ``(1, 2, 2, 2)``.\n        temporal_strides (Sequence[int]):\n            Temporal strides of residual blocks of each stage.\n            Default: ``(1, 1, 1, 1)``.\n        dilations (Sequence[int]): Dilation of each stage.\n            Default: ``(1, 1, 1, 1)``.\n        conv1_kernel (Sequence[int]): Kernel size of the first conv layer.\n            Default: ``(3, 7, 7)``.\n        conv1_stride_s (int): Spatial stride of the first conv layer.\n            Default: 2.\n        conv1_stride_t (int): Temporal stride of the first conv layer.\n            Default: 1.\n        pool1_stride_s (int): Spatial stride of the first pooling layer.\n            Default: 2.\n        pool1_stride_t (int): Temporal stride of the first pooling layer.\n            Default: 1.\n        with_pool2 (bool): Whether to use pool2. Default: True.\n        style (str): `pytorch` or `caffe`. If set to \"pytorch\", the stride-two\n            layer is the 3x3 conv layer, otherwise the stride-two layer is\n            the first 1x1 conv layer. Default: 'pytorch'.\n        frozen_stages (int): Stages to be frozen (all param fixed). -1 means\n            not freezing any parameters. Default: -1.\n        inflate (Sequence[int]): Inflate Dims of each block.\n            Default: (1, 1, 1, 1).\n        inflate_style (str): ``3x1x1`` or ``3x3x3``. which determines the\n            kernel sizes and padding strides for conv1 and conv2 in each block.\n            Default: '3x1x1'.\n        conv_cfg (dict): Config for conv layers. required keys are ``type``\n            Default: ``dict(type='Conv3d')``.\n        norm_cfg (dict): Config for norm layers. required keys are ``type`` and\n            ``requires_grad``.\n            Default: ``dict(type='BN3d', requires_grad=True)``.\n        act_cfg (dict): Config dict for activation layer.\n            Default: ``dict(type='ReLU', inplace=True)``.\n        norm_eval (bool): Whether to set BN layers to eval mode, namely, freeze\n            running stats (mean and var). Default: False.\n        with_cp (bool): Use checkpoint or not. Using checkpoint will save some\n            memory while slowing down the training speed. Default: False.\n        non_local (Sequence[int]): Determine whether to apply non-local module\n            in the corresponding block of each stages. Default: (0, 0, 0, 0).\n        non_local_cfg (dict): Config for non-local module. Default: ``dict()``.\n        zero_init_residual (bool):\n            Whether to use zero initialization for residual block,\n            Default: True.\n        kwargs (dict, optional): Key arguments for \"make_res_layer\".\n    \"\"\"\n\n    arch_settings = {\n        18: (BasicBlock3d, (2, 2, 2, 2)),\n        34: (BasicBlock3d, (3, 4, 6, 3)),\n        50: (Bottleneck3d, (3, 4, 6, 3)),\n        101: (Bottleneck3d, (3, 4, 23, 3)),\n        152: (Bottleneck3d, (3, 8, 36, 3)),\n    }\n\n    def __init__(\n        self,\n        depth,\n        pretrained,\n        stage_blocks=None,\n        pretrained2d=True,\n        in_channels=3,\n        num_stages=4,\n        base_channels=64,\n        out_indices=(\n            0,\n            1,\n            2,\n            3,\n        ),\n        spatial_strides=(1, 2, 2, 2),\n        temporal_strides=(1, 1, 1, 1),\n        dilations=(1, 1, 1, 1),\n        conv1_kernel=(3, 7, 7),\n        conv1_stride_s=2,\n        conv1_stride_t=1,\n        pool1_stride_s=2,\n        pool1_stride_t=1,\n        with_pool1=True,\n        with_pool2=True,\n        style=\"pytorch\",\n        frozen_stages=-1,\n        inflate=(1, 1, 1, 1),\n        inflate_style=\"3x1x1\",\n        conv_cfg=dict(type=\"Conv3d\"),\n        norm_cfg=dict(type=\"BN3d\", requires_grad=True),\n        act_cfg=dict(type=\"ReLU\", inplace=True),\n        norm_eval=False,\n        with_cp=False,\n        non_local=(0, 0, 0, 0),\n        non_local_cfg=dict(),\n        zero_init_residual=True,\n        drop_path_rate=0.0,\n        **kwargs,\n    ):\n        super().__init__()\n        if depth not in self.arch_settings:\n            raise KeyError(f\"invalid depth {depth} for resnet\")\n        self.depth = depth\n        self.pretrained = pretrained\n        self.pretrained2d = pretrained2d\n        self.in_channels = in_channels\n        self.base_channels = base_channels\n        self.num_stages = num_stages\n        assert 1 <= num_stages <= 4\n        self.stage_blocks = stage_blocks\n        self.out_indices = out_indices\n        assert max(out_indices) < num_stages\n        self.spatial_strides = spatial_strides\n        self.temporal_strides = temporal_strides\n        self.dilations = dilations\n        assert (\n            len(spatial_strides)\n            == len(temporal_strides)\n            == len(dilations)\n            == num_stages\n        )\n        if self.stage_blocks is not None:\n            assert len(self.stage_blocks) == num_stages\n\n        self.conv1_kernel = conv1_kernel\n        self.conv1_stride_s = conv1_stride_s\n        self.conv1_stride_t = conv1_stride_t\n        self.pool1_stride_s = pool1_stride_s\n        self.pool1_stride_t = pool1_stride_t\n        self.with_pool1 = with_pool1\n        self.with_pool2 = with_pool2\n        self.style = style\n        self.frozen_stages = frozen_stages\n        self.stage_inflations = _ntuple(num_stages)(inflate)\n        self.non_local_stages = _ntuple(num_stages)(non_local)\n        self.inflate_style = inflate_style\n        self.conv_cfg = conv_cfg\n        self.norm_cfg = norm_cfg\n        self.act_cfg = act_cfg\n        self.norm_eval = norm_eval\n        self.with_cp = with_cp\n        self.zero_init_residual = zero_init_residual\n\n        self.block, stage_blocks = self.arch_settings[depth]\n\n        if self.stage_blocks is None:\n            self.stage_blocks = stage_blocks[:num_stages]\n\n        self.inplanes = self.base_channels\n\n        self.non_local_cfg = non_local_cfg\n\n        self._make_stem_layer()\n\n        self.res_layers = []\n        for i, num_blocks in enumerate(self.stage_blocks):\n            spatial_stride = spatial_strides[i]\n            temporal_stride = temporal_strides[i]\n            dilation = dilations[i]\n            planes = self.base_channels * 2**i\n            res_layer = self.make_res_layer(\n                self.block,\n                self.inplanes,\n                planes,\n                num_blocks,\n                spatial_stride=spatial_stride,\n                temporal_stride=temporal_stride,\n                dilation=dilation,\n                style=self.style,\n                norm_cfg=self.norm_cfg,\n                conv_cfg=self.conv_cfg,\n                act_cfg=self.act_cfg,\n                non_local=self.non_local_stages[i],\n                non_local_cfg=self.non_local_cfg,\n                inflate=self.stage_inflations[i],\n                inflate_style=self.inflate_style,\n                with_cp=with_cp,\n                drop_path_rate=drop_path_rate,\n                **kwargs,\n            )\n            self.inplanes = planes * self.block.expansion\n            layer_name = f\"layer{i + 1}\"\n            self.add_module(layer_name, res_layer)\n            self.res_layers.append(layer_name)\n\n        self.feat_dim = (\n            self.block.expansion\n            * self.base_channels\n            * 2 ** (len(self.stage_blocks) - 1)\n        )\n\n    @staticmethod\n    def make_res_layer(\n        block,\n        inplanes,\n        planes,\n        blocks,\n        spatial_stride=1,\n        temporal_stride=1,\n        dilation=1,\n        style=\"pytorch\",\n        inflate=1,\n        inflate_style=\"3x1x1\",\n        non_local=0,\n        non_local_cfg=dict(),\n        norm_cfg=None,\n        act_cfg=None,\n        conv_cfg=None,\n        with_cp=False,\n        drop_path_rate=0.0,\n        **kwargs,\n    ):\n        \"\"\"Build residual layer for ResNet3D.\n        Args:\n            block (nn.Module): Residual module to be built.\n            inplanes (int): Number of channels for the input feature\n                in each block.\n            planes (int): Number of channels for the output feature\n                in each block.\n            blocks (int): Number of residual blocks.\n            spatial_stride (int | Sequence[int]): Spatial strides in\n                residual and conv layers. Default: 1.\n            temporal_stride (int | Sequence[int]): Temporal strides in\n                residual and conv layers. Default: 1.\n            dilation (int): Spacing between kernel elements. Default: 1.\n            style (str): ``pytorch`` or ``caffe``. If set to ``pytorch``,\n                the stride-two layer is the 3x3 conv layer, otherwise\n                the stride-two layer is the first 1x1 conv layer.\n                Default: ``pytorch``.\n            inflate (int | Sequence[int]): Determine whether to inflate\n                for each block. Default: 1.\n            inflate_style (str): ``3x1x1`` or ``3x3x3``. which determines\n                the kernel sizes and padding strides for conv1 and conv2\n                in each block. Default: '3x1x1'.\n            non_local (int | Sequence[int]): Determine whether to apply\n                non-local module in the corresponding block of each stages.\n                Default: 0.\n            non_local_cfg (dict): Config for non-local module.\n                Default: ``dict()``.\n            conv_cfg (dict | None): Config for norm layers. Default: None.\n            norm_cfg (dict | None): Config for norm layers. Default: None.\n            act_cfg (dict | None): Config for activate layers. Default: None.\n            with_cp (bool | None): Use checkpoint or not. Using checkpoint\n                will save some memory while slowing down the training speed.\n                Default: False.\n        Returns:\n            nn.Module: A residual layer for the given config.\n        \"\"\"\n        inflate = inflate if not isinstance(inflate, int) else (inflate,) * blocks\n        non_local = (\n            non_local if not isinstance(non_local, int) else (non_local,) * blocks\n        )\n        assert len(inflate) == blocks and len(non_local) == blocks\n        downsample = None\n        if spatial_stride != 1 or inplanes != planes * block.expansion:\n            downsample = ConvModule(\n                inplanes,\n                planes * block.expansion,\n                kernel_size=1,\n                stride=(temporal_stride, spatial_stride, spatial_stride),\n                bias=False,\n                conv_cfg=conv_cfg,\n                norm_cfg=norm_cfg,\n                act_cfg=None,\n            )\n\n        layers = []\n        layers.append(\n            block(\n                inplanes,\n                planes,\n                spatial_stride=spatial_stride,\n                temporal_stride=temporal_stride,\n                dilation=dilation,\n                downsample=downsample,\n                style=style,\n                inflate=(inflate[0] == 1),\n                inflate_style=inflate_style,\n                non_local=(non_local[0] == 1),\n                non_local_cfg=non_local_cfg,\n                norm_cfg=norm_cfg,\n                conv_cfg=conv_cfg,\n                act_cfg=act_cfg,\n                with_cp=with_cp,\n                **kwargs,\n            )\n        )\n        inplanes = planes * block.expansion\n        for i in range(1, blocks):\n            block_dpr = drop_path_rate * i / (blocks - 1)\n            layers.append(\n                block(\n                    inplanes,\n                    planes,\n                    spatial_stride=1,\n                    temporal_stride=1,\n                    dilation=dilation,\n                    style=style,\n                    inflate=(inflate[i] == 1),\n                    inflate_style=inflate_style,\n                    non_local=(non_local[i] == 1),\n                    non_local_cfg=non_local_cfg,\n                    norm_cfg=norm_cfg,\n                    conv_cfg=conv_cfg,\n                    act_cfg=act_cfg,\n                    with_cp=with_cp,\n                    drop_path=DropPath(block_dpr) if block_dpr > 0.0 else None,\n                    **kwargs,\n                )\n            )\n\n        return nn.Sequential(*layers)\n\n    @staticmethod\n    def _inflate_conv_params(\n        conv3d, state_dict_2d, module_name_2d, inflated_param_names\n    ):\n        \"\"\"Inflate a conv module from 2d to 3d.\n        Args:\n            conv3d (nn.Module): The destination conv3d module.\n            state_dict_2d (OrderedDict): The state dict of pretrained 2d model.\n            module_name_2d (str): The name of corresponding conv module in the\n                2d model.\n            inflated_param_names (list[str]): List of parameters that have been\n                inflated.\n        \"\"\"\n        weight_2d_name = module_name_2d + \".weight\"\n\n        conv2d_weight = state_dict_2d[weight_2d_name]\n        kernel_t = conv3d.weight.data.shape[2]\n\n        new_weight = conv2d_weight.data.unsqueeze(2).expand_as(conv3d.weight) / kernel_t\n        conv3d.weight.data.copy_(new_weight)\n        inflated_param_names.append(weight_2d_name)\n\n        if getattr(conv3d, \"bias\") is not None:\n            bias_2d_name = module_name_2d + \".bias\"\n            conv3d.bias.data.copy_(state_dict_2d[bias_2d_name])\n            inflated_param_names.append(bias_2d_name)\n\n    @staticmethod\n    def _inflate_bn_params(bn3d, state_dict_2d, module_name_2d, inflated_param_names):\n        \"\"\"Inflate a norm module from 2d to 3d.\n        Args:\n            bn3d (nn.Module): The destination bn3d module.\n            state_dict_2d (OrderedDict): The state dict of pretrained 2d model.\n            module_name_2d (str): The name of corresponding bn module in the\n                2d model.\n            inflated_param_names (list[str]): List of parameters that have been\n                inflated.\n        \"\"\"\n        for param_name, param in bn3d.named_parameters():\n            param_2d_name = f\"{module_name_2d}.{param_name}\"\n            param_2d = state_dict_2d[param_2d_name]\n            if param.data.shape != param_2d.shape:\n                warnings.warn(\n                    f\"The parameter of {module_name_2d} is not\"\n                    \"loaded due to incompatible shapes. \"\n                )\n                return\n\n            param.data.copy_(param_2d)\n            inflated_param_names.append(param_2d_name)\n\n        for param_name, param in bn3d.named_buffers():\n            param_2d_name = f\"{module_name_2d}.{param_name}\"\n            # some buffers like num_batches_tracked may not exist in old\n            # checkpoints\n            if param_2d_name in state_dict_2d:\n                param_2d = state_dict_2d[param_2d_name]\n                param.data.copy_(param_2d)\n                inflated_param_names.append(param_2d_name)\n\n    @staticmethod\n    def _inflate_weights(self, logger):\n        \"\"\"Inflate the resnet2d parameters to resnet3d.\n        The differences between resnet3d and resnet2d mainly lie in an extra\n        axis of conv kernel. To utilize the pretrained parameters in 2d model,\n        the weight of conv2d models should be inflated to fit in the shapes of\n        the 3d counterpart.\n        Args:\n            logger (logging.Logger): The logger used to print\n                debugging infomation.\n        \"\"\"\n\n        state_dict_r2d = _load_checkpoint(self.pretrained)\n        if \"state_dict\" in state_dict_r2d:\n            state_dict_r2d = state_dict_r2d[\"state_dict\"]\n\n        inflated_param_names = []\n        for name, module in self.named_modules():\n            if isinstance(module, ConvModule):\n                # we use a ConvModule to wrap conv+bn+relu layers, thus the\n                # name mapping is needed\n                if \"downsample\" in name:\n                    # layer{X}.{Y}.downsample.conv->layer{X}.{Y}.downsample.0\n                    original_conv_name = name + \".0\"\n                    # layer{X}.{Y}.downsample.bn->layer{X}.{Y}.downsample.1\n                    original_bn_name = name + \".1\"\n                else:\n                    # layer{X}.{Y}.conv{n}.conv->layer{X}.{Y}.conv{n}\n                    original_conv_name = name\n                    # layer{X}.{Y}.conv{n}.bn->layer{X}.{Y}.bn{n}\n                    original_bn_name = name.replace(\"conv\", \"bn\")\n                if original_conv_name + \".weight\" not in state_dict_r2d:\n                    logger.warning(\n                        f\"Module not exist in the state_dict_r2d\"\n                        f\": {original_conv_name}\"\n                    )\n                else:\n                    shape_2d = state_dict_r2d[original_conv_name + \".weight\"].shape\n                    shape_3d = module.conv.weight.data.shape\n                    if shape_2d != shape_3d[:2] + shape_3d[3:]:\n                        logger.warning(\n                            f\"Weight shape mismatch for \"\n                            f\": {original_conv_name} : \"\n                            f\"3d weight shape: {shape_3d}; \"\n                            f\"2d weight shape: {shape_2d}. \"\n                        )\n                    else:\n                        self._inflate_conv_params(\n                            module.conv,\n                            state_dict_r2d,\n                            original_conv_name,\n                            inflated_param_names,\n                        )\n\n                if original_bn_name + \".weight\" not in state_dict_r2d:\n                    logger.warning(\n                        f\"Module not exist in the state_dict_r2d\"\n                        f\": {original_bn_name}\"\n                    )\n                else:\n                    self._inflate_bn_params(\n                        module.bn,\n                        state_dict_r2d,\n                        original_bn_name,\n                        inflated_param_names,\n                    )\n\n        # check if any parameters in the 2d checkpoint are not loaded\n        remaining_names = set(state_dict_r2d.keys()) - set(inflated_param_names)\n        if remaining_names:\n            logger.info(\n                f\"These parameters in the 2d checkpoint are not loaded\"\n                f\": {remaining_names}\"\n            )\n\n    def inflate_weights(self, logger):\n        self._inflate_weights(self, logger)\n\n    def _make_stem_layer(self):\n        \"\"\"Construct the stem layers consists of a conv+norm+act module and a\n        pooling layer.\"\"\"\n        self.conv1 = ConvModule(\n            self.in_channels,\n            self.base_channels,\n            kernel_size=self.conv1_kernel,\n            stride=(self.conv1_stride_t, self.conv1_stride_s, self.conv1_stride_s),\n            padding=tuple([(k - 1) // 2 for k in _triple(self.conv1_kernel)]),\n            bias=False,\n            conv_cfg=self.conv_cfg,\n            norm_cfg=self.norm_cfg,\n            act_cfg=self.act_cfg,\n        )\n\n        self.maxpool = nn.MaxPool3d(\n            kernel_size=(1, 3, 3),\n            stride=(self.pool1_stride_t, self.pool1_stride_s, self.pool1_stride_s),\n            padding=(0, 1, 1),\n        )\n\n        self.pool2 = nn.MaxPool3d(kernel_size=(2, 1, 1), stride=(2, 1, 1))\n\n    def _freeze_stages(self):\n        \"\"\"Prevent all the parameters from being optimized before\n        ``self.frozen_stages``.\"\"\"\n        if self.frozen_stages >= 0:\n            self.conv1.eval()\n            for param in self.conv1.parameters():\n                param.requires_grad = False\n\n        for i in range(1, self.frozen_stages + 1):\n            m = getattr(self, f\"layer{i}\")\n            m.eval()\n            for param in m.parameters():\n                param.requires_grad = False\n\n    @staticmethod\n    def _init_weights(self, pretrained=None):\n        \"\"\"Initiate the parameters either from existing checkpoint or from\n        scratch.\n        Args:\n            pretrained (str | None): The path of the pretrained weight. Will\n                override the original `pretrained` if set. The arg is added to\n                be compatible with mmdet. Default: None.\n        \"\"\"\n        if pretrained:\n            self.pretrained = pretrained\n        if isinstance(self.pretrained, str):\n            logger = logging.Logger()\n            logger.info(f\"load model from: {self.pretrained}\")\n\n            if self.pretrained2d:\n                # Inflate 2D model into 3D model.\n                self.inflate_weights(logger)\n\n            else:\n                # Directly load 3D model.\n                load_checkpoint(self, self.pretrained, strict=False, logger=logger)\n\n        elif self.pretrained is None:\n            for m in self.modules():\n                if isinstance(m, nn.Conv3d):\n                    kaiming_init(m)\n                elif isinstance(m, _BatchNorm):\n                    constant_init(m, 1)\n\n            if self.zero_init_residual:\n                for m in self.modules():\n                    if isinstance(m, Bottleneck3d):\n                        constant_init(m.conv3.bn, 0)\n                    elif isinstance(m, BasicBlock3d):\n                        constant_init(m.conv2.bn, 0)\n        else:\n            raise TypeError(\"pretrained must be a str or None\")\n\n    def init_weights(self, pretrained=None):\n        self._init_weights(self, pretrained)\n\n    def forward(self, x):\n        \"\"\"Defines the computation performed at every call.\n        Args:\n            x (torch.Tensor): The input data.\n        Returns:\n            torch.Tensor: The feature of the input\n            samples extracted by the backbone.\n        \"\"\"\n        x = self.conv1(x)\n        outs = [x]\n\n        if self.with_pool1:\n            x = self.maxpool(x)\n        for i, layer_name in enumerate(self.res_layers):\n            res_layer = getattr(self, layer_name)\n            x = res_layer(x)\n            # print(i, x.shape)\n            if i == 0 and self.with_pool2:\n                x = self.pool2(x)\n            if i in self.out_indices:\n                outs.append(x)\n        if len(outs) == 1:\n            return outs[0]\n\n        return tuple(outs)\n\n    def train(self, mode=True):\n        \"\"\"Set the optimization status when training.\"\"\"\n        super().train(mode)\n        self._freeze_stages()\n        if mode and self.norm_eval:\n            for m in self.modules():\n                if isinstance(m, _BatchNorm):\n                    m.eval()\n\n\nclass ResNet3dLayer(nn.Module):\n    \"\"\"ResNet 3d Layer.\n    Args:\n        depth (int): Depth of resnet, from {18, 34, 50, 101, 152}.\n        pretrained (str | None): Name of pretrained model.\n        pretrained2d (bool): Whether to load pretrained 2D model.\n            Default: True.\n        stage (int): The index of Resnet stage. Default: 3.\n        base_channels (int): Channel num of stem output features. Default: 64.\n        spatial_stride (int): The 1st res block's spatial stride. Default 2.\n        temporal_stride (int): The 1st res block's temporal stride. Default 1.\n        dilation (int): The dilation. Default: 1.\n        style (str): `pytorch` or `caffe`. If set to \"pytorch\", the stride-two\n            layer is the 3x3 conv layer, otherwise the stride-two layer is\n            the first 1x1 conv layer. Default: 'pytorch'.\n        all_frozen (bool): Frozen all modules in the layer. Default: False.\n        inflate (int): Inflate Dims of each block. Default: 1.\n        inflate_style (str): ``3x1x1`` or ``3x3x3``. which determines the\n            kernel sizes and padding strides for conv1 and conv2 in each block.\n            Default: '3x1x1'.\n        conv_cfg (dict): Config for conv layers. required keys are ``type``\n            Default: ``dict(type='Conv3d')``.\n        norm_cfg (dict): Config for norm layers. required keys are ``type`` and\n            ``requires_grad``.\n            Default: ``dict(type='BN3d', requires_grad=True)``.\n        act_cfg (dict): Config dict for activation layer.\n            Default: ``dict(type='ReLU', inplace=True)``.\n        norm_eval (bool): Whether to set BN layers to eval mode, namely, freeze\n            running stats (mean and var). Default: False.\n        with_cp (bool): Use checkpoint or not. Using checkpoint will save some\n            memory while slowing down the training speed. Default: False.\n        zero_init_residual (bool):\n            Whether to use zero initialization for residual block,\n            Default: True.\n        kwargs (dict, optional): Key arguments for \"make_res_layer\".\n    \"\"\"\n\n    def __init__(\n        self,\n        depth,\n        pretrained,\n        pretrained2d=True,\n        stage=3,\n        base_channels=64,\n        spatial_stride=2,\n        temporal_stride=1,\n        dilation=1,\n        style=\"pytorch\",\n        all_frozen=False,\n        inflate=1,\n        inflate_style=\"3x1x1\",\n        conv_cfg=dict(type=\"Conv3d\"),\n        norm_cfg=dict(type=\"BN3d\", requires_grad=True),\n        act_cfg=dict(type=\"ReLU\", inplace=True),\n        norm_eval=False,\n        with_cp=False,\n        zero_init_residual=True,\n        **kwargs,\n    ):\n        super().__init__()\n        self.arch_settings = ResNet3d.arch_settings\n        assert depth in self.arch_settings\n\n        self.make_res_layer = ResNet3d.make_res_layer\n        self._inflate_conv_params = ResNet3d._inflate_conv_params\n        self._inflate_bn_params = ResNet3d._inflate_bn_params\n        self._inflate_weights = ResNet3d._inflate_weights\n        self._init_weights = ResNet3d._init_weights\n\n        self.depth = depth\n        self.pretrained = pretrained\n        self.pretrained2d = pretrained2d\n        self.stage = stage\n        # stage index is 0 based\n        assert 0 <= stage <= 3\n        self.base_channels = base_channels\n\n        self.spatial_stride = spatial_stride\n        self.temporal_stride = temporal_stride\n        self.dilation = dilation\n\n        self.style = style\n        self.all_frozen = all_frozen\n\n        self.stage_inflation = inflate\n        self.inflate_style = inflate_style\n        self.conv_cfg = conv_cfg\n        self.norm_cfg = norm_cfg\n        self.act_cfg = act_cfg\n        self.norm_eval = norm_eval\n        self.with_cp = with_cp\n        self.zero_init_residual = zero_init_residual\n\n        block, stage_blocks = self.arch_settings[depth]\n        stage_block = stage_blocks[stage]\n        planes = 64 * 2**stage\n        inplanes = 64 * 2 ** (stage - 1) * block.expansion\n\n        res_layer = self.make_res_layer(\n            block,\n            inplanes,\n            planes,\n            stage_block,\n            spatial_stride=spatial_stride,\n            temporal_stride=temporal_stride,\n            dilation=dilation,\n            style=self.style,\n            norm_cfg=self.norm_cfg,\n            conv_cfg=self.conv_cfg,\n            act_cfg=self.act_cfg,\n            inflate=self.stage_inflation,\n            inflate_style=self.inflate_style,\n            with_cp=with_cp,\n            **kwargs,\n        )\n\n        self.layer_name = f\"layer{stage + 1}\"\n        self.add_module(self.layer_name, res_layer)\n\n    def inflate_weights(self, logger):\n        self._inflate_weights(self, logger)\n\n    def _freeze_stages(self):\n        \"\"\"Prevent all the parameters from being optimized before\n        ``self.frozen_stages``.\"\"\"\n        if self.all_frozen:\n            layer = getattr(self, self.layer_name)\n            layer.eval()\n            for param in layer.parameters():\n                param.requires_grad = False\n\n    def init_weights(self, pretrained=None):\n        self._init_weights(self, pretrained)\n\n    def forward(self, x):\n        \"\"\"Defines the computation performed at every call.\n        Args:\n            x (torch.Tensor): The input data.\n        Returns:\n            torch.Tensor: The feature of the input\n            samples extracted by the backbone.\n        \"\"\"\n        res_layer = getattr(self, self.layer_name)\n        out = res_layer(x)\n        return out\n\n    def train(self, mode=True):\n        \"\"\"Set the optimization status when training.\"\"\"\n        super().train(mode)\n        self._freeze_stages()\n        if mode and self.norm_eval:\n            for m in self.modules():\n                if isinstance(m, _BatchNorm):\n                    m.eval()\n\n\nclass CSNBottleneck3d(Bottleneck3d):\n    \"\"\"Channel-Separated Bottleneck Block.\n    This module is proposed in\n    \"Video Classification with Channel-Separated Convolutional Networks\"\n    Link: https://arxiv.org/pdf/1711.11248.pdf\n    Args:\n        inplanes (int): Number of channels for the input in first conv3d layer.\n        planes (int): Number of channels produced by some norm/conv3d layers.\n        bottleneck_mode (str): Determine which ways to factorize a 3D\n            bottleneck block using channel-separated convolutional networks.\n                If set to 'ip', it will replace the 3x3x3 conv2 layer with a\n                1x1x1 traditional convolution and a 3x3x3 depthwise\n                convolution, i.e., Interaction-preserved channel-separated\n                bottleneck block.\n                If set to 'ir', it will replace the 3x3x3 conv2 layer with a\n                3x3x3 depthwise convolution, which is derived from preserved\n                bottleneck block by removing the extra 1x1x1 convolution,\n                i.e., Interaction-reduced channel-separated bottleneck block.\n            Default: 'ir'.\n        args (position arguments): Position arguments for Bottleneck.\n        kwargs (dict, optional): Keyword arguments for Bottleneck.\n    \"\"\"\n\n    def __init__(self, inplanes, planes, *args, bottleneck_mode=\"ir\", **kwargs):\n        super(CSNBottleneck3d, self).__init__(inplanes, planes, *args, **kwargs)\n        self.bottleneck_mode = bottleneck_mode\n        conv2 = []\n        if self.bottleneck_mode == \"ip\":\n            conv2.append(\n                ConvModule(\n                    planes,\n                    planes,\n                    1,\n                    stride=1,\n                    bias=False,\n                    conv_cfg=self.conv_cfg,\n                    norm_cfg=self.norm_cfg,\n                    act_cfg=None,\n                )\n            )\n        conv2_kernel_size = self.conv2.conv.kernel_size\n        conv2_stride = self.conv2.conv.stride\n        conv2_padding = self.conv2.conv.padding\n        conv2_dilation = self.conv2.conv.dilation\n        conv2_bias = bool(self.conv2.conv.bias)\n        self.conv2 = ConvModule(\n            planes,\n            planes,\n            conv2_kernel_size,\n            stride=conv2_stride,\n            padding=conv2_padding,\n            dilation=conv2_dilation,\n            bias=conv2_bias,\n            conv_cfg=self.conv_cfg,\n            norm_cfg=self.norm_cfg,\n            act_cfg=self.act_cfg,\n            groups=planes,\n        )\n        conv2.append(self.conv2)\n        self.conv2 = nn.Sequential(*conv2)\n\n\nclass ResNet3dCSN(ResNet3d):\n    \"\"\"ResNet backbone for CSN.\n    Args:\n        depth (int): Depth of ResNetCSN, from {18, 34, 50, 101, 152}.\n        pretrained (str | None): Name of pretrained model.\n        temporal_strides (tuple[int]):\n            Temporal strides of residual blocks of each stage.\n            Default: (1, 2, 2, 2).\n        conv1_kernel (tuple[int]): Kernel size of the first conv layer.\n            Default: (3, 7, 7).\n        conv1_stride_t (int): Temporal stride of the first conv layer.\n            Default: 1.\n        pool1_stride_t (int): Temporal stride of the first pooling layer.\n            Default: 1.\n        norm_cfg (dict): Config for norm layers. required keys are `type` and\n            `requires_grad`.\n            Default: dict(type='BN3d', requires_grad=True, eps=1e-3).\n        inflate_style (str): `3x1x1` or `3x3x3`. which determines the kernel\n            sizes and padding strides for conv1 and conv2 in each block.\n            Default: '3x3x3'.\n        bottleneck_mode (str): Determine which ways to factorize a 3D\n            bottleneck block using channel-separated convolutional networks.\n                If set to 'ip', it will replace the 3x3x3 conv2 layer with a\n                1x1x1 traditional convolution and a 3x3x3 depthwise\n                convolution, i.e., Interaction-preserved channel-separated\n                bottleneck block.\n                If set to 'ir', it will replace the 3x3x3 conv2 layer with a\n                3x3x3 depthwise convolution, which is derived from preserved\n                bottleneck block by removing the extra 1x1x1 convolution,\n                i.e., Interaction-reduced channel-separated bottleneck block.\n            Default: 'ip'.\n        kwargs (dict, optional): Key arguments for \"make_res_layer\".\n    \"\"\"\n\n    def __init__(\n        self,\n        depth,\n        pretrained,\n        temporal_strides=(1, 2, 2, 2),\n        conv1_kernel=(3, 7, 7),\n        conv1_stride_t=1,\n        pool1_stride_t=1,\n        norm_cfg=dict(type=\"BN3d\", requires_grad=True, eps=1e-3),\n        inflate_style=\"3x3x3\",\n        bottleneck_mode=\"ir\",\n        bn_frozen=False,\n        drop_path_rate=0.0,\n        **kwargs,\n    ):\n        self.arch_settings = {\n            # 18: (BasicBlock3d, (2, 2, 2, 2)),\n            # 34: (BasicBlock3d, (3, 4, 6, 3)),\n            50: (CSNBottleneck3d, (3, 4, 6, 3)),\n            101: (CSNBottleneck3d, (3, 4, 23, 3)),\n            152: (CSNBottleneck3d, (3, 8, 36, 3)),\n        }\n        self.bn_frozen = bn_frozen\n        if bottleneck_mode not in [\"ip\", \"ir\"]:\n            raise ValueError(\n                f'Bottleneck mode must be \"ip\" or \"ir\",' f\"but got {bottleneck_mode}.\"\n            )\n        super(ResNet3dCSN, self).__init__(\n            depth,\n            pretrained,\n            temporal_strides=temporal_strides,\n            conv1_kernel=conv1_kernel,\n            conv1_stride_t=conv1_stride_t,\n            pool1_stride_t=pool1_stride_t,\n            norm_cfg=norm_cfg,\n            inflate_style=inflate_style,\n            bottleneck_mode=bottleneck_mode,\n            drop_path_rate=drop_path_rate,\n            **kwargs,\n        )\n\n    def train(self, mode=True):\n        super(ResNet3d, self).train(mode)\n        self._freeze_stages()\n        if mode and self.norm_eval:\n            for m in self.modules():\n                if isinstance(m, _BatchNorm):\n                    m.eval()\n                    if self.bn_frozen:\n                        for param in m.parameters():\n                            param.requires_grad = False\n\n                            \n                            \nclass InkDetResNet3dCSNModel(nn.Module):\n    def __init__(\n        self,\n        pretrained: bool = False,\n        depth: str = \"50\",\n        bottleneck_mode: str = \"ir\",\n        drop_path_rate: float = 0.0,\n        in_chans: int = 1,\n        preprocess_in_model: bool = False,\n        start_z: int = 8,\n        end_z: int = -8,\n        shift_z: int = 2,\n        sampling_z: int = 1,\n        num_class: int = 1,\n    ):\n        super().__init__()\n        self.backbone = ResNet3dCSN(\n            pretrained2d=False,\n            in_channels=in_chans,\n            pretrained=None,\n            depth=int(depth),\n            with_pool2=False,\n            bottleneck_mode=bottleneck_mode,\n            norm_eval=False,\n            zero_init_residual=False,\n            drop_path_rate=drop_path_rate,\n        )\n        self.in_chans = in_chans\n        self.head = nn.Sequential(\n            nn.BatchNorm2d(2048),\n            nn.Conv2d(\n                2048,\n                512,\n                1,\n            ),\n            nn.ReLU(inplace=True),\n            nn.BatchNorm2d(512),\n            nn.Conv2d(\n                512,\n                num_class,\n                1,\n            ),\n        )\n        self.preprocess_in_model = preprocess_in_model\n        self.start_z = start_z\n        self.end_z = end_z\n        self.shift_z = shift_z\n        self.sampling_z = sampling_z\n        if pretrained:\n            load_weights = {\n                \"50\": {\n                    \"ir\": \"../../libs/pretrained/vmz_ircsn_ig65m_pretrained_r50_32x2x1_58e_kinetics400_rgb_20210617-86d33018.pth\"\n                },\n                \"152\": {\n                    \"ir\": \"../../libs/pretrained/ircsn_ig65m-pretrained-r152_8xb12-32x2x1-58e_kinetics400-rgb_20220811-c7a3cc5b.pth\"\n                },\n            }\n            load_weight = load_weights[depth][bottleneck_mode]\n            rank_zero_info(f\"load pretrained weight {load_weight}!!!\")\n            state_dict = torch.load(load_weight, map_location=\"cpu\")  # load checkpoint\n            if \"state_dict\" in state_dict.keys():\n                state_dict = state_dict[\"state_dict\"]\n            state_dict = intersect_dicts(\n                state_dict, self.state_dict(), exclude=[]\n            )  # intersect\n            self.load_state_dict(state_dict, strict=False)  # load\n            print(\n                \"Transferred %g/%g items from %s\"\n                % (len(state_dict), len(self.state_dict()), load_weight)\n            )  # report\n            del state_dict\n            gc.collect()\n\n    def preprocess(self, img):\n        if self.training and np.random.rand() < 0.5:\n            shift = np.random.randint(-self.shift_z, self.shift_z + 1)\n        else:\n            shift = 0\n        img = img[:, self.start_z + shift : self.end_z + shift]\n        img = img[:: self.sampling_z]\n        return img\n\n    def forward_image_feats(self, img):\n        if self.preprocess_in_model:\n            img = self.preprocess(img)\n        mean = img.mean(dim=(1, 2, 3), keepdim=True)\n        std = img.std(dim=(1, 2, 3), keepdim=True) + 1e-6\n        img = (img - mean) / std\n        bs, ch, h, w = img.shape\n        assert ch % self.in_chans == 0\n        groups_3d = ch // self.in_chans\n        img = img.reshape((bs, groups_3d, self.in_chans, h, w))\n\n        if self.training:\n            ch_arr = list(range(img.shape[2]))\n            ch_arr = [\n                random.sample(ch_arr, len(ch_arr)) if np.random.rand() < 0.2 else ch_arr\n                for _ in range(img.shape[0])\n            ]\n            for i, ca in enumerate(ch_arr):\n                img[i] = img[i, :, ca]\n        img = img.permute(0, 2, 1, 3, 4).contiguous()\n        img_feat = self.backbone(img)[-1].amax(2)\n        return img_feat\n\n    def forward(\n        self,\n        img: torch.Tensor,\n    ):\n        \"\"\"\n        img: (bs, ch, h, w)\n        \"\"\"\n        img_feat = self.forward_image_feats(img)\n        return self.head(img_feat)\n    \n    \n                            \nclass InkDetLightningModel3d(pl.LightningModule):\n    def __init__(\n        self,\n        valid_fragment_id: str,\n        pretrained: bool = False,\n        depth: str = \"50\",\n        bottleneck_mode: str = \"ir\",\n        drop_path_rate: float = 0.0,\n        in_chans: int = 1,\n        preprocess_in_model: bool = False,\n        start_z: int = 8,\n        end_z: int = -8,\n        sampling_z: int = 1,\n        shift_z: int = 2,\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    ) -> None:\n        super().__init__()\n        self.__build_model(\n            pretrained=pretrained,\n            depth=depth,\n            bottleneck_mode=bottleneck_mode,\n            drop_path_rate=drop_path_rate,\n            in_chans=in_chans,\n            preprocess_in_model=preprocess_in_model,\n            start_z=start_z,\n            end_z=end_z,\n            sampling_z=sampling_z,\n            shift_z=shift_z,\n        )\n        self.save_hyperparameters()\n\n    def __build_model(\n        self,\n        pretrained: bool = False,\n        depth: str = \"50\",\n        bottleneck_mode: str = \"ir\",\n        drop_path_rate: float = 0.0,\n        in_chans: int = 1,\n        preprocess_in_model: bool = False,\n        start_z: int = 8,\n        end_z: int = -8,\n        sampling_z: int = 1,\n        shift_z: int = 2,\n    ):\n        self.model = InkDetResNet3dCSNModel(\n            pretrained=pretrained,\n            depth=depth,\n            bottleneck_mode=bottleneck_mode,\n            drop_path_rate=drop_path_rate,\n            in_chans=in_chans,\n            preprocess_in_model=preprocess_in_model,\n            start_z=start_z,\n            end_z=end_z,\n            sampling_z=sampling_z,\n            shift_z=shift_z,\n            num_class=1,\n        )\n        self.model_ema = ModelEmaV2(self.model, decay=0.99)\n        \n    def predict(self, volume):\n        output = torch.sigmoid(self.model_ema.module(volume))\n        output = F.interpolate(\n            output.float(),\n            scale_factor=32,\n            mode=\"bilinear\",\n            align_corners=True,\n        ).half()\n        return output","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-25T15:49:53.256973Z","iopub.execute_input":"2023-05-25T15:49:53.257777Z","iopub.status.idle":"2023-05-25T15:50:08.851429Z","shell.execute_reply.started":"2023-05-25T15:49:53.25772Z","shell.execute_reply":"2023-05-25T15:50:08.850212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"args = {\n    \"seed\": 2023,\n    \"image_size\": 256,\n    \"num_workers\": 2,\n    \"batch_size\":6,#8\n    \"models_conf\":\n    {   #fold0\n        \"/kaggle/input/inkdet-weights/exp055/convnext_tiny_split3d3x9csn_l6_mixup_ep30/fold0\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp055/convnext_tiny_split3d5x7csn_mixup_ep30/fold0\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp055/resnetrs50_split3d3x9csn_l6_mixup_ep30/fold0\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp055/resnetrs50_split3d5x7csn_mixup_ep30/fold0\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp055/swinv2_tiny_window8_256_split3d3x9csn_l6_mixup_ep30/fold0\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp055/swinv2_tiny_window8_256_split3d5x7csn_mixup_ep30/fold0\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp056/resnet3d50csnir1x32_mixup_ep30/fold0\": {\"infer_size\": 192, \"model\": InkDetLightningModel3d},\n        \"/kaggle/input/inkdet-weights/exp056/resnet3d152csnir1x24_mixup_ep30/fold0\": {\"infer_size\": 192, \"model\": InkDetLightningModel3d},\n        \"/kaggle/input/inkdet-weights-ron/exp055/resnext50d_32x4d_split3d3x9csn_l6_mixup_ep15/fold0\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n#         \"/kaggle/input/inkdet-weights-ron/exp055/resnext50d_32x4d_split3d5x7csn_mixup_ep15/fold0\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        #fold1\n        \"/kaggle/input/inkdet-weights/exp055/convnext_tiny_split3d3x9csn_l6_mixup_ep30/fold1\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp055/convnext_tiny_split3d5x7csn_mixup_ep30/fold1\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp055/resnetrs50_split3d3x9csn_l6_mixup_ep30/fold1\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp055/resnetrs50_split3d5x7csn_mixup_ep30/fold1\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp055/swinv2_tiny_window8_256_split3d3x9csn_l6_mixup_ep30/fold1\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp055/swinv2_tiny_window8_256_split3d5x7csn_mixup_ep30/fold1\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n        \"/kaggle/input/inkdet-weights/exp056/resnet3d50csnir1x32_mixup_ep30/fold1\": {\"infer_size\": 192, \"model\": InkDetLightningModel3d},\n        \"/kaggle/input/inkdet-weights/exp056/resnet3d152csnir1x24_mixup_ep30/fold1\": {\"infer_size\": 192, \"model\": InkDetLightningModel3d},\n        \"/kaggle/input/inkdet-weights-ron/exp055/resnext50d_32x4d_split3d3x9csn_l6_mixup_ep15/fold1\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n#         \"/kaggle/input/inkdet-weights-ron/exp055/resnext50d_32x4d_split3d5x7csn_mixup_ep15/fold1\": {\"infer_size\": 256, \"model\": InkDetLightningModel},\n    },\n    \"threshold\": 0.93,\n}\npl.seed_everything(args[\"seed\"])\nwarnings.simplefilter(\"ignore\")\npreds = []\n\nmodels_conf = args[\"models_conf\"]\nckpt_paths = [glob(f\"{logdir}/**/best_fbeta.ckpt\",recursive=True)[0] for logdir in models_conf]\nprint(f\"ckpt_path = {ckpt_paths}\")\nmodels = [models_conf[logdir][\"model\"].load_from_checkpoint(\n    glob(f\"{logdir}/**/best_fbeta.ckpt\",recursive=True)[0], valid_fragment_id=fragment_ids[0], pretrained=False, preprocess_in_model=True,\n).half().eval().to(device=device) for logdir in models_conf]\n\n\ndef inference_single(model, volume, infer_size, img_size):\n    pad = (img_size - infer_size) // 2\n    if pad > 0:\n        new_volume = volume[..., pad:-pad, pad:-pad]\n    else:\n        new_volume = volume\n    output = model.predict(new_volume)\n    new_output = torch.zeros((output.shape[0], output.shape[1], output.shape[2] + pad * 2, output.shape[2] + pad * 2)).to(device)\n    \n    if pad > 0:\n        new_output[..., pad:-pad, pad:-pad] = output\n        mask = torch.zeros((output.shape[0], output.shape[1], output.shape[2] + pad * 2, output.shape[2] + pad * 2)).to(device)\n        mask[..., pad:-pad, pad:-pad] = torch.ones_like(output)\n    else:\n        new_output = output\n        mask = torch.ones((output.shape[0], output.shape[1], output.shape[2] + pad * 2, output.shape[2] + pad * 2)).to(device)\n    return new_output, mask\n\ndef inference(fragment_id):\n    valid_volume_paths = np.asarray(\n                    sorted(\n                        glob(\n                            f\"vesuvius_patches_{PATCH_SIZE}/test/{fragment_id}/surface_volume/**/*.npy\",\n                            recursive=True,\n                        )\n                    )\n                )\n    dataloader = InkDetDataModule(\n        train_volume_paths=valid_volume_paths,\n        valid_volume_paths=valid_volume_paths,\n        image_size=args[\"image_size\"],\n        num_workers=args[\"num_workers\"],\n        batch_size=args[\"batch_size\"],\n        preprocess_in_model=True,\n    ).test_dataloader()\n    fragment_mask = np.array(\n        Image.open(\n            f\"/kaggle/input/vesuvius-challenge-ink-detection/test/{fragment_id}/mask.png\"\n        ).convert(\"1\")\n    ) > 0\n    p_valid = np.stack([np.zeros_like(fragment_mask, dtype=np.float16) for _ in models]) # (models, 1, h, w)\n    count_pix = np.stack([np.zeros_like(fragment_mask, dtype=np.uint8) for _ in models]) # (models, 1, h, w)\n    count = 0\n    for batch in tqdm(dataloader):\n        volume_cpu, _, x, y = batch # volume_cpu: (bs, ch, h, w)\n        tta_ops = (np.arange(len(volume_cpu)) + count) % 4\n        count += len(volume_cpu)\n        for i, j in zip(range(len(volume_cpu)), tta_ops):\n            if j == 1:\n                volume_cpu[i] = volume_cpu[i].flip(1)\n            elif j == 2:\n                volume_cpu[i] = volume_cpu[i].flip(2)\n            elif j == 3:\n                volume_cpu[i] = volume_cpu[i].flip(1).flip(2)\n        volume = volume_cpu.to(device)\n        with torch.no_grad():\n            pred_batch = []\n            mask_batch = []\n            for model, logdir in zip(models, models_conf):\n                pred_batch_single, mask_batch_single = inference_single(model, volume, models_conf[logdir][\"infer_size\"], args[\"image_size\"])\n                pred_batch.append(pred_batch_single)\n                mask_batch.append(mask_batch_single)\n            pred_batch = torch.stack(pred_batch).permute((1, 0, 2, 3, 4)) # (models, bs, 1, h, w) -> (bs, models, 1, h, w)\n            mask_batch = torch.stack(mask_batch).permute((1, 0, 2, 3, 4)) # (models, bs, 1, h, w) -> (bs, models, 1, h, w)\n        for i, j in zip(range(len(pred_batch)), tta_ops): # pred_batch: (bs, models, ch, h, w)\n            if j == 1:\n                pred_batch[i] = pred_batch[i].flip(2)\n            elif j == 2:\n                pred_batch[i] = pred_batch[i].flip(3)\n            elif j == 3:\n                pred_batch[i] = pred_batch[i].flip(2).flip(3)\n        for xi, yi, pi, mi in zip(\n            x,\n            y,\n            pred_batch,\n            mask_batch,\n        ):\n            patch = pi.to(torch.float16).detach().cpu().numpy()\n            mask = mi.to(torch.uint8).cpu().numpy()\n            y_lim, x_lim = fragment_mask[..., \n                yi * PATCH_SIZE : yi * PATCH_SIZE + volume.shape[-2],\n                xi * PATCH_SIZE : xi * PATCH_SIZE + volume.shape[-1],\n            ].shape\n            p_valid[..., \n                yi * PATCH_SIZE : yi * PATCH_SIZE + volume.shape[-2],\n                xi * PATCH_SIZE : xi * PATCH_SIZE + volume.shape[-1],\n            ] += patch[..., 0, :y_lim, :x_lim]\n            count_pix[..., \n                yi * PATCH_SIZE : yi * PATCH_SIZE + volume.shape[-2],\n                xi * PATCH_SIZE : xi * PATCH_SIZE + volume.shape[-1],\n            ] += mask[..., 0, :y_lim, :x_lim]\n        del volume_cpu, pred_batch, mask_batch, volume\n        gc.collect()\n        torch.cuda.empty_cache()\n    count_pix *= fragment_mask[None, ...]\n    p_valid /= count_pix\n    p_valid = np.stack([np.nan_to_num(p, posinf=0, neginf=0) for p in p_valid])\n    p_valid *= fragment_mask[None, ...]\n    print([(c.max(), c.mean()) for c in count_pix], [(p.max(), p.mean()) for p in p_valid])\n    count_pix = count_pix.mean(0) > 0\n    p_valid = p_valid.mean(0)\n    p_valid_tmp = p_valid.reshape(-1)[np.where(count_pix.reshape(-1))]\n    p_valid = p_valid > np.quantile(p_valid_tmp, args[\"threshold\"])\n    del dataloader, fragment_mask, count_pix, p_valid_tmp\n    gc.collect()\n    torch.cuda.empty_cache()\n    return p_valid\n\nfor valid_idx in fragment_ids:\n    print(valid_idx)\n    preds.append(inference(valid_idx))","metadata":{"execution":{"iopub.status.busy":"2023-05-25T15:50:08.853434Z","iopub.execute_input":"2023-05-25T15:50:08.855407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle(output):\n    flat_img = output.astype(np.uint8).flatten()\n    starts = np.array((flat_img[:-1] == 0) & (flat_img[1:] == 1))\n    ends = np.array((flat_img[:-1] == 1) & (flat_img[1:] == 0))\n    starts_ix = np.where(starts)[0] + 2\n    ends_ix = np.where(ends)[0] + 2\n    lengths = ends_ix - starts_ix\n    return \" \".join(map(str, sum(zip(starts_ix, lengths), ())))\n\ndef fast_rle(img):\n    flat_img = img.flatten().astype(np.uint8)\n\n    starts = np.array((flat_img[:-1] == 0) & (flat_img[1:] == 1))\n    ends = np.array((flat_img[:-1] == 1) & (flat_img[1:] == 0))\n    starts_ix = np.where(starts)[0] + 2\n    ends_ix = np.where(ends)[0] + 2\n    lengths = ends_ix - starts_ix\n    predicted_arr = np.stack([starts_ix, lengths]).T.flatten()\n    f = StringIO()\n    np.savetxt(f, predicted_arr.reshape(1, -1), delimiter=\" \", fmt=\"%d\")\n    predicted = f.getvalue().strip()\n    return predicted","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom collections import defaultdict\nimport pandas as pd\n\nsubmission = defaultdict(list)\n# for fragment_id, fragment_name in enumerate(test_fragments):\nfor fragment_id, pred_image in zip(fragment_ids, preds):\n    plt.imshow(pred_image)\n    plt.show()\n    submission[\"Id\"].append(fragment_id)\n    submission[\"Predicted\"].append(fast_rle(pred_image))\n\npd.DataFrame.from_dict(submission).to_csv(\"/kaggle/working/submission.csv\", index=False)\npd.DataFrame.from_dict(submission)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}