{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/pretrainedmodels/pretrainedmodels-0.7.4')\nsys.path.append('/kaggle/input/efficientnet-pytorch/EfficientNet-PyTorch-master')\nsys.path.append('/kaggle/input/timm-pytorch-image-models/pytorch-image-models-master')\nsys.path.append('/kaggle/input/segmentation-models-pytorch/segmentation_models.pytorch-master')\nfrom sklearn.metrics import roc_auc_score, accuracy_score, f1_score, log_loss\nimport pickle\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.cuda.amp import autocast, GradScaler\nimport warnings\nimport sys\nimport pandas as pd\nimport os\nimport gc\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nfrom itertools import combinations\nimport cv2\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport argparse\nimport importlib\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam, SGD, AdamW\n\nimport datetime\nimport segmentation_models_pytorch as smp\nfrom segmentation_models_pytorch.base import SegmentationHead\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\nimport torch.nn.functional as F\nsys.path.append(\"/kaggle/input/resnet3d\")\nfrom resnet3d import generate_model\nimport filecmp","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:21.527624Z","iopub.execute_input":"2023-06-02T02:41:21.528455Z","iopub.status.idle":"2023-06-02T02:41:29.145299Z","shell.execute_reply.started":"2023-06-02T02:41:21.528415Z","shell.execute_reply":"2023-06-02T02:41:29.144087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    # exp\n    comp_name = 'vesuvius'\n    comp_dir_path = '/kaggle/input/'\n    comp_folder_name = 'vesuvius-challenge-ink-detection'\n    comp_dataset_path = f'{comp_dir_path}{comp_folder_name}/'\n    exp_name = \"craniotabesmayoress\"\n    print(exp_name)\n\n    # Models\n    model = '3D-2D' # 2.5D, 3D-2D, ink-classifier\n    slices = (20, 36)\n    in_chans = slices[1] - slices[0]\n    target_size = 1\n\n    # 2.5D\n    encoder = 'efficientnet_b0'\n    decoder = 'Unet' # Unet, Unet++, DeepLabV3+\n    depth = 4\n    decoder_channels = [128, 64, 32, 16]\n    drop_rate = 0.3\n    drop_path_rate = 0.2\n    use_pool = True\n\n    # 3D-2D\n    pooler = \"attention\" # projection, attention, mean, conv, conv_attention, max\n    encoder3d = \"resnet\" # 2D, resnet, custom\n    decoder3d2d = \"FPN\"\n    encoder_depth = 34\n    weight_path3d = f\"r3d{encoder_depth}_KM_200ep.pth\"\n    use_denoiser = True\n\n    # ink-classifier\n    region_size = 64\n    region_stride = region_size//4\n    feature_pooler = \"mean\" # mean, max, None, gem\n    \n    # training\n    size = 224\n    tile_size = 224\n    sampling_method = \"adaptive_stride_random\"\n    filter_method = \"mean\"\n    stride = tile_size // 4\n    sample_stride = tile_size // 4\n    neg_sample_size = 8\n    n_folds = 5\n    folds = list(combinations(range(1, n_folds + 1), 1))\n    train_dir = f\"/media/fql/Data/Kaggle/Vesuvius Challenge - Ink Detection/data/sampled/{size}*{size}_stride{stride}_negsample{neg_sample_size}__slice{slices}method({sampling_method})/\"  \n    train_batch_size = 16\n    valid_batch_size = train_batch_size * 2\n    use_amp = True\n    grad_clipping = True\n\n    scheduler = \"gradualwarmup\"\n    num_cycles=0.5\n    num_warmup_steps_rate=0.1\n    epochs = 5\n    \n    if scheduler == \"gradualwarmup\":\n        warmup_factor = 10\n        lr = 1e-4 / warmup_factor\n    else:\n        lr = 1e-4\n    loss = \"BCE\"\n\n    pretrained = True\n    inf_weight = 'best'  # 'best'\n    min_lr = 1e-6\n    weight_decay = 1e-3\n    max_grad_norm = 1000\n    num_workers = 0\n    seed = 42\n\n    # output\n    outputs_path = f'/home/fql/Kaggle/Vesuvius Challenge - Ink Detection/experiments/{exp_name}/'\n\n    model_dir = outputs_path + \\\n        f'{comp_name}-models/'\n\n    log_dir = outputs_path + 'logs/'\n    log_path = log_dir + f'{exp_name}'\n    \n    # augmentations\n    train_aug_list = [\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomBrightnessContrast(p=0.5),\n        A.ShiftScaleRotate(p=0.5),\n        A.OneOf([\n                A.GaussNoise(var_limit=[10, 50]),\n                A.GaussianBlur(),\n                A.MotionBlur(),\n                ], p=0.5),\n        A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.5),\n        A.CoarseDropout(max_holes=1, max_width=int(size * 0.3), max_height=int(size * 0.3), \n                        mask_fill_value=0, p=0.5),\n        A.Normalize(\n            mean= [0] * in_chans,\n            std= [1] * in_chans\n        ),\n        ToTensorV2(transpose_mask=True),\n    ]\n\n    valid_aug_list = [\n        A.Normalize(\n            mean= [0] * in_chans,\n            std= [1] * in_chans\n        ),\n        ToTensorV2(transpose_mask=True),\n    ]","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.147346Z","iopub.execute_input":"2023-06-02T02:41:29.147693Z","iopub.status.idle":"2023-06-02T02:41:29.172498Z","shell.execute_reply.started":"2023-06-02T02:41:29.147662Z","shell.execute_reply":"2023-06-02T02:41:29.171240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IS_DEBUG = False\nmode = 'train' if IS_DEBUG else 'test'      \nTH = 0.52\nTTA = True","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.174407Z","iopub.execute_input":"2023-06-02T02:41:29.175128Z","iopub.status.idle":"2023-06-02T02:41:29.182535Z","shell.execute_reply.started":"2023-06-02T02:41:29.175087Z","shell.execute_reply":"2023-06-02T02:41:29.181349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.186767Z","iopub.execute_input":"2023-06-02T02:41:29.187087Z","iopub.status.idle":"2023-06-02T02:41:29.253530Z","shell.execute_reply.started":"2023-06-02T02:41:29.187058Z","shell.execute_reply":"2023-06-02T02:41:29.252330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## helper","metadata":{}},{"cell_type":"code","source":"# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\ndef rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    \n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\n\ndef fbeta_numpy(targets, preds, beta=0.5, smooth=1e-5):\n    \"\"\"\n    https://www.kaggle.com/competitions/vesuvius-challenge-ink-detection/discussion/397288\n    \"\"\"\n    y_true_count = targets.sum()\n    ctp = preds[targets == 1].sum()\n    cfp = preds[targets == 0].sum()\n    beta_squared = beta * beta\n\n    c_precision = ctp / (ctp + cfp + smooth)\n    c_recall = ctp / (y_true_count + smooth)\n    dice = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall + smooth)\n\n    return dice\n\n\ndef calc_fbeta(mask, mask_pred, th):\n    mask = mask.astype(int).flatten()\n    mask_pred = mask_pred.flatten()\n    dice = fbeta_numpy(mask, (mask_pred >= th).astype(int), beta=0.5)\n\n    return dice\n\n\ndef calc_cv(mask_gts, mask_preds, orig_sizes):\n    best_th = 0\n    best_dice = 0\n    ths = np.array(range(30, 90 + 1, 5)) / 100\n\n    for th in ths:\n        orig_h, orig_w = orig_sizes\n        mask_gt = mask_gts[:orig_h, :orig_w]\n        mask_pred = mask_preds[:orig_h, :orig_w]\n        dice = calc_fbeta(mask_gt, mask_pred, th)\n        print(f\"th: {th} dice: {dice}\")\n\n        if dice > best_dice:\n            best_dice = dice\n            best_th = th\n\n    return best_dice, best_th","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.255492Z","iopub.execute_input":"2023-06-02T02:41:29.255883Z","iopub.status.idle":"2023-06-02T02:41:29.272513Z","shell.execute_reply.started":"2023-06-02T02:41:29.255840Z","shell.execute_reply":"2023-06-02T02:41:29.271224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## dataset","metadata":{}},{"cell_type":"code","source":"def read_image(fragment_id):\n    images = []\n\n    # idxs = range(65)\n    #mid = 65 // 2\n    #start = mid - CFG.in_chans // 2\n    #end = mid + CFG.in_chans // 2\n    idxs = range(CFG.slices[0], CFG.slices[1])\n\n    for i in idxs:\n        \n        image = cv2.imread(CFG.comp_dataset_path + f\"{mode}/{fragment_id}/surface_volume/{i:02}.tif\", 0)\n\n        pad0 = (CFG.tile_size - image.shape[0] % CFG.tile_size)\n        pad1 = (CFG.tile_size - image.shape[1] % CFG.tile_size)\n\n        image = np.pad(image, [(0, pad0), (0, pad1)], constant_values=0)\n\n        images.append(image)\n    images = np.stack(images, axis=2)\n    \n    return images","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.274203Z","iopub.execute_input":"2023-06-02T02:41:29.275467Z","iopub.status.idle":"2023-06-02T02:41:29.285757Z","shell.execute_reply.started":"2023-06-02T02:41:29.275422Z","shell.execute_reply":"2023-06-02T02:41:29.284896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(data, cfg):\n    if data == 'train':\n        aug = A.Compose(cfg.train_aug_list)\n    elif data == 'valid':\n        aug = A.Compose(cfg.valid_aug_list)\n\n    # print(aug)\n    return aug\n\nclass CustomDataset(Dataset):\n    def __init__(self, images, cfg, labels=None, transform=None):\n        self.images = images\n        self.cfg = cfg\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        # return len(self.xyxys)\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        # x1, y1, x2, y2 = self.xyxys[idx]\n        image = self.images[idx]\n        data = self.transform(image=image)\n        image = data['image']\n        if CFG.model != \"2.5D\":\n            image = torch.unsqueeze(image, dim=0)\n        return image\n    \nclass InkDataset(Dataset):\n    def __init__(self, volumes, labels, transform=None, mode=\"train\"):\n        self.volumes = volumes\n        if mode == \"train\":\n            self.labels = torch.Tensor(labels).float()\n        self.transform = transform\n        self.mode = mode\n\n    def __len__(self):\n        return len(self.volumes)\n\n    def __getitem__(self, idx):\n        image = self.volumes[idx]\n        if self.transform:\n            data = self.transform(image=image)\n            image = data['image']\n        image = torch.unsqueeze(image, dim=0)\n        if self.mode == \"train\":\n            label = self.labels[idx]\n            label = torch.unsqueeze(label, dim=0)\n            return image, label\n        else:\n            return image","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.287999Z","iopub.execute_input":"2023-06-02T02:41:29.289171Z","iopub.status.idle":"2023-06-02T02:41:29.304360Z","shell.execute_reply.started":"2023-06-02T02:41:29.289126Z","shell.execute_reply":"2023-06-02T02:41:29.303409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_test_dataset(fragment_id):\n    test_images = read_image(fragment_id)\n    \n    x1_list = list(range(0, test_images.shape[1]-CFG.tile_size+1, CFG.stride))\n    y1_list = list(range(0, test_images.shape[0]-CFG.tile_size+1, CFG.stride))\n    \n    test_images_list = []\n    xyxys = []\n    for y1 in y1_list:\n        for x1 in x1_list:\n            y2 = y1 + CFG.tile_size\n            x2 = x1 + CFG.tile_size\n            image = test_images[y1:y2, x1:x2]\n            #image = np.pad(image, [(16,16), (16,16), (0,0)])\n            test_images_list.append(image)\n            xyxys.append((x1, y1, x2, y2))\n    xyxys = np.stack(xyxys)\n            \n    test_dataset = CustomDataset(test_images_list, CFG, transform=get_transforms(data='valid', cfg=CFG))\n    \n    test_loader = DataLoader(test_dataset,\n                          batch_size=CFG.valid_batch_size,\n                          shuffle=False,\n                          num_workers=2, pin_memory=True, drop_last=False)\n    \n    return test_loader, xyxys\n\ndef get_ink_data_infer(test_dir, size, stride=0):\n    mask_path = f\"{test_dir}mask.png\"\n    mask = cv2.imread(mask_path, 0) / 255.\n    images = []\n    idxs = range(CFG.slices[0], CFG.slices[1])\n    for i in idxs:\n        image = cv2.imread(f\"{test_dir}surface_volume/{i:02}.tif\", 0)\n        images.append(image)\n    images = np.stack(images, axis=2)\n    radius = int(size // 2)\n    # Create a Boolean array mask of the same shape as the mask, initially all True\n    not_border = np.zeros(mask.shape, dtype=bool)\n    not_border[radius:mask.shape[0] - radius, radius:mask.shape[1] - radius] = True\n    arr_mask = np.array(mask) * not_border\n    pixels = np.argwhere(arr_mask)\n    if stride != 0:\n        sparse_mask = np.zeros(mask.shape, dtype=bool)\n        sparse_mask[::stride, ::stride] = True\n        pixels = np.argwhere(sparse_mask * arr_mask)\n    else:\n        pixels = np.argwhere(arr_mask)\n\n    volumes = []\n    for y, x in pixels:\n        subvolume = images[y - radius:y + radius, x - radius:x + radius, :]\n        volumes.append(subvolume)\n    return volumes, pixels, mask","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.306294Z","iopub.execute_input":"2023-06-02T02:41:29.307090Z","iopub.status.idle":"2023-06-02T02:41:29.325230Z","shell.execute_reply.started":"2023-06-02T02:41:29.307050Z","shell.execute_reply":"2023-06-02T02:41:29.324103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model","metadata":{}},{"cell_type":"code","source":"# poolers and decoders\nclass VolumeProjection(torch.nn.Module):\n    def __init__(self, input_shape, output_shape):\n        super().__init__()\n        self.fc = torch.nn.Linear(int(np.prod(input_shape)), int(np.prod(output_shape)))\n        self.flatten = torch.nn.Flatten(start_dim=2)\n        self.output_shape = output_shape\n\n    def forward(self, x):\n        y = self.flatten(x)\n        y = self.fc(y)\n        y = y.view((y.shape[0], y.shape[1], self.output_shape[0], self.output_shape[1]))\n        return y\n\n\nclass AttentionPool(torch.nn.Module):\n    def __init__(self, depth, height, width):\n        super().__init__()\n        self.attention_weights = nn.Parameter(torch.ones(1, 1, depth, height, width))\n        self.softmax = nn.Softmax(dim=2)\n\n    def forward(self, x):\n        # Apply softmax along the depth dimension to obtain attention weights\n        attention_weights = self.softmax(self.attention_weights)\n        # Perform attention pooling by multiplying the attention weights with the input tensor\n        pooled_output = torch.mul(attention_weights, x)\n        # Sum the pooled output along the depth dimension\n        pooled_output = torch.sum(pooled_output, dim=2)\n        return pooled_output\n\n\nclass Subvolume3DcnnEncoder(nn.Module):\n    def __init__(self, batch_norm_momentum, filters):\n        super().__init__()\n        strides = [1, 2, 2, 2]\n        filter_sizes = [1] + filters\n        filter_list_pairs = list(zip(filter_sizes[:-1], filter_sizes[1:]))  # [(1, 16), (16, 32), (32, 64), (64, 128)]\n        self.conv_layers = nn.Sequential(\n            *[nn.Sequential(\n                nn.Conv3d(chan_in, chan_out, kernel_size=3, stride=stride, padding=1),\n                nn.ReLU(),\n                nn.BatchNorm3d(num_features=filter_, momentum=batch_norm_momentum)\n            )\n                for (chan_in, chan_out), stride, filter_ in zip(filter_list_pairs, strides, filters)])\n        self.apply(self.init_weight)\n\n    @staticmethod\n    def init_weight(m):\n        if isinstance(m, nn.Conv3d):\n            nn.init.xavier_uniform_(m.weight)\n            nn.init.zeros_(m.bias)\n\n    def forward(self, x):\n        return self.conv_layers(x)\n\n\nclass GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        return self.gem(x, p=self.p, eps=self.eps)\n        \n    def gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool3d(x.clamp(min=eps).pow(p), (x.size(-3), x.size(-2), x.size(-1))).pow(1./p)\n\n\nclass Decoder(nn.Module):\n    def __init__(self, encoder_dims, upscale):\n        super().__init__()\n        self.convs = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(\n                    in_channels=encoder_dims[i] + encoder_dims[i - 1],\n                    out_channels=encoder_dims[i - 1],\n                    kernel_size=3,\n                    stride=1,\n                    padding=1,\n                    bias=False\n                ),\n                nn.BatchNorm2d(encoder_dims[i - 1]),\n                nn.ReLU(inplace=True)\n            ) for i in range(1, len(encoder_dims))])\n\n        self.logit = nn.Conv2d(encoder_dims[0], 1, 1, 1, 0)\n        self.up = nn.Upsample(scale_factor=upscale, mode=\"bilinear\")\n\n    def forward(self, feature_maps):\n        for i in range(len(feature_maps) - 1, 0, -1):\n            f_up = F.interpolate(feature_maps[i], scale_factor=2, mode=\"bilinear\")\n            f = torch.cat([feature_maps[i - 1], f_up], dim=1)\n            f_down = self.convs[i - 1](f)\n            feature_maps[i - 1] = f_down\n\n        mask = self.logit(feature_maps[0])\n        mask = self.up(mask)\n        return mask\n\n\n# 2.5D\nclass CustomModel(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n\n        self.cfg = cfg\n        if cfg.use_pool:\n            in_chans = 1\n        else:\n            in_chans = cfg.in_chans\n            \n        self.encoder = timm.create_model(\n            cfg.encoder,\n            in_chans=in_chans,\n            features_only=True,\n            drop_rate=cfg.drop_rate,\n            drop_path_rate=cfg.drop_path_rate,\n            out_indices=tuple(range(cfg.depth)),\n            pretrained=False\n        )\n\n        temp = self.encoder(torch.rand((1, 1, cfg.size, cfg.size)))\n        \n        encoder_channels = [in_chans] + self.encoder.feature_info.channels()\n        decoder_channels = cfg.decoder_channels\n        \n        if cfg.use_pool:\n            self.pooler = nn.ModuleList([AttentionPool(self.cfg.in_chans, x.shape[-2], x.shape[-1]) for x in temp])\n\n        if cfg.decoder == \"Unet\":\n            self.decoder = smp.decoders.unet.decoder.UnetDecoder(\n                encoder_channels=encoder_channels,\n                decoder_channels=decoder_channels,\n                n_blocks=cfg.depth,\n            )\n            self.segmentation_head = SegmentationHead(\n                in_channels=decoder_channels[-1],\n                out_channels=cfg.target_size,\n                activation=None,\n                kernel_size=3,\n            )\n        elif cfg.decoder == \"Unet++\":\n            self.decoder = smp.decoders.unetplusplus.decoder.UnetPlusPlusDecoder(\n                encoder_channels=encoder_channels,\n                decoder_channels=decoder_channels,\n                n_blocks=cfg.depth,\n            )\n            self.segmentation_head = SegmentationHead(\n                in_channels=decoder_channels[-1],\n                out_channels=cfg.target_size,\n                activation=None,\n                kernel_size=3,\n            )\n        elif cfg.decoder == \"DeepLabV3+\":\n            self.decoder = smp.decoders.deeplabv3.decoder.DeepLabV3PlusDecoder(\n                encoder_channels=encoder_channels[:cfg.depth + 1],\n            )\n            self.segmentation_head = SegmentationHead(\n                in_channels=self.decoder.out_channels,\n                out_channels=cfg.target_size,\n                activation=None,\n                kernel_size=1,\n                upsampling=4,\n            )\n            \n        if cfg.use_denoiser:\n            self.denoiser = smp.Unet(\n                encoder_name=\"tu-resnet10t\", # \"tu-resnet10t\" \"resnet18\"\n                encoder_weights=\"imagenet\",\n                in_channels=1,\n                classes=1,\n                activation=None,\n            )\n\n    def get_features(self, x):\n        feat_maps = self.encoder(x)\n        return feat_maps\n\n    def forward(self, x):\n        if self.cfg.use_pool:\n            bs = x.shape[0]\n            x = x.view((-1, 1, x.shape[-2], x.shape[-1]))\n            feat_maps = self.get_features(x)\n            feat_maps_pooled = []\n            for i, f in enumerate(feat_maps):\n                f = f.view((bs, -1, self.cfg.in_chans, f.shape[-2], f.shape[-1]))\n                o = self.pooler[i](f)\n                feat_maps_pooled.append(o)\n        else:\n            feat_maps_pooled = self.get_features(x)\n        feat_maps_pooled = [x] + feat_maps_pooled\n        decoder_output = self.decoder(*feat_maps_pooled)\n        masks = self.segmentation_head(decoder_output)\n        \n        if self.cfg.use_denoiser:\n            noise = self.denoiser(masks)\n            masks = masks - noise\n        \n        return masks\n\n\n# 3D-2D\nclass SegModel(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n        if cfg.encoder3d == \"resnet\":\n            # https://github.com/kenshohara/3D-ResNets-PyTorch\n            self.encoder = generate_model(model_depth=cfg.encoder_depth, n_input_channels=1)\n\n        temp = self.encoder(torch.rand((1, 1, cfg.in_chans, cfg.size, cfg.size)))\n        \n        if cfg.decoder3d2d == \"Custom\":\n            self.decoder = Decoder(encoder_dims=[64, 128, 256, 512], upscale=4)\n        elif cfg.decoder3d2d == \"FPN\":\n            self.decoder = smp.decoders.fpn.decoder.FPNDecoder(\n                encoder_channels=[64, 128, 256, 512],\n                encoder_depth=4,\n                pyramid_channels=256,\n                segmentation_channels=128,\n                dropout=0.2,\n                merge_policy=\"add\",\n            )\n            self.segmentation_head = SegmentationHead(\n                in_channels=128,\n                out_channels=cfg.target_size,\n                activation=None,\n                kernel_size=1,\n                upsampling=4,\n            )\n\n        if cfg.pooler == \"projection\":\n            self.poolers = nn.ModuleList([VolumeProjection(x.shape[2:], x.shape[3:]) for x in temp])\n            \n        elif cfg.pooler == \"conv\":\n            self.poolers = nn.ModuleList(\n                [\n                    nn.Sequential(\n                        nn.Conv2d(in_channels=x.shape[2] * x.shape[1], out_channels=x.shape[1], kernel_size=1),\n                        nn.BatchNorm2d(x.shape[1]),\n                        nn.ReLU(inplace=True)\n                    ) for x in temp\n                ]\n            )\n            \n        elif cfg.pooler == \"attention\":\n            self.poolers = nn.ModuleList([AttentionPool(x.shape[2], x.shape[3], x.shape[4]) for x in temp])\n\n        elif cfg.pooler == \"conv_attention\":\n            self.poolers = nn.ModuleList([\n                nn.Sequential(\n                    nn.Conv3d(x.shape[1], x.shape[1], kernel_size=3, padding=1),\n                    nn.ReLU(inplace=True)\n                ) for x in temp\n            ])\n            \n        if cfg.use_denoiser:\n            self.denoiser = smp.Unet(\n                encoder_name=\"tu-resnet10t\", # \"tu-resnet10t\" \"resnet18\"\n                encoder_weights=None,\n                in_channels=1,\n                classes=1,\n                activation=None,\n            )\n\n    def forward(self, x):\n        feat_maps = self.encoder(x)\n\n        if self.cfg.pooler == \"projection\" or self.cfg.pooler == \"attention\":\n            feat_maps_pooled = []\n            for i, x in enumerate(feat_maps):\n                pooled = self.poolers[i](x)\n                feat_maps_pooled.append(pooled)\n\n        elif self.cfg.pooler == \"conv\":\n            feat_maps_pooled = []\n            for i, x in enumerate(feat_maps):\n                x = x.view((x.shape[0], x.shape[1] * x.shape[2], x.shape[3], x.shape[4]))\n                pooled = self.poolers[i](x)\n                feat_maps_pooled.append(pooled)\n\n        elif self.cfg.pooler == \"mean\":\n            feat_maps_pooled = [torch.mean(f, dim=2) for f in feat_maps]\n            \n        elif self.cfg.pooler == \"max\":\n            feat_maps_pooled = [torch.max(f, dim=2) for f in feat_maps]\n\n        elif self.cfg.pooler == \"conv_attention\":\n            feat_maps_pooled = []\n            for i, x in enumerate(feat_maps):\n                w = self.poolers[i](x)\n                w = F.softmax(w, 2)\n                pooled = (w * x).sum(2)\n                feat_maps_pooled.append(pooled)\n               \n        if self.cfg.decoder3d2d == \"Custom\":        \n            pred_mask = self.decoder(feat_maps_pooled)\n        else:\n            pred_mask = self.decoder(*feat_maps_pooled)\n            pred_mask = self.segmentation_head(pred_mask)\n        if self.cfg.use_denoiser:\n            noise = self.denoiser(pred_mask)\n            pred_mask = pred_mask - noise\n        \n        return pred_mask\n\n    def load_pretrained_weights(self, state_dict):\n        # Convert 3 channel weights to single channel\n        # ref - https://timm.fast.ai/models#Case-1:-When-the-number-of-input-channels-is-1\n        conv1_weight = state_dict['conv1.weight']\n        state_dict['conv1.weight'] = conv1_weight.sum(dim=1, keepdim=True)\n        self.encoder.load_state_dict(state_dict, strict=False)\n\n\n# ink-classifier\nclass LinearInkDecoder(nn.Module):\n    def __init__(self, cfg, input_shape):\n        super().__init__()\n        self.fc = nn.Linear(input_shape, 1)\n        if cfg.feature_pooler == \"mean\":\n            self.pool = nn.AdaptiveAvgPool3d(1)\n        elif cfg.feature_pooler == \"max\":\n            self.pool = nn.AdaptiveMaxPool3d(1)\n    def forward(self, x):\n        x = self.pool(x).squeeze(dim=-1).squeeze(dim=-1).squeeze(dim=-1)\n        return self.fc(x)\n\n\nclass InkClassifier(nn.Module):\n    def __init__(self, cfg, batch_norm_momentum=0.1, filters=[16, 32, 64, 128]):\n        super().__init__()\n        self.cfg = cfg\n        if cfg.encoder3d == \"custom\":\n            self.encoder = Subvolume3DcnnEncoder(batch_norm_momentum, filters)\n            dim = filters[-1]\n            self.decoder = LinearInkDecoder(cfg, dim)\n        elif cfg.encoder3d == \"resnet\":\n            self.encoder = generate_model(model_depth=CFG.encoder_depth, n_input_channels=1)\n            dim = 512\n            if cfg.feature_pooler in [\"mean\", \"max\"]:\n                self.decoder = LinearInkDecoder(cfg, dim)\n            \n            elif cfg.feature_pooler == \"gem\":\n                self.pool = nn.Sequential(\n                    GeM(),\n                    nn.Flatten(),\n                )\n                temp = self.pool(self.encoder(torch.rand((1, 1, cfg.in_chans, cfg.region_size, cfg.region_size)))[-1])\n                dim = temp.shape[-1]\n                self.decoder = nn.Sequential(\n                    GeM(),\n                    nn.Flatten(),\n                    nn.Linear(dim, cfg.target_size),\n                )\n                \n            elif cfg.feature_pooler == \"None\":\n                self.pool = nn.Flatten()\n                temp = self.pool(self.encoder(torch.rand((1, 1, CFG.in_chans, cfg.region_size, cfg.region_size)))[-1])\n                dim = temp.shape[-1]\n                self.decoder = nn.Sequential(\n                    nn.Flatten(),\n                    nn.Linear(dim, cfg.target_size),\n                )\n        elif cfg.encoder3d == \"2D\":\n            self.encoder = timm.create_model(\n                cfg.encoder,\n                in_chans=1,\n                num_classes=0,\n                drop_rate=cfg.drop_rate,\n                drop_path_rate=cfg.drop_path_rate,\n                pretrained=False\n            )\n            dim = self.encoder(torch.rand(1, 1, cfg.region_size, cfg.region_size)).shape[-1]\n            self.decoder = nn.Linear(dim * cfg.in_chans, cfg.target_size)\n\n    def forward(self, x):\n        if self.cfg.encoder3d == \"custom\":\n            x = self.encoder(x)\n            return self.decoder(x)\n        elif self.cfg.encoder3d == \"resnet\":\n            x = self.encoder(x)[-1]\n            return self.decoder(x)\n        elif self.cfg.encoder3d == \"2D\":\n            bs = x.shape[0]\n            x = x.view(-1, 1, self.cfg.region_size, self.cfg.region_size)\n            x = self.encoder(x)\n            x = x.view(bs, -1)\n            return self.decoder(x)\n\n    def load_pretrained_weights(self, state_dict):\n        # Convert 3 channel weights to single channel\n        # ref - https://timm.fast.ai/models#Case-1:-When-the-number-of-input-channels-is-1\n        conv1_weight = state_dict['conv1.weight']\n        state_dict['conv1.weight'] = conv1_weight.sum(dim=1, keepdim=True)\n        self.encoder.load_state_dict(state_dict, strict=False)\n\n\ndef build_model(cfg):\n    if cfg.model == \"2.5D\":\n        model = CustomModel(cfg)\n    elif cfg.model == \"3D-2D\":\n        model = SegModel(cfg)\n    elif cfg.model == \"ink-classifier\":\n        model = InkClassifier(cfg, filters=[16, 32, 64, 128])\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.329121Z","iopub.execute_input":"2023-06-02T02:41:29.329832Z","iopub.status.idle":"2023-06-02T02:41:29.419861Z","shell.execute_reply.started":"2023-06-02T02:41:29.329791Z","shell.execute_reply":"2023-06-02T02:41:29.418729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EnsembleModel:\n    def __init__(self, use_tta=False):\n        self.models = []\n        self.use_tta = use_tta\n\n    def __call__(self, x):\n        temps = []\n        if self.use_tta:\n            xs = [\n                x,\n                torch.rot90(x, k=1, dims=(-2, -1)),\n                torch.rot90(x, k=2, dims=(-2, -1)),\n                torch.rot90(x, k=3, dims=(-2, -1)),\n            ]\n            for x in xs:\n                temp = []\n                for m in self.models:\n                    temp.append(torch.sigmoid(m(x)))\n                temps.append(torch.mean(torch.stack(temp), dim=0))\n            if CFG.model != \"ink-classifier\":\n                temps = [\n                    temps[0],\n                    torch.rot90(temps[1], k=-1, dims=(-2, -1)),\n                    torch.rot90(temps[2], k=-2, dims=(-2, -1)),\n                    torch.rot90(temps[3], k=-3, dims=(-2, -1)),\n                ]\n        else:\n            for m in self.models:\n                temps.append(torch.sigmoid(m(x)))\n            \n        out = torch.mean(torch.stack(temps), dim=0).to('cpu').numpy()\n        return out\n\n    def add_model(self, model):\n        self.models.append(model)\n\ndef build_ensemble_model(tta):\n    model = EnsembleModel(tta)\n    for fold in range(1, CFG.n_folds+1):\n    #for fold in range(5, 6):\n        _model = build_model(CFG)\n        model_path = f'/kaggle/input/{CFG.exp_name}/{CFG.model}_fold{fold}_best.pth'\n        if torch.cuda.is_available():\n            torch.load(model_path)\n            state = torch.load(model_path)['model']\n        else:\n            state = torch.load(model_path, map_location=torch.device('cpu'))['model']\n        _model.load_state_dict(state)\n        _model = nn.DataParallel(_model)\n        _model.to(device)\n        _model.eval()\n        \n        model.add_model(_model)\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.424386Z","iopub.execute_input":"2023-06-02T02:41:29.424766Z","iopub.status.idle":"2023-06-02T02:41:29.442714Z","shell.execute_reply.started":"2023-06-02T02:41:29.424733Z","shell.execute_reply":"2023-06-02T02:41:29.441711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if mode == 'test':\n    fragment_ids = sorted(os.listdir(CFG.comp_dataset_path + mode))\nelse:\n    fragment_ids = [3]","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.445634Z","iopub.execute_input":"2023-06-02T02:41:29.446535Z","iopub.status.idle":"2023-06-02T02:41:29.457403Z","shell.execute_reply.started":"2023-06-02T02:41:29.446492Z","shell.execute_reply":"2023-06-02T02:41:29.456300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## main","metadata":{}},{"cell_type":"code","source":"QUICK_SAVE = True\n\nsample_submission_flag = filecmp.cmp(\n    \"../input/vesuvius-challenge-ink-detection/test/a/surface_volume/00.tif\",\n    \"../input/vcid-file-check/00.tif\",\n    shallow=True\n)\n\nif sample_submission_flag and QUICK_SAVE and not (IS_DEBUG):\n    df_sub = pd.read_csv(\"../input/vesuvius-challenge-ink-detection/sample_submission.csv\")\n    df_sub.to_csv(\"submission.csv\", index=False)\nelse:\n    model = build_ensemble_model(TTA)\n    results = []\n    def ink_infer_fn(loader, model, device, pixels, mask):\n        out = np.zeros_like(mask).astype(\"float\")\n        mask_count = np.zeros_like(mask)\n        mask_count += (1 - mask)\n        radius = CFG.region_size//2\n        for i, (images) in tqdm(enumerate(loader), total=len(loader), disable=not IS_DEBUG):\n            images = images.to(device)\n            batch_size = images.size(0)\n\n            with torch.no_grad():\n                preds = model(images)\n                for j, value in enumerate(preds):\n                    y, x = pixels[(i * batch_size) + j]\n                    out[y-radius:y+radius, x-radius: x+radius] += value\n                    mask_count[y-radius:y+radius, x-radius: x+radius] += np.ones((CFG.region_size, CFG.region_size))\n\n        out /= mask_count\n        out *= mask\n        return out\n    \n    for fragment_id in fragment_ids:\n        if CFG.model == \"3D-2D\" or CFG.model == \"2.5D\":\n            test_loader, xyxys = make_test_dataset(fragment_id)\n            binary_mask = cv2.imread(CFG.comp_dataset_path + f\"{mode}/{fragment_id}/mask.png\", 0)\n            binary_mask = (binary_mask / 255).astype(int)\n            ori_h = binary_mask.shape[0]\n            ori_w = binary_mask.shape[1]\n            # mask = mask / 255\n            pad0 = (CFG.tile_size - binary_mask.shape[0] % CFG.tile_size)\n            pad1 = (CFG.tile_size - binary_mask.shape[1] % CFG.tile_size)\n            binary_mask = np.pad(binary_mask, [(0, pad0), (0, pad1)], constant_values=0)\n            mask_pred = np.zeros(binary_mask.shape)\n            mask_count = np.zeros(binary_mask.shape)\n            for step, (images) in tqdm(enumerate(test_loader), total=len(test_loader), disable=not IS_DEBUG):\n                images = images.to(device)\n                batch_size = images.size(0)\n                with torch.no_grad():\n                    y_preds = model(images)\n\n                start_idx = step * CFG.valid_batch_size\n                end_idx = start_idx + batch_size\n                for i, (x1, y1, x2, y2) in enumerate(xyxys[start_idx:end_idx]):\n                    temp = y_preds[i]\n                    mask_pred[y1:y2, x1:x2] += temp.squeeze(0)\n                    mask_count[y1:y2, x1:x2] += np.ones((CFG.tile_size, CFG.tile_size))\n\n            mask_pred /= mask_count\n\n            mask_pred = mask_pred[:ori_h, :ori_w]\n            binary_mask = binary_mask[:ori_h, :ori_w]\n\n            mask_pred = (mask_pred >= TH).astype(int)\n            mask_pred *= binary_mask\n            inklabels_rle = rle(mask_pred)\n            results.append((fragment_id, inklabels_rle))\n\n            del mask_pred, mask_count\n            del test_loader\n\n            gc.collect()\n            torch.cuda.empty_cache()\n\n        elif CFG.model == \"ink-classifier\":\n            test_dir = CFG.comp_dataset_path + f\"{mode}/{fragment_id}/\"\n            volumes, pixels, mask = get_ink_data_infer(test_dir, CFG.region_size, CFG.region_stride)\n            if IS_DEBUG:\n                label_path = f\"{test_dir}inklabels.png\"\n                label = cv2.imread(label_path, 0) / 255.\n            dataset = InkDataset(volumes, _, get_transforms(data='valid', cfg=CFG), mode = \"test\")\n            loader = DataLoader(\n                dataset,\n                batch_size=CFG.valid_batch_size,\n                shuffle=False,\n                num_workers=CFG.num_workers,\n                pin_memory=True,\n                drop_last=False\n            )\n            pred = ink_infer_fn(loader, model, device, pixels, mask)\n        \n            if IS_DEBUG:\n                calc_cv(label, pred, label.shape)\n            pred = (pred >= TH).astype(int)\n            \n            inklabels_rle = rle(pred)\n            results.append((fragment_id, inklabels_rle))\n            \n            del pred, loader\n            gc.collect()\n            torch.cuda.empty_cache()\n            \n    sub = pd.DataFrame(results, columns=['Id', 'Predicted'])\n    sample_sub = pd.read_csv(CFG.comp_dataset_path + 'sample_submission.csv')\n    sample_sub = pd.merge(sample_sub[['Id']], sub, on='Id', how='left')\n    sample_sub.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-06-02T02:41:29.460150Z","iopub.execute_input":"2023-06-02T02:41:29.461128Z","iopub.status.idle":"2023-06-02T02:43:32.437816Z","shell.execute_reply.started":"2023-06-02T02:41:29.461085Z","shell.execute_reply":"2023-06-02T02:43:32.435886Z"},"trusted":true},"execution_count":null,"outputs":[]}]}