{"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":"\"\"\" \n Single script that can handle ensemble of different types of models (e.g. Unet, ResNet3D, SE-Resnet3D) and different data pipelines and generate submission file\n\"\"\"\n\n# {{{ Module Imports for all models\n\n# Specify System Paths for Kaggle Notebook\nimport 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')\nsys.path.append('/kaggle/input/segmentation-models-pytorch-v2')\nsys.path.append('/kaggle/input/resnet3d')\nsys.path.append('/kaggle/input/seresnet3d')\nsys.path.append('/kaggle/input/einops/einops-master')\n\n# Generic Imports\nimport pickle\nimport warnings\nimport pandas as pd\nimport os\nimport gc\nimport sys\nimport math\nimport time\nimport random\nimport datetime\nimport importlib\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nfrom functools import partial\nimport hashlib\n\n# Computer Vision\nimport cv2\nimport PIL.Image as Image\nimport imageio\n\n# ML modules\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\nfrom sklearn.metrics import roc_auc_score, accuracy_score, f1_score, log_loss, fbeta_score\nimport cupy as xp\nfrom einops import rearrange, reduce, repeat\n\n# Torch\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\n# TIMM\nimport timm\nfrom timm.models.resnet import resnet34d\n\n# Plotting\n#import matplotlib\n#matplotlib.use('TkAgg')\nimport matplotlib.pyplot as plt\n\n# Monitoring\nfrom tqdm.auto import tqdm\n\n# Data Augmentation\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\n# Logging\nfrom logging import getLogger, INFO, FileHandler, Formatter, StreamHandler\n\n# Model Imports\nimport segmentation_models_pytorch as smp\nfrom segmentation_models_pytorch.decoders.unet.decoder import UnetDecoder\nfrom segmentation_models_pytorch.encoders import get_encoder\nfrom resnet3d import generate_model\nfrom seresnet3d import Resnet3d\n\n# For importing pre-trained model. See https://github.com/pytorch/pytorch/issues/33288\nimport ssl\nssl._create_default_https_context = ssl._create_unverified_context\n# }}}\n\n# {{{ Core Config for all models\n\nclass Config:\n\n    # === Core Paths === \n    comp_name = 'vesuvius'\n    root_path = '/kaggle/input/'\n    comp_name = 'vesuvius-challenge-ink-detection'\n    dataset_path = f'{root_path}{comp_name}/'\n\n    target_size = 1 # target classes\n\n    num_workers = 2\n    seed = 38\n    use_tta: bool = True\n    use_denoising: bool = False\n    use_th_search = True\n\n    device_ids = [0,1]\n    #device_ids = [1]\n\n# }}}\n\n# {{{ Functions for all models\n\n# 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    # pixels = (pixels >= thr).astype(int)\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# L1/Hessian denoising\n# Reference https://www.kaggle.com/code/brettolsen/improving-performance-with-l1-hessian-denoising\n\ndelta_lookup = {\n    \"xx\": xp.array([[1, -2, 1]], dtype=float),\n    \"yy\": xp.array([[1], [-2], [1]], dtype=float),\n    \"xy\": xp.array([[1, -1], [-1, 1]], dtype=float),\n}\n\ndef operate_derivative(img_shape, pair):\n    assert len(img_shape) == 2\n    delta = delta_lookup[pair]\n    fft = xp.fft.fftn(delta, img_shape)\n    return fft * xp.conj(fft)\n\ndef soft_threshold(vector, threshold):\n    return xp.sign(vector) * xp.maximum(xp.abs(vector) - threshold, 0)\n\ndef back_diff(input_image, dim):\n    assert dim in (0, 1)\n    r, n = xp.shape(input_image)\n    size = xp.array((r, n))\n    position = xp.zeros(2, dtype=int)\n    temp1 = xp.zeros((r+1, n+1), dtype=float)\n    temp2 = xp.zeros((r+1, n+1), dtype=float)\n    \n    temp1[position[0]:size[0], position[1]:size[1]] = input_image\n    temp2[position[0]:size[0], position[1]:size[1]] = input_image\n    \n    size[dim] += 1\n    position[dim] += 1\n    temp2[position[0]:size[0], position[1]:size[1]] = input_image\n    temp1 -= temp2\n    size[dim] -= 1\n    return temp1[0:size[0], 0:size[1]]\n\ndef forward_diff(input_image, dim):\n    assert dim in (0, 1)\n    r, n = xp.shape(input_image)\n    size = xp.array((r, n))\n    position = xp.zeros(2, dtype=int)\n    temp1 = xp.zeros((r+1, n+1), dtype=float)\n    temp2 = xp.zeros((r+1, n+1), dtype=float)\n        \n    size[dim] += 1\n    position[dim] += 1\n\n    temp1[position[0]:size[0], position[1]:size[1]] = input_image\n    temp2[position[0]:size[0], position[1]:size[1]] = input_image\n    \n    size[dim] -= 1\n    temp2[0:size[0], 0:size[1]] = input_image\n    temp1 -= temp2\n    size[dim] += 1\n    return -temp1[position[0]:size[0], position[1]:size[1]]\n\ndef iter_deriv(input_image, b, scale, mu, dim1, dim2):\n    g = back_diff(forward_diff(input_image, dim1), dim2)\n    d = soft_threshold(g + b, 1 / mu)\n    b = b + (g - d)\n    L = scale * back_diff(forward_diff(d - b, dim2), dim1)\n    return L, b\n\ndef iter_xx(*args):\n    return iter_deriv(*args, dim1=1, dim2=1)\n\ndef iter_yy(*args):\n    return iter_deriv(*args, dim1=0, dim2=0)\n\ndef iter_xy(*args):\n    return iter_deriv(*args, dim1=0, dim2=1)\n\ndef iter_sparse(input_image, bsparse, scale, mu):\n    d = soft_threshold(input_image + bsparse, 1 / mu)\n    bsparse = bsparse + (input_image - d)\n    Lsparse = scale * (d - bsparse)\n    return Lsparse, bsparse\n\ndef denoise_image(input_image, iter_num=100, fidelity=150, sparsity_scale=10, continuity_scale=0.5, mu=1):\n    image_size = xp.shape(input_image)\n    #print(\"Initialize denoising\")\n    norm_array = (\n        operate_derivative(image_size, \"xx\") + \n        operate_derivative(image_size, \"yy\") + \n        2 * operate_derivative(image_size, \"xy\")\n    )\n    norm_array += (fidelity / mu) + sparsity_scale ** 2\n    b_arrays = {\n        \"xx\": xp.zeros(image_size, dtype=float),\n        \"yy\": xp.zeros(image_size, dtype=float),\n        \"xy\": xp.zeros(image_size, dtype=float),\n        \"L1\": xp.zeros(image_size, dtype=float),\n    }\n    g_update = xp.multiply(fidelity / mu, input_image)\n    for i in tqdm(range(iter_num), total=iter_num):\n        #print(f\"Starting iteration {i+1}\")\n        g_update = xp.fft.fftn(g_update)\n        if i == 0:\n            g = xp.fft.ifftn(g_update / (fidelity / mu)).real\n        else:\n            g = xp.fft.ifftn(xp.divide(g_update, norm_array)).real\n        g_update = xp.multiply((fidelity / mu), input_image)\n        \n        #print(\"XX update\")\n        L, b_arrays[\"xx\"] = iter_xx(g, b_arrays[\"xx\"], continuity_scale, mu)\n        g_update += L\n        \n        #print(\"YY update\")\n        L, b_arrays[\"yy\"] = iter_yy(g, b_arrays[\"yy\"], continuity_scale, mu)\n        g_update += L\n        \n        #print(\"XY update\")\n        L, b_arrays[\"xy\"] = iter_xy(g, b_arrays[\"xy\"], 2 * continuity_scale, mu)\n        g_update += L\n        \n        #print(\"L1 update\")\n        L, b_arrays[\"L1\"] = iter_sparse(g, b_arrays[\"L1\"], sparsity_scale, mu)\n        g_update += L\n        \n    g_update = xp.fft.fftn(g_update)\n    g = xp.fft.ifftn(xp.divide(g_update, norm_array)).real\n    \n    g[g < 0] = 0\n    g -= g.min()\n    g /= g.max()\n    return g\n\n# }}}\n\n# {{{ Dataset\n\ndef read_image_and_binary_mask(fragment_id):\n    images = []\n\n    start = cfg.start_slice\n    end = cfg.end_slice\n    idxs = range(start, end)\n\n    for i in tqdm(idxs):\n        \n        image = cv2.imread(cfg.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    binary_mask = cv2.imread(cfg.dataset_path + f\"{mode}/{fragment_id}/mask.png\", 0)\n\n    binary_mask = np.pad(binary_mask, [(0, pad0), (0, pad1)], constant_values=0)\n    binary_mask = (binary_mask / 255).astype(int)\n    \n    return images, binary_mask\n    \ndef get_transforms(cfg):\n\n    aug = A.Compose([\n        A.Resize(cfg.input_size, cfg.input_size),\n        A.Normalize(\n            mean= [0] * cfg.in_chans,\n            std= [1] * cfg.in_chans\n        ),\n        ToTensorV2(transpose_mask=True),\n    ])\n\n    return aug\n\ndef normalization(x: torch.Tensor) -> torch.Tensor:\n    \"\"\"input.shape=(batch,f1,f2,...)\"\"\"\n\n    #[batch,f1,f2]->dim[1,2]\n    dim = list(range(1, x.ndim))\n    mean = x.mean(dim = dim,keepdim = True)\n    std = x.std(dim = dim, keepdim = True)\n    return (x - mean) / (std + 1e-9)\n\nclass CustomDataset(Dataset):\n\n    def __init__(self, images, cfg, labels=None, transform=None, xys=None):\n        self.images = images\n        self.cfg = cfg\n        self.labels = labels\n        self.transform = transform\n        self.xys = xys\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        \n        if cfg.pipeline_version == '0.1':\n            image = image.astype(np.float32) / np.iinfo(image.dtype).max\n            image = image * np.iinfo(np.uint16).max\n            image = image/65535.0\n            image = (image - 0.45)/0.225\n            data = self.transform(image=image)\n            image = data['image']\n        elif cfg.pipeline_version == '0.2':\n            image = image.astype(np.float32) / np.iinfo(image.dtype).max\n            image = image * np.iinfo(np.uint16).max\n            image = image/image.max()*255\n            image = torch.from_numpy(image).permute(2,0,1).to(torch.float32) / 255\n            image[image > 0.78] = 0.78\n        else:\n            data = self.transform(image=image)\n            image = data['image']\n\n        return image, self.xys[idx]\n\ndef make_test_dataset(fragment_id):\n    test_images, binary_mask = read_image_and_binary_mask(fragment_id) # 'a' and 'b'\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\n            if binary_mask[y1:y1+cfg.tile_size, x1:x1+cfg.tile_size].max() > 0:\n                y2 = y1 + cfg.tile_size\n                x2 = x1 + cfg.tile_size\n\n                if np.all(test_images[y1:y2, x1:x2]==0):\n                    continue\n\n                test_images_list.append(test_images[y1:y2, x1:x2])\n                xyxys.append((x1, y1, x2, y2))\n\n    xyxys = np.stack(xyxys)\n            \n    test_dataset = CustomDataset(test_images_list, cfg, transform=get_transforms(cfg), xys=xyxys)\n    \n    test_loader = DataLoader(test_dataset,\n                          batch_size=cfg.batch_size,\n                          shuffle=False,\n                          num_workers=cfg.num_workers, pin_memory=True, drop_last=False)\n    \n    return test_loader, xyxys\n\ndef TTA(x: torch.Tensor, model: nn.Module):\n    #x.shape=(batch,c,h,w)\n\n    shape=x.shape\n\n    if cfg.use_chan_tta:\n\n        # Channel TTA\n        tta_chan_stride = 5\n        num_split_chans = (cfg.in_chans - cfg.z_dims) // tta_chan_stride + 1 # (25 - 20) // 5 + 1 = 2\n        if num_split_chans != 1:\n            x = [x[:, tta_chan_stride*i:cfg.z_dims + tta_chan_stride*i] for i in range(num_split_chans)]\n            x = torch.cat(x,dim=0)\n\n        # 90 degree Rotation TTA\n        x = [torch.rot90(x,k=i,dims=(-2,-1)) for i in range(4)]\n        x = torch.cat(x,dim=0)\n\n        x = model(x)\n        x = torch.sigmoid(x)\n\n        x = x.reshape(4,shape[0]*num_split_chans, *x.shape[1:])\n        x = [torch.rot90(x[i],k=-i,dims=(-2,-1)) for i in range(4)]\n        x = torch.stack(x,dim=0)\n        x = x.mean(0) # [32,224,224] \n\n        x = x.reshape(num_split_chans, shape[0], *x.shape[1:])\n        x = x.mean(0)\n\n    else:\n        x = [torch.rot90(x,k=i,dims=(-2,-1)) for i in range(4)]\n        x = torch.cat(x,dim=0)\n        x = model(x)\n        x = torch.sigmoid(x)\n        x = x.reshape(4,shape[0],*shape[2:])\n        x = [torch.rot90(x[i],k=-i,dims=(-2,-1)) for i in range(4)]\n        x = torch.stack(x,dim=0)\n        x = x.mean(0) # [32,224,224] \n        x = x.unsqueeze(1) # Make output [32,1,224,224,]\n\n    return x\n# }}}\n\n# {{{ Inference Mask\ndef create_inference_mask():\n    \"\"\"\n    It is used to mask out the edge pixels for model prediction\n\n    If is of dimension input_size x input_size with zeros on the border with edge_size\n    \"\"\"\n    mask = torch.zeros((cfg.input_size, cfg.input_size))\n    effective_pred_size = cfg.input_size - cfg.edge_size*2 \n    start = cfg.edge_size\n    end = start + cfg.input_size - cfg.edge_size*2 \n    mask[start:end, start:end] = 1\n\n    return mask\n# }}}\n\n# {{{ Model Definition\n\"\"\"\nResNet3D Model\n- Encoder is a 3D ResNet model. The architecture has been modified to remove temporal downsampling between blocks.\n- A 2D decoder is used for predicting the segmentation map.\n- The encoder feature maps are average pooled over the Z dimension before passing it to the decoder -> i.e. transform from 3D to 2D\n\"\"\"\n\nclass Decoder(nn.Module):\n    def __init__(self, encoder_dims, upscale):\n        super().__init__()\n\n        # List of Conv blocks - each block contains a 2d Conv, BN, and ReLU\n        # There are 4 blocks in total\n        self.convs = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(encoder_dims[i]+encoder_dims[i-1], encoder_dims[i-1], 3, 1, 1, bias=False),\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): # In reverse order!\n\n            # Upsample feature_map by a factor of 2 using bilinear interpolation\n            f_up = F.interpolate(feature_maps[i], scale_factor=2, mode=\"bilinear\")\n\n            # Concatenate wth feature maps of previous layers\n            f = torch.cat([feature_maps[i-1], f_up], dim=1)\n\n            # Forward pass through convolution layers \n            f_down = self.convs[i-1](f)\n\n            # Save feature maps of previous layer\n            feature_maps[i-1] = f_down\n\n        x = self.logit(feature_maps[0])\n        mask = self.up(x)\n        return mask\n\n\nclass SegModel(nn.Module):\n\n    def __init__(self, model_depth=34):\n        super().__init__()\n        if cfg.model_depth == 152:\n            self.encoder = generate_model_v2(model_depth=model_depth, n_input_channels=1)\n            self.decoder = Decoder(encoder_dims=[256, 512, 1024, 2048], upscale=4)\n        else:\n            self.encoder = generate_model(model_depth=model_depth, n_input_channels=1)\n            self.decoder = Decoder(encoder_dims=[64, 128, 256, 512], upscale=4)\n        \n    def forward(self, x):\n        if x.ndim==4:\n            x=x[:,None] # Add an extra dimension\n\n        feat_maps = self.encoder(x)\n\n        # Transform 3D to 2D by avg pooling over depth dimension\n        feat_maps_pooled = [torch.mean(f, dim=2) for f in feat_maps]\n\n        pred_mask = self.decoder(feat_maps_pooled)\n        return pred_mask\n\nclass SE_Resnet3D(nn.Module):\n\n    def __init__(self, model_depth=101, SE=True):\n        super().__init__()\n        model = Resnet3d(model_depth, SE=True)\n        self.encoder = model.encoder\n        self.decoder = model.decoder\n\n    def forward(self, x):\n\n        x = normalization(x.reshape(-1,*x.shape[2:])).reshape(x.shape)\n\n        if x.ndim==4:\n            x=x[:,None]\n\n        feat_maps = self.encoder.get_each_layer_features(x) \n        pred_mask = self.decoder.forward(feat_maps) \n\n        return pred_mask\n\n\"\"\"\nUnet Model with Attention Pooling along the z (depth) dimension\n\"\"\"\n\nclass SmpUnetDecoder(UnetDecoder):\n    \"\"\"\n    Customized Unet Decoder. But why need to reinvent the wheel?\n    \"\"\"\n\n    def __init__(self, **kwargs):\n        super(SmpUnetDecoder, self).__init__(**kwargs)\n\n    def forward(self, encoder):\n        feature = encoder[::-1]  # reverse channels to start from head of encoder\n        head = feature[0] # Obtain the bottleneck feature of the encoder (e.g. Resnet)\n        skip = feature[1:] + [None] # Features from encoder to be concatenated to the decoder features as skip connections\n        d = self.center(head) # Identity\n\n        decoder = []\n\n        # Iterate through the decoder blocks\n        for i, decoder_block in enumerate(self.blocks):\n            s = skip[i]\n            d = decoder_block(d, s) \n            decoder.append(d)\n\n        last = d\n        return last\n\nclass Unet_Attn_Pooling(nn.Module):\n\n    def __init__(self, backbone):\n        super().__init__()\n\n        self.backbone = backbone\n\n        if self.backbone == \"resnet34d\":\n\n            self.encoder_channels = [64, 64, 128, 256, 512] # Standard Resnet34 out_channels\n            self.decoder_channels = [256, 128, 64, 32, 16] \n\n            #self.encoder = resnet34d(pretrained=True, in_chans=cfg.z_dims) # Timm resnet34d as the Unet Encoder\n            self.encoder = timm.create_model(\"resnet34d\", pretrained=False, in_chans=cfg.z_dims) # Timm resnet34d as the Unet Encoder\n            self.z_offsets = [0,2,4]\n\n        elif self.backbone == \"se_resnext50_32x4d\":\n\n            self.encoder_channels = [64, 256, 512, 1024, 2048] \n            self.decoder_channels = [256, 128, 64, 32, 16] \n\n            self.z_offsets = [0,2,4]\n\n            # Get Encoder from segmentation_models_pytorch\n            self.encoder = get_encoder(\n                cfg.backbone,\n                in_channels=cfg.z_dims,\n                depth=5,\n                #weights=\"imagenet\",\n                weights=None,\n            )\n\n        elif self.backbone == \"mit_b3\":\n\n            self.encoder_channels = [0, 64, 128, 320, 512] \n            self.decoder_channels = [256, 128, 64, 32, 16] \n\n            self.z_offsets = [0,2,4,6,8]\n\n            # Get Encoder from segmentation_models_pytorch\n            self.encoder = get_encoder(\n                cfg.backbone,\n                in_channels=cfg.z_dims,\n                depth=5,\n                #weights=\"imagenet\",\n                weights=None,\n            )\n\n        self.decoder = SmpUnetDecoder(\n            encoder_channels=[0] + self.encoder_channels, # (3, 64, 64, 128, 256, 512)\n            decoder_channels=self.decoder_channels, # (256, 128, 64, 32, 16)\n            n_blocks=5,\n            use_batchnorm=True,\n            center=False,\n            attention_type=None,\n        )\n\n        self.logit = nn.Conv2d(self.decoder_channels[-1], cfg.target_size, kernel_size=1)\n\n        if self.backbone == \"mit_b3\":\n\n            # Dirty hack to avoid 0 channel buy for nn.Conv2d\n            dummy_encoder_channels = [1, 64, 128, 320, 512] \n\n            # Attention Pooling\n            self.pooling_weight = nn.ModuleList([\n                nn.Sequential(\n                    nn.Conv2d(channel, channel, kernel_size=3, padding=1),\n                    nn.ReLU(inplace=True),\n                ) for channel in dummy_encoder_channels\n            ])\n\n        else:\n\n            # Attention Pooling\n            self.pooling_weight = nn.ModuleList([\n                nn.Sequential(\n                    nn.Conv2d(channel, channel, kernel_size=3, padding=1),\n                    nn.ReLU(inplace=True),\n                ) for channel in self.encoder_channels\n            ])\n\n    def forward(self, x):\n\n        B, C, H, W = x.shape # (8, 12, 224, 224)\n        K = len(self.z_offsets)\n        x = torch.cat([x[:,i:i+cfg.z_dims,:,:] for i in self.z_offsets], 0) # shape: [24, 6, 224, 224]\n\n        if cfg.backbone == \"resnet34d\":\n\n            # Foward pass through Resnet and obtain feature maps\n            feat_maps = []\n\n            x = self.encoder.conv1(x) # [24, 64, 112, 112]\n            x = self.encoder.bn1(x)\n            x = self.encoder.act1(x)\n            feat_maps.append(x)\n\n            x = F.avg_pool2d(x, kernel_size=2, stride=2) # [24, 64, 56, 56]\n\n            x = self.encoder.layer1(x) # [24, 64, 56, 56]\n            feat_maps.append(x)\n\n            x = self.encoder.layer2(x) # [24, 128, 28, 28]\n            feat_maps.append(x)\n\n            x = self.encoder.layer3(x) # [24, 256, 14, 14]\n            feat_maps.append(x)\n            \n            x = self.encoder.layer4(x) # [24, 512, 7, 7]\n            feat_maps.append(x)\n\n        else:\n\n            feat_maps = self.encoder(x)\n\n            # Exclude the first feature\n            feat_maps = feat_maps[1:]\n\n        # Attention Pooling across slices (z)\n        for idx in range(len(feat_maps)):\n            feat = feat_maps[idx] # [24, 64, 112, 112]\n            if feat.shape[1] != 0:\n                attn_map = self.pooling_weight[idx](feat)\n                _, c, h, w = attn_map.shape # [24, 64, 112, 112]\n                attn_map = rearrange(attn_map, '(K B) c h w -> B K c h w', K=K, B=B, h=h, w=w) #f.reshape(B, K, c, h, w) # [8, 3, 64, 112, 112]\n                feat = rearrange(feat, '(K B) c h w -> B K c h w', K=K, B=B, h=h, w=w) #e.reshape(B, K, c, h, w) # [8, 3, 64, 112, 112]\n                attn_weight = F.softmax(attn_map, 1) # [8, 3, 64, 112, 112]\n                feat = (attn_weight * feat).sum(1) # [8, 64, 112, 112]\n                feat_maps[idx] = feat\n            else:\n                _, c, h, w = feat.shape # [24, 0, 112, 112]\n                feat = rearrange(feat, '(K B) c h w -> B K c h w', K=K, B=B, h=h, w=w) #e.reshape(B, K, c, h, w) # [8, 3, 0, 112, 112]\n                feat_maps[idx] = feat.sum(1) # [8, 0, 112, 112]\n\n        top_feat = self.decoder(feat_maps) # [8, 16, 224, 224]\n\n        logit = self.logit(top_feat) # [8, 1, 224, 224]\n\n        return logit\n\nclass CustomModel(nn.Module):\n    def __init__(self, cfg, weight=None):\n        super().__init__()\n        self.cfg = cfg\n\n        if cfg.use_attn_pooling:\n            self.encoder = Unet_Attn_Pooling(cfg.backbone)\n\n        elif cfg.backbone == \"resnet3d\":\n            self.encoder = SegModel(model_depth=cfg.model_depth)\n\n        elif cfg.backbone[:3] == \"mit\":\n            self.encoder = smp.Unet(\n                encoder_name=cfg.backbone, \n                encoder_weights=weight,\n                classes=cfg.target_size,\n                activation=None,\n            )\n            #if cfg.in_chans == 6:\n            #    out_channels=self.encoder.encoder.patch_embed1.proj.out_channels\n            #    self.encoder.encoder.patch_embed1.proj=nn.Conv2d(cfg.in_chans,out_channels,7,4,3)\n        else:\n            self.encoder = smp.Unet(\n                encoder_name=cfg.backbone, \n                encoder_weights=weight,\n                in_channels=cfg.in_chans, # 6 -> middle 6 slices out of a total of 65\n                classes=cfg.target_size,\n                activation=None,\n            )\n\n    def forward(self, image):\n\n        if cfg.backbone[:3] == \"mit\" and cfg.in_chans == 6:\n            input_1, input_2 = image.split(3, dim=1)\n            output_1 = self.encoder(input_1)\n            output_2 = self.encoder(input_2)\n            output = (output_1 + output_2) / 2\n        else:\n            output = self.encoder(image)\n            output = output.squeeze(-1)\n\n        return output\n    \n\ndef build_model(cfg, model_path):\n    print('model_name', cfg.model_name)\n    print('backbone', cfg.backbone)\n\n    if cfg.pipeline_version == '0.1':\n        model = SegModel(cfg.model_depth)\n    elif cfg.pipeline_version == '0.2':\n        model = SE_Resnet3D(model_depth=101, SE=True)\n    else:\n        model = CustomModel(cfg)\n\n    print(f'Load model from path: {model_path}')\n    if cfg.pipeline_version == '0.2':\n        model.load_state_dict(torch.load(model_path,map_location=\"cpu\"))\n    else:\n        model.load_state_dict(torch.load(model_path,map_location=\"cpu\")['model'])\n    \n    return model\n\nclass EnsembleModel(nn.Module):\n    def __init__(self, weights=None):\n        super().__init__()\n        self.models = nn.ModuleList()\n        self.weights = torch.tensor(weights)[:,None,None,None,None]\n\n    def __call__(self, x):\n        \"\"\"\n        Weighted average of models\n        \"\"\"\n        x = [model(x) for model in self.models]\n        x = torch.stack(x, dim=0)\n        x = x * self.weights.to(device=x.device)\n        return torch.sum(x, dim=0)\n\n    def add_model(self, model):\n        self.models.append(model)\n\ndef build_ensemble_model(cfg, model_path, cv_folds, weights):\n\n    model = EnsembleModel(weights)\n\n    for fold in cv_folds:\n        path = model_path.format(fold=fold, backbone=cfg.backbone)\n        _model = build_model(cfg, path)\n        model.add_model(_model)\n    \n    return model\n# }}}\n\n# Main \n\n# Global Config\ncfg = Config\n\nmode = 'test'\nfragment_ids = sorted(os.listdir(cfg.dataset_path + mode)) # ['a','b']\n\n#device = torch.device('cuda:1' if torch.cuda.is_available() else 'cpu')\n\nmodel_template_list = [{\n    'alias': 'seresnet3d-101',\n    'weight': 0.45,\n    'config': {\n        'model_name': 'seresnet3d_101',\n        'backbone': 'se_resnet3d_101',\n        'in_chans': 25,\n        'z_dims': 20,\n        'start_slice': 15,\n        'end_slice': 40,\n        'input_size': 256,\n        'tile_size': 256,\n        'edge_size': 32, # Do not predict on edge pixels\n        'stride': (256-32*2) // 4, # 256 - 32*2 = 192\n        'batch_size': 8,\n        'pipeline_version': '0.2',\n        'inference_mode': 'eval',\n        'use_attn_pooling': False,\n        'use_chan_tta': True,\n        'use_inference_mask': True,\n        'model': {\n            'ensemble': False,\n            'model_path': '/kaggle/input/se-resnet101-737-model/SE-resnet3d-101_epoch_14_tl0.345_vl0.335_vp0.737.pth'\n        }\n    }\n}, {\n    'alias': 'resnet3d',\n    'weight': 0.2,\n    'config': {\n        'model_name': 'resnet3d_224',\n        'backbone': 'resnet3d',\n        'in_chans': 18,\n        'start_slice': 22,\n        'end_slice': 40,\n        'model_depth': 34,\n        'input_size': 224,\n        'tile_size': 224,\n        'edge_size': 16,\n        'stride': (224-16*2) // 2, \n        'batch_size': 8, \n        'pipeline_version': '1.0',\n        'inference_mode': 'train',\n        'use_attn_pooling': False,\n        'use_chan_tta': False,\n        'use_inference_mask': True,\n        'model': {\n            'ensemble': True,\n            'folds': [1,3,4,5],\n            'weights': [0.3,0.3,0.2,0.2],\n            'model_path': '/kaggle/input/resnet34-224-models/resnet34-224-models/Unet_resnet3d_fold{fold}.pth'\n        }\n    }\n}, {\n    'alias': 'resnet3d',\n    'weight': 0.1,\n    'config': {\n        'model_name': 'resnet3d_192',\n        'backbone': 'resnet3d',\n        'in_chans': 18,\n        'start_slice': 22,\n        'end_slice': 40,\n        'model_depth': 34,\n        'input_size': 192,\n        'tile_size': 192,\n        'edge_size': 16,\n        'stride': (192-16*2) // 2,\n        'batch_size': 8,\n        'pipeline_version': '1.0',\n        'inference_mode': 'train',\n        'use_attn_pooling': False,\n        'use_chan_tta': False,\n        'use_inference_mask': True,\n        'model': {\n            'ensemble': True,\n            'folds': [1,2,3],\n            'weights': [0.4,0.3,0.3],\n            'model_path': '/kaggle/input/resnet34-192-models/resnet34-192-models/Unet_resnet3d_fold{fold}.pth'\n        }\n    }\n}, {\n    'alias': 'mit-b3',\n    'weight': 0.2,\n    'config': {\n        'model_name': 'mit_b3_attnPool',\n        'backbone': 'mit_b3',\n        'in_chans': 11,\n        'z_dims': 3,\n        'start_slice': 25,\n        'end_slice': 36,\n        'input_size': 224,\n        'tile_size': 224,\n        'edge_size': 32,\n        'stride': (224-32*2) // 2,\n        'batch_size': 8,\n        'pipeline_version': '1.0',\n        'inference_mode': 'eval',\n        'use_attn_pooling': True,\n        'use_chan_tta': False,\n        'use_inference_mask': True,\n        'model': {\n            'ensemble': True,\n            'folds': [1,3,4,5],\n            'weights': [0.3,0.3,0.2,0.2],\n            'model_path': '/kaggle/input/mit-b3-attnpool-exp048-models/mit-b3-attnPool-exp048-models/Unet_mit_b3_fold{fold}.pth'\n        }\n    }\n}, {\n    'alias': 'unet',\n    'weight': 0.15,\n    'config': {\n        'model_name': 'mit_b3_6chans',\n        'backbone': 'mit_b3',\n        'in_chans': 6,\n        'start_slice': 29,\n        'end_slice': 35,\n        'input_size': 224,\n        'tile_size': 224,\n        'edge_size': 24,\n        'stride': (224-24*2) // 2,\n        'batch_size': 8,\n        'pipeline_version': '1.0',\n        'inference_mode': 'eval',\n        'use_attn_pooling': False,\n        'use_chan_tta': False,\n        'use_inference_mask': True,\n        'model': {\n            'ensemble': True,\n            'folds': [1,2,3],\n            'weights': [0.4,0.3,0.3],\n            'model_path': '/kaggle/input/mit-b3-models/mit-b3-models/Unet_mit_b3_fold{fold}.pth'\n        }\n    }\n}]\n#model_template_list.append(model_template_list.pop(0))\n\n\npred_masks_dict = defaultdict(list)\n\nfixed_TH = 0.5 # Arbitrary confidence threshold\n\nensemble_weights = []\n\na_file = cfg.dataset_path + f\"test/a/mask.png\"\nwith open(a_file,'rb') as f:\n    hash_md5 = hashlib.md5(f.read()).hexdigest()\nis_skip_test = hash_md5 == '0b0fffdc0e88be226673846a143bb3e0'\n\nis_skip_test = False\n\nif is_skip_test:\n    submit_df = pd.DataFrame({\n        'Id': ['a', 'b'],\n        'Predicted':['1 2', '1 2']\n    })\n    submit_df.to_csv('submission.csv', index=False)\n\nelse:\n\n    pred_paths_dict = defaultdict(list) # Store prediction on disk\n    for model_template in model_template_list:\n\n        # Register model specific config\n        for k, v in model_template['config'].items():\n            setattr(cfg,k,v)\n\n        # Initialize model\n        model_path = model_template['config']['model']['model_path']\n        if model_template['config']['model']['ensemble']:\n            folds = model_template['config']['model']['folds']\n            weights = model_template['config']['model']['weights']\n            model = build_ensemble_model(cfg, model_path, folds, weights)\n        else:\n            model = build_model(cfg, model_path)\n            \n        model = nn.DataParallel(model, device_ids=Config.device_ids)\n        model = model.cuda()\n        if cfg.inference_mode == 'eval':\n            model.eval()\n\n        # Register ensemnble weights\n        ensemble_weights.append(model_template['weight'])\n\n        for fragment_id in fragment_ids:\n\n            test_loader, xyxys = make_test_dataset(fragment_id)\n\n            # Load mask for test fragment\n            binary_mask = cv2.imread(cfg.dataset_path + f\"{mode}/{fragment_id}/mask.png\", 0)\n            binary_mask = (binary_mask / 255).astype(int)\n\n            # Save height and weight for the original mask prior to padding\n            ori_h = binary_mask.shape[0]\n            ori_w = binary_mask.shape[1]\n\n            # Pad mask such that it is divisible by 224 which is the patch size\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\n            binary_mask = np.pad(binary_mask, [(0, pad0), (0, pad1)], constant_values=0)\n            \n            mask_pred = torch.zeros(binary_mask.shape,device=\"cuda:1\")\n            mask_count = torch.zeros(binary_mask.shape,device=\"cuda:1\")\n            if cfg.use_inference_mask:\n                inference_mask = create_inference_mask().to(\"cuda:1\")\n\n            for step, (images, xys) in tqdm(enumerate(test_loader), total=len(test_loader)):\n                images = images.cuda()\n\n                with torch.no_grad():\n                    with autocast():\n                        if cfg.use_tta:\n                            y_preds = TTA(images, model)\n                        else:\n                            y_preds = model(images)\n                            y_preds = torch.sigmoid(y_preds)\n                \n                if cfg.use_inference_mask:\n                    y_preds = (y_preds.to(\"cuda:1\").squeeze(1) * inference_mask[None]) # shape=(batch,H,W)\n                else:\n                    y_preds = (y_preds.to(\"cuda:1\").squeeze(1))\n\n                for k, (x1, y1, x2, y2) in enumerate(xys):\n                    if cfg.use_inference_mask:\n                        mask_pred[y1:y2, x1:x2] += y_preds[k]\n                        mask_count[y1:y2, x1:x2] += inference_mask\n                    else:\n                        mask_pred[y1:y2, x1:x2] += y_preds[k]\n                        mask_count[y1:y2, x1:x2] += 1\n\n            print(f'mask_count_min: {mask_count.min()}')\n            mask_pred /= (mask_count + 1e-7)\n\n            # Cut out the region with original height and width\n            mask_pred = mask_pred[:ori_h, :ori_w]\n            binary_mask = binary_mask[:ori_h, :ori_w]\n\n            mask_pred *= torch.from_numpy(binary_mask).to(\"cuda:1\")\n            mask_pred = mask_pred.cpu().numpy()\n            \n            if cfg.use_th_search:\n\n                # Store the prediction on disk for thresholding computation in the next stage\n                mask_pred = (mask_pred*255).astype(np.uint8) \n                path = f\"_{cfg.model_name}_{fragment_id}\"+\".png\"\n                cv2.imwrite(path, mask_pred, [cv2.IMWRITE_PNG_COMPRESSION,3])\n                pred_paths_dict[cfg.model_name].append(path)\n\n            else:\n\n                pred_masks_dict[fragment_id].append(mask_pred)\n\n                fig, axes = plt.subplots(1, 3, figsize=(15, 8))\n                axes[0].imshow(mask_count[:ori_h, :ori_w].cpu())\n                axes[1].imshow(mask_pred)\n\n                mask_pred = (mask_pred >= fixed_TH).astype(int)\n\n                axes[2].imshow(mask_pred)\n                plt.show()\n                #plt.savefig(f\"pred_{cfg.model_name}_{fragment_id}_fixed_th.png\")\n\n                plt.clf()\n                fig.clear()\n                plt.close(fig)\n\n            del mask_pred, mask_count, binary_mask, xyxys, test_loader, images, y_preds\n            gc.collect()\n            torch.cuda.empty_cache()\n\n        del model\n        gc.collect()\n        torch.cuda.empty_cache()\n\n    # Generate Submission\n\n    if cfg.use_th_search:\n\n        total_weight = np.sum(ensemble_weights)\n\n        print('Start averaging pred')\n\n        # Read prediction stored on disk\n        pred_outputs = []\n        for model_idx, (model_name, pred_paths) in enumerate(pred_paths_dict.items()):\n            print(f\"Reading prediction file from {model_name}\")\n            for fragment_idx, pred_path in enumerate(pred_paths):\n                pred_img = cv2.imread(pred_path, 0).astype(np.uint16)\n\n                if model_idx == 0:\n                    pred_outputs.append(pred_img * ensemble_weights[0] / total_weight)\n                else:\n                    pred_outputs[fragment_idx] += pred_img * ensemble_weights[model_idx] / total_weight\n\n        # Normalize averaged outputs\n\n        print('Start normalizing pred')\n\n        for fragment_idx, pred_output in enumerate(pred_outputs): \n            shape = pred_output.shape\n            pred_output = pred_output.flatten()\n            cache = np.zeros(pred_output.shape[0], dtype=np.uint16)\n            cache[pred_output.argsort()] = (np.arange(pred_output.shape[0]) / pred_output.shape[0]*65535.0).astype(np.uint16)\n            pred_output = cache.reshape(shape)\n            pred_outputs[fragment_idx] = pred_output\n\n            del cache\n            gc.collect()\n\n        print('start th search')\n\n        # Search for Threshold\n        th_percentile = 0.03\n\n        TH = [output.flatten() for output in pred_outputs] \n        TH = np.concatenate(TH)\n        TH.sort()\n        TH:float = TH[-int(len(TH)*th_percentile)]\n\n        results = []\n        for fragment_id, mask_pred in zip(fragment_ids, pred_outputs):\n\n            fig, axes = plt.subplots(1, 2, figsize=(15, 8))\n            axes[0].imshow(mask_pred)\n        \n            mask_pred = (mask_pred >= TH).astype(np.uint8)\n            results.append((fragment_id, rle(mask_pred)))\n\n            axes[1].imshow(mask_pred)\n            plt.show()\n            #plt.savefig(f\"pred_ensemble_{fragment_id}_dynamic_th.png\")\n\n            plt.clf()\n            fig.clear()\n            plt.close(fig)\n            \n        del pred_outputs\n        gc.collect()\n        torch.cuda.empty_cache()\n\n\n        print('done')\n\n    else:\n\n        results = []\n\n        for fragment_id, pred_list in pred_masks_dict.items():\n\n            avg_pred = np.average(pred_list, axis=0, weights=ensemble_weights)\n\n            fig, axes = plt.subplots(1, 2, figsize=(15, 8))\n            axes[0].imshow(avg_pred)\n\n            avg_pred = (avg_pred >= fixed_TH).astype(int)\n\n            axes[1].imshow(avg_pred)\n            plt.show()\n            #plt.savefig(f\"pred_ensemble_{fragment_id}_fixed_th.png\")\n\n            inklabels_rle = rle(avg_pred)\n\n            results.append((fragment_id, inklabels_rle))\n\n            del avg_pred\n            gc.collect()\n            torch.cuda.empty_cache()\n            \n        del pred_masks_dict\n        gc.collect()\n        torch.cuda.empty_cache()\n\n    sub = pd.DataFrame(results, columns=['Id', 'Predicted'])\n\n    sample_sub = pd.read_csv(cfg.dataset_path + 'sample_submission.csv')\n    sample_sub = pd.merge(sample_sub[['Id']], sub, on='Id', how='left')\n\n    sample_sub.to_csv(\"submission.csv\", index=False)","metadata":{"_uuid":"ad439d37-e7f0-4a86-a5be-bfde770976d8","_cell_guid":"1b0bdeb9-05ac-48eb-b861-1ef7f427aeb7","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-14T14:40:31.278393Z","iopub.execute_input":"2023-06-14T14:40:31.278759Z","iopub.status.idle":"2023-06-14T14:41:51.298597Z","shell.execute_reply.started":"2023-06-14T14:40:31.278724Z","shell.execute_reply":"2023-06-14T14:41:51.295203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}